Skip to content

wan_s2v

Wan2.2-S2V-14B: audio-driven video generation.

Ported from the official native implementation (Wan-Video/Wan2.2, wan/modules/s2v/). There is no diffusers implementation of S2V -- the PR that would have added one (huggingface/diffusers#12258) was closed unmerged -- so the official repo is the only reference, and the checkpoint ships in native Wan naming (blocks.0.self_attn.q) rather than diffusers naming.

Three things make S2V structurally different from every other Wan variant here, and none of them can be expressed by configuring WanTransformer3DModel:

  1. Heterogeneous sequence. Tokens are [video | reference image | motion]. Only the leading video_len video tokens are denoised and returned; the rest are conditioning context carried through the tower.
  2. Two-segment modulation. With zero_timestep, video tokens are modulated by the real timestep while ref/motion tokens are modulated by a fixed zero timestep -- they are already clean, so they must not be treated as noisy. Every modulation site in the block therefore splits at seg_idx.
  3. Precomputed heterogeneous RoPE. Each span gets its own positional range (motion frames sit at negative time offsets); frequencies are built once for the whole sequence rather than derived from a single grid.

Sequence parallelism is not wired up in this first version (upstream shards pre_compute_freqs alongside the hidden states); single-GPU and FSDP-style weight sharding work. See the S2V section of the support matrix.

Classes

fastvideo.models.dits.wan_s2v.FramePackMotioner

FramePackMotioner(inner_dim: int = 5120, num_heads: int = 40, zip_frame_buckets: tuple[int, int, int] = (1, 2, 16), drop_mode: str = 'drop')

Bases: Module

Compress past motion frames at three temporal scales.

Recent frames keep full detail (proj), older ones are downsampled 2x (proj_2x) and 4x (proj_4x) -- more history for fewer tokens, the same trade a video codec makes. Buckets are [nearest, mid, farthest].

Source code in fastvideo/models/dits/wan_s2v.py
def __init__(self, inner_dim: int = 5120, num_heads: int = 40,
             zip_frame_buckets: tuple[int, int, int] = (1, 2, 16),
             drop_mode: str = "drop") -> None:
    super().__init__()
    assert inner_dim % num_heads == 0 and (inner_dim // num_heads) % 2 == 0
    self.proj = nn.Conv3d(16, inner_dim, kernel_size=(1, 2, 2), stride=(1, 2, 2))
    self.proj_2x = nn.Conv3d(16, inner_dim, kernel_size=(2, 4, 4), stride=(2, 4, 4))
    self.proj_4x = nn.Conv3d(16, inner_dim, kernel_size=(4, 8, 8), stride=(4, 8, 8))
    # Plain ints, not a tensor: the loader builds this model under
    # torch.device("meta"), where any tensor attribute becomes unreadable
    # (.item()/.sum() raise) and only registered buffers can be rematerialised.
    self.zip_frame_buckets = tuple(zip_frame_buckets)
    self.inner_dim = inner_dim
    self.num_heads = num_heads
    self.drop_mode = drop_mode
    # Non-persistent: derived from config, never stored in the checkpoint.
    self.register_buffer("freqs", rope_freqs(inner_dim // num_heads), persistent=False)

fastvideo.models.dits.wan_s2v.HeadS2V

HeadS2V(dim: int, out_dim: int, patch_size: tuple[int, int, int], eps: float = 1e-06)

Bases: Module

Final norm + projection back to patch space.

Source code in fastvideo/models/dits/wan_s2v.py
def __init__(self, dim: int, out_dim: int, patch_size: tuple[int, int, int], eps: float = 1e-6) -> None:
    super().__init__()
    self.patch_size = patch_size
    self.norm = FP32LayerNorm(dim, eps, elementwise_affine=False)
    self.head = nn.Linear(dim, math.prod(patch_size) * out_dim)
    self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)

fastvideo.models.dits.wan_s2v.WanS2VAttentionBlock

WanS2VAttentionBlock(dim: int, ffn_dim: int, num_heads: int, qk_norm: bool = True, cross_attn_norm: bool = True, eps: float = 1e-06)

Bases: Module

One of the 40 blocks.

Cannot inherit WanTransformerBlock: every modulation site here is segment-aware (see the module docstring).

Source code in fastvideo/models/dits/wan_s2v.py
def __init__(self, dim: int, ffn_dim: int, num_heads: int, qk_norm: bool = True,
             cross_attn_norm: bool = True, eps: float = 1e-6) -> None:
    super().__init__()
    self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
    self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=False)
    # fp32 like norm1/norm2: upstream's WanLayerNorm computes in fp32 and casts
    # back, and FastVideo's own Wan port does the same for this norm.
    self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
    self.self_attn = WanS2VSelfAttention(dim, num_heads, qk_norm, eps)
    self.cross_attn = WanS2VCrossAttention(dim, num_heads, qk_norm, eps)
    self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh", bias=True)
    # Native Wan calls this "modulation"; diffusers calls it scale_shift_table.
    self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)

fastvideo.models.dits.wan_s2v.WanS2VCrossAttention

WanS2VCrossAttention(dim: int, num_heads: int, qk_norm: bool = True, eps: float = 1e-06)

Bases: WanS2VSelfAttention

Text cross-attention: same projections as self-attention, no RoPE.

Source code in fastvideo/models/dits/wan_s2v.py
def __init__(self, dim: int, num_heads: int, qk_norm: bool = True, eps: float = 1e-6) -> None:
    super().__init__()
    assert dim % num_heads == 0
    self.num_heads = num_heads
    self.head_dim = dim // num_heads
    self.to_q = ReplicatedLinear(dim, dim)
    self.to_k = ReplicatedLinear(dim, dim)
    self.to_v = ReplicatedLinear(dim, dim)
    self.to_out = ReplicatedLinear(dim, dim)
    self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
    self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
    self.attn = LocalAttention(num_heads=num_heads, head_size=self.head_dim, dropout_rate=0,
                               softmax_scale=None, causal=False,
                               supported_attention_backends=S2V_ATTENTION_BACKENDS)

fastvideo.models.dits.wan_s2v.WanS2VSelfAttention

WanS2VSelfAttention(dim: int, num_heads: int, qk_norm: bool = True, eps: float = 1e-06)

Bases: Module

Self-attention over the full heterogeneous sequence, with precomputed RoPE.

Source code in fastvideo/models/dits/wan_s2v.py
def __init__(self, dim: int, num_heads: int, qk_norm: bool = True, eps: float = 1e-6) -> None:
    super().__init__()
    assert dim % num_heads == 0
    self.num_heads = num_heads
    self.head_dim = dim // num_heads
    self.to_q = ReplicatedLinear(dim, dim)
    self.to_k = ReplicatedLinear(dim, dim)
    self.to_v = ReplicatedLinear(dim, dim)
    self.to_out = ReplicatedLinear(dim, dim)
    self.norm_q = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
    self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
    self.attn = LocalAttention(num_heads=num_heads, head_size=self.head_dim, dropout_rate=0,
                               softmax_scale=None, causal=False,
                               supported_attention_backends=S2V_ATTENTION_BACKENDS)

fastvideo.models.dits.wan_s2v.WanS2VTransformer3DModel

WanS2VTransformer3DModel(config: WanS2VConfig, hf_config: dict[str, Any], **kwargs)

Bases: BaseDiT

Wan2.2-S2V-14B transformer.

Source code in fastvideo/models/dits/wan_s2v.py
def __init__(self, config: WanS2VConfig, hf_config: dict[str, Any], **kwargs) -> None:
    super().__init__(config=config, hf_config=hf_config)
    arch = config.arch_config
    dim = arch.hidden_size
    self.hidden_size = dim
    self.num_attention_heads = arch.num_attention_heads
    self.rope_max_seq_len = arch.rope_max_seq_len
    self.num_channels_latents = arch.num_channels_latents
    self.out_channels = arch.out_channels
    self.patch_size = arch.patch_size
    self.text_len = arch.text_len
    self.freq_dim = arch.freq_dim
    self.zero_timestep = arch.zero_timestep
    self.enable_adain = arch.enable_adain
    self.adain_mode = arch.adain_mode
    self.add_last_motion = arch.add_last_motion

    self.patch_embedding = nn.Conv3d(arch.in_channels, dim, kernel_size=arch.patch_size, stride=arch.patch_size)
    self.text_embedding = MLP(arch.text_dim, dim, output_dim=dim, act_type="gelu_pytorch_tanh", bias=True)
    self.time_embedding = MLP(arch.freq_dim, dim, output_dim=dim, act_type="silu", bias=True)
    # Named (not indexed) so the child is `time_projection.linear`, which is
    # what param_names_mapping targets; nn.Sequential would name it ".1".
    self.time_projection = nn.Sequential(
        OrderedDict([("act", nn.SiLU()), ("linear", ReplicatedLinear(dim, dim * 6))]))

    self.blocks = nn.ModuleList([
        WanS2VAttentionBlock(dim, arch.ffn_dim, arch.num_attention_heads, qk_norm=True,
                             cross_attn_norm=arch.cross_attn_norm, eps=arch.eps)
        for _ in range(arch.num_layers)
    ])
    self.head = HeadS2V(dim, arch.out_channels, arch.patch_size, arch.eps)

    self.cond_encoder = None
    if arch.cond_dim > 0:
        self.cond_encoder = nn.Conv3d(arch.cond_dim, dim, kernel_size=arch.patch_size, stride=arch.patch_size)
    self.casual_audio_encoder = CausalAudioEncoder(
        dim=arch.audio_dim, num_layers=25, out_dim=dim, num_token=arch.num_audio_token,
        need_global=arch.enable_adain)
    self.audio_injector = AudioInjector(
        dim=dim, num_heads=arch.num_attention_heads, inject_layers=arch.audio_inject_layers,
        enable_adain=arch.enable_adain, adain_dim=dim, eps=arch.eps)
    # 3 token kinds: 0 = noisy video, 1 = reference image, 2 = motion.
    self.trainable_cond_mask = nn.Embedding(3, dim)
    if arch.enable_framepack:
        self.frame_packer = FramePackMotioner(
            inner_dim=dim, num_heads=arch.num_attention_heads, zip_frame_buckets=(1, 2, 16),
            drop_mode=arch.framepack_drop_mode)

    # Non-persistent: derived from config, absent from the checkpoint.
    self.register_buffer("freqs", rope_freqs(dim // arch.num_attention_heads, arch.rope_max_seq_len),
                         persistent=False)

Methods:

fastvideo.models.dits.wan_s2v.WanS2VTransformer3DModel.forward
forward(hidden_states: Tensor | list[Tensor], encoder_hidden_states: Tensor | list[Tensor], timestep: Tensor, ref_latents: Tensor | list[Tensor] | None = None, motion_latents: Tensor | list[Tensor] | None = None, cond_states: Tensor | list[Tensor] | None = None, audio_input: Tensor | None = None, motion_frames: tuple[int, int] = (17, 5), add_last_motion: int = 2, drop_motion_frames: bool = False, **kwargs) -> list[Tensor]

Denoise one step of an audio-driven video.

hidden_states [B, C, T, H, W] or list of [C, T, H, W] noisy video latents ref_latents reference-image latents; required, the model is image-conditioned motion_latents previously generated frames; None on the first clip cond_states pose/control latents, or None when unused audio_input [B, 25, C_a, T_a] stacked wav2vec2 hidden states; required

Returns a batched [B, C, T, H, W] fp32 tensor when hidden_states came in batched (the DenoisingStage path -- CFG arithmetic and scheduler.step need a tensor), or a list of per-sample tensors when it came in as a list (the reference-implementation path).

Source code in fastvideo/models/dits/wan_s2v.py
def forward(self, hidden_states: torch.Tensor | list[torch.Tensor],
            encoder_hidden_states: torch.Tensor | list[torch.Tensor], timestep: torch.Tensor,
            ref_latents: torch.Tensor | list[torch.Tensor] | None = None,
            motion_latents: torch.Tensor | list[torch.Tensor] | None = None,
            cond_states: torch.Tensor | list[torch.Tensor] | None = None,
            audio_input: torch.Tensor | None = None,
            motion_frames: tuple[int, int] = (17, 5), add_last_motion: int = 2,
            drop_motion_frames: bool = False, **kwargs) -> list[torch.Tensor]:
    """Denoise one step of an audio-driven video.

    hidden_states  [B, C, T, H, W] or list of [C, T, H, W] noisy video latents
    ref_latents    reference-image latents; required, the model is image-conditioned
    motion_latents previously generated frames; None on the first clip
    cond_states    pose/control latents, or None when unused
    audio_input    [B, 25, C_a, T_a] stacked wav2vec2 hidden states; required

    Returns a batched [B, C, T, H, W] fp32 tensor when ``hidden_states`` came
    in batched (the DenoisingStage path -- CFG arithmetic and scheduler.step
    need a tensor), or a list of per-sample tensors when it came in as a
    list (the reference-implementation path).
    """
    batched_input = isinstance(hidden_states, torch.Tensor)
    hidden_states = self._as_sample_list(hidden_states)
    ref_latents = self._as_sample_list(ref_latents)
    motion_latents = self._as_sample_list(motion_latents)
    cond_states = self._as_sample_list(cond_states)
    if isinstance(encoder_hidden_states, torch.Tensor):
        encoder_hidden_states = list(encoder_hidden_states)
    elif isinstance(encoder_hidden_states, list) and encoder_hidden_states and \
            encoder_hidden_states[0].dim() == 3:
        # DenoisingStage passes a list with one batched [B, L, C] tensor per
        # text encoder (same convention wanvideo.py unwraps with [0]).
        encoder_hidden_states = list(encoder_hidden_states[0])
    if audio_input is None or ref_latents is None:
        raise ValueError(
            "Wan-S2V needs both a reference image and audio. Got "
            f"ref_latents={'set' if ref_latents is not None else 'None'}, "
            f"audio_input={'set' if audio_input is not None else 'None'}. The pipeline supplies "
            "these from batch.image_latent and batch.audio_embeds.")
    if motion_latents is None:
        # First clip: no history yet. The official runner drops motion tokens
        # here (drop_first_motion=True in its config); the zeros only exist so
        # the frame packer has an input to (then discard) -- their values are
        # never attended to.
        drop_motion_frames = True
        motion_latents = [torch.zeros_like(u[:, :1]) for u in hidden_states]
    if cond_states is None and self.cond_encoder is not None:
        # The reference always encodes a cond tensor (zeros when unused), and
        # cond_encoder has a bias -- skipping it entirely would shift every
        # video token relative to the official implementation.
        cond_states = [torch.zeros_like(u) for u in hidden_states]

    # Every conditioning input above is batch-1 by construction (one
    # reference image, one motion history, one audio track) while
    # ``hidden_states`` carries ``num_videos_per_prompt`` samples. Repeat them
    # so each sample is conditioned: the zips below would otherwise truncate
    # to the shorter list and silently collapse the whole batch to 1, while
    # ``context``/``audio_emb`` keep their own batch and then mismatch.
    n_samples = len(hidden_states)
    ref_latents = self._repeat_batch1(ref_latents, n_samples)
    motion_latents = self._repeat_batch1(motion_latents, n_samples)
    if cond_states is not None:
        cond_states = self._repeat_batch1(cond_states, n_samples)
    if audio_input.size(0) == 1 and n_samples > 1:
        audio_input = audio_input.repeat(n_samples, 1, 1, 1)

    add_last_motion = int(self.add_last_motion) * add_last_motion
    audio_emb, audio_emb_global = self._embed_audio(audio_input, motion_frames)
    freqs = self.freqs.to(self.patch_embedding.weight.device)

    # 1. video tokens (+ pose conditioning added in patch space)
    x = [self.patch_embedding(u.unsqueeze(0)) for u in hidden_states]
    if cond_states is not None and self.cond_encoder is not None:
        x = [x_ + self.cond_encoder(c.unsqueeze(0)) for x_, c in zip(x, cond_states, strict=True)]
    original_grid_sizes = torch.stack([torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
    x = [u.flatten(2).transpose(1, 2) for u in x]
    video_len = x[0].size(1)
    grid_sizes = [[torch.zeros_like(original_grid_sizes), original_grid_sizes, original_grid_sizes]]

    # 2. reference-image tokens, parked far away in time (see REF_TIME_INDEX)
    ref = [self.patch_embedding(r.unsqueeze(0)) for r in ref_latents]
    bsz, h, w = len(ref), ref[0].shape[3], ref[0].shape[4]
    grid_sizes.append([
        torch.tensor([REF_TIME_INDEX, 0, 0]).unsqueeze(0).repeat(bsz, 1),
        torch.tensor([REF_TIME_INDEX + 1, h, w]).unsqueeze(0).repeat(bsz, 1),
        torch.tensor([1, h, w]).unsqueeze(0).repeat(bsz, 1),
    ])
    x = [torch.cat([u, r.flatten(2).transpose(1, 2)], dim=1) for u, r in zip(x, ref, strict=True)]

    # 3. token-kind mask: 0 video, 1 reference (2 = motion, tagged on append)
    mask = [torch.zeros([1, u.shape[1]], dtype=torch.long, device=u.device) for u in x]
    for m in mask:
        m[:, video_len:] = 1

    # 4. RoPE for video+ref, then append motion tokens (which carry their own)
    stacked = torch.cat(x)
    rope = rope_precompute(
        stacked.detach().view(stacked.size(0), stacked.size(1), self.num_attention_heads,
                              self.hidden_size // self.num_attention_heads), grid_sizes, freqs)
    x, rope = [u.unsqueeze(0) for u in stacked], [u.unsqueeze(0) for u in rope]
    x, rope, mask = self._inject_motion(x, rope, mask, motion_latents, drop_motion_frames, add_last_motion)

    x = torch.cat(x, dim=0)
    rope = torch.cat(rope, dim=0)
    x = x + self.trainable_cond_mask(torch.cat(mask, dim=0)).to(x.dtype)

    # 5. conditioning embeddings
    e, block_e = self._timestep_embedding(timestep, video_len)
    context = self.text_embedding(
        torch.stack([
            torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in encoder_hidden_states
        ]))

    # 6. the tower, with audio re-injected at 12 of the 40 blocks
    for idx, block in enumerate(self.blocks):
        x = block(x, block_e, context, rope)
        x = self.audio_injector(x, idx, audio_emb, audio_emb_global, video_len, self.adain_mode)

    # 7. only the video span is denoised output; ref/motion are context
    out = [u.float() for u in self.unpatchify(self.head(x[:, :video_len], e), original_grid_sizes)]
    # Batched callers (DenoisingStage) need a tensor back for CFG arithmetic
    # and scheduler.step; list callers may have per-sample shapes, keep lists.
    return torch.stack(out) if batched_input else out
fastvideo.models.dits.wan_s2v.WanS2VTransformer3DModel.materialize_non_persistent_buffers
materialize_non_persistent_buffers(device: device, dtype: dtype | None = None) -> None

Rebuild the RoPE tables after meta-device construction.

TransformerLoader builds the DiT under torch.device("meta") and then streams checkpoint weights in. Non-persistent buffers are not in the checkpoint, so they stay on meta until this hook (called by fsdp_load) recreates them with real storage. Complex dtype is deliberate -- these are rotation factors, not activations, and must not follow the model dtype.

Source code in fastvideo/models/dits/wan_s2v.py
def materialize_non_persistent_buffers(self, device: torch.device, dtype: torch.dtype | None = None) -> None:
    """Rebuild the RoPE tables after meta-device construction.

    TransformerLoader builds the DiT under ``torch.device("meta")`` and then
    streams checkpoint weights in. Non-persistent buffers are not in the
    checkpoint, so they stay on meta until this hook (called by fsdp_load)
    recreates them with real storage. Complex dtype is deliberate -- these
    are rotation factors, not activations, and must not follow the model dtype.
    """
    head_dim = self.hidden_size // self.num_attention_heads
    if self.freqs.is_meta:
        self.freqs = rope_freqs(head_dim, self.rope_max_seq_len).to(device)
    packer = getattr(self, "frame_packer", None)
    if packer is not None and packer.freqs.is_meta:
        packer.freqs = rope_freqs(packer.inner_dim // packer.num_heads).to(device)

Functions:

fastvideo.models.dits.wan_s2v.rope_apply

rope_apply(x: Tensor, freqs: Tensor) -> Tensor

Rotate q/k of shape [B, L, N, D] by precomputed complex frequencies.

Source code in fastvideo/models/dits/wan_s2v.py
def rope_apply(x: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor:
    """Rotate q/k of shape [B, L, N, D] by precomputed complex frequencies."""
    x_complex = torch.view_as_complex(x.to(torch.float64).unflatten(3, (-1, 2)))
    return torch.view_as_real(x_complex * freqs[:, :x.size(1)]).flatten(3).float()

fastvideo.models.dits.wan_s2v.rope_freqs

rope_freqs(head_dim: int, max_seq_len: int = 1024) -> Tensor

(time, height, width) RoPE tables for one head; time absorbs the remainder of head_dim.

Source code in fastvideo/models/dits/wan_s2v.py
def rope_freqs(head_dim: int, max_seq_len: int = 1024) -> torch.Tensor:
    """(time, height, width) RoPE tables for one head; time absorbs the remainder of head_dim."""
    spatial = 2 * (head_dim // 6)
    return torch.cat([rope_params(max_seq_len, band) for band in (head_dim - 2 * spatial, spatial, spatial)], dim=1)

fastvideo.models.dits.wan_s2v.rope_precompute

rope_precompute(x: Tensor, grid_sizes: list, freqs: Tensor) -> Tensor

Build per-token complex RoPE frequencies for a heterogeneous sequence.

grid_sizes is a list of spans laid out back-to-back in sequence order; each span is [start, end, extent] where every element is a [B, 3] tensor of (frame, height, width). A negative start frame means the span sits in the past (motion frames), which is encoded by walking the time axis backwards and conjugating the temporal band rather than by indexing negatively.

Source code in fastvideo/models/dits/wan_s2v.py
def rope_precompute(x: torch.Tensor, grid_sizes: list, freqs: torch.Tensor) -> torch.Tensor:
    """Build per-token complex RoPE frequencies for a heterogeneous sequence.

    ``grid_sizes`` is a list of spans laid out back-to-back in sequence order;
    each span is ``[start, end, extent]`` where every element is a ``[B, 3]``
    tensor of (frame, height, width). A negative start frame means the span sits
    in the *past* (motion frames), which is encoded by walking the time axis
    backwards and conjugating the temporal band rather than by indexing
    negatively.
    """
    b, s, n, c = x.size(0), x.size(1), x.size(2), x.size(3) // 2
    freqs_t, freqs_h, freqs_w = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
    output = torch.view_as_complex(x.detach().reshape(b, s, n, -1, 2).to(torch.float64))

    offset = 0
    for span_start, span_end, extent in grid_sizes:
        seq_len = 0
        for i in range(span_start.shape[0]):
            f_o, h_o, w_o = span_start[i]
            t_f, t_h, t_w = extent[i]
            seq_f, seq_h, seq_w = (int(v) for v in span_end[i] - span_start[i])
            seq_len = seq_f * seq_h * seq_w
            if seq_len <= 0 or t_f <= 0:
                continue
            past = f_o < 0
            direction = -1 if past else 1
            f_idx = _sample_positions(direction * f_o.item(), direction * t_f.item(), seq_f)
            band_t = freqs_t[f_idx].conj() if past else freqs_t[f_idx]
            output[i, offset:offset + seq_len] = torch.cat([
                band_t.view(seq_f, 1, 1, -1).expand(seq_f, seq_h, seq_w, -1),
                freqs_h[_sample_positions(h_o.item(), t_h.item(), seq_h)].view(1, seq_h, 1, -1).expand(
                    seq_f, seq_h, seq_w, -1),
                freqs_w[_sample_positions(w_o.item(), t_w.item(), seq_w)].view(1, 1, seq_w, -1).expand(
                    seq_f, seq_h, seq_w, -1),
            ], dim=-1).reshape(seq_len, 1, -1)
        offset += seq_len
    return output