Skip to content

vllm_omni.diffusion.models.seedvr2.nadit

SeedVR2 DiT: window geometry, sequence parallelism, rotary embeddings and NaDiT.

DEFAULT_WINDOW module-attribute

DEFAULT_WINDOW: Size3 = (4, 3, 3)

DEFAULT_WINDOW_METHODS module-attribute

DEFAULT_WINDOW_METHODS: tuple[str, ...] = (
    "720pwin_by_size_bysize",
    "720pswin_by_size_bysize",
)

GEOMETRY_VERSION module-attribute

GEOMETRY_VERSION = 1

INT32_MAX module-attribute

INT32_MAX = 2 ** 31 - 1

LANG_THETA module-attribute

LANG_THETA = 10000.0

MAX_TEMPORAL_WINDOW module-attribute

MAX_TEMPORAL_WINDOW = 30

PLANNER_VERSION module-attribute

PLANNER_VERSION = 1

REFERENCE_AREA module-attribute

REFERENCE_AREA = 45 * 80

SEEDVR2_3B_CONFIG module-attribute

SEEDVR2_3B_CONFIG: dict = {
    "vid_in_channels": 33,
    "vid_out_channels": 16,
    "vid_dim": 2560,
    "txt_in_dim": 5120,
    "heads": 20,
    "head_dim": 128,
    "expand_ratio": 4,
    "norm_eps": 1e-05,
    "patch_size": (1, 2, 2),
    "num_layers": 32,
    "mm_layers": 10,
    "window": DEFAULT_WINDOW,
    "window_method": DEFAULT_WINDOW_METHODS,
    "rope_dim": 128,
    "vid_out_norm": True,
}

Size3 module-attribute

Size3 = tuple[int, int, int]

WINDOW_METHODS module-attribute

WINDOW_METHODS: dict[
    str, Callable[[Size3, Size3], WindowSlices]
] = {
    "720pwin_by_size_bysize": make_720p_windows,
    "720pswin_by_size_bysize": make_720p_shifted_windows,
}

WindowSlices module-attribute

WindowSlices = list[tuple[slice, slice, slice]]

logger module-attribute

logger = init_logger(__name__)

AdaSingle

Bases: Module

Timestep modulation with per-layer shift / scale / gate parameters.

dim instance-attribute

dim = int(dim)

emb_dim instance-attribute

emb_dim = int(emb_dim)

layers instance-attribute

layers = list(layers)

forward

forward(
    hid: Tensor, emb: Tensor, layer: str, mode: str
) -> Tensor

slice

slice(
    emb: Tensor, layer: str
) -> tuple[Tensor, Tensor, Tensor]

(shift, scale, gate) embedding slices for layer, each [1, dim].

DeviceWindowRedistributionPlan dataclass

Device-side materialisation of :class:WindowRedistributionPlan.

input_split_sizes instance-attribute

input_split_sizes: tuple[int, ...]

network_exchange_required instance-attribute

network_exchange_required: bool

num_recv_rows instance-attribute

num_recv_rows: int

output_split_sizes instance-attribute

output_split_sizes: tuple[int, ...]

recv_to_dst_indices instance-attribute

recv_to_dst_indices: Tensor

send_indices instance-attribute

send_indices: Tensor

LocalWindowContext dataclass

Everything one rank needs to run local window attention for one layout.

global_windows instance-attribute

global_windows: int

joint_cu_seqlens instance-attribute

joint_cu_seqlens: Tensor

joint_len property

joint_len: int

joint_lengths property

joint_lengths: Tensor

joint_order instance-attribute

joint_order: Tensor

layout_key instance-attribute

layout_key: WindowLayoutKey

local_windows instance-attribute

local_windows: int

max_joint_len property

max_joint_len: int

num_video_tokens property

num_video_tokens: int

text_len instance-attribute

text_len: int

txt_src instance-attribute

txt_src: Tensor

vid_src instance-attribute

vid_src: Tensor

video_cu_seqlens instance-attribute

video_cu_seqlens: Tensor

window_shapes instance-attribute

window_shapes: Tensor

MMArg

Pair of per-stream values (video / text) used by :class:MMModule.

txt instance-attribute

txt = txt

vid instance-attribute

vid = vid

MMModule

Bases: Module

Apply a module to the video and text streams, optionally sharing weights.

shared_weights=True builds a single submodule used for both streams (the reference's all prefix in the checkpoint); vid_only=True drops the text branch entirely.

all instance-attribute

all = factory(dims.vid)

shared_weights instance-attribute

shared_weights = shared_weights

txt instance-attribute

txt = None if vid_only else factory(dims.txt)

vid instance-attribute

vid = factory(dims.vid)

vid_only instance-attribute

vid_only = vid_only

forward

forward(vid: Tensor, txt: Tensor | None, *args, **kwargs)

NaDiTOutput

Container matching the reference NaDiTOutput.

vid_sample instance-attribute

vid_sample = vid_sample

NaMMRotaryEmbedding3d

Bases: Module

Reference mmrope3d: axial 3-D RoPE for windows plus 1-D RoPE for text.

axis_dim instance-attribute

axis_dim = self.rotary_dim // self.num_axes

mm class-attribute instance-attribute

mm = True

num_axes instance-attribute

num_axes = int(num_axes)

rot_dim instance-attribute

rot_dim = self.axis_dim * self.num_axes

rotary_dim instance-attribute

rotary_dim = int(rotary_dim)

forward

forward(
    vid_q: Tensor,
    vid_k: Tensor,
    vid_freqs: Tensor,
    txt_q: Tensor,
    txt_k: Tensor,
    txt_freqs: Tensor,
) -> tuple[Tensor, Tensor, Tensor, Tensor]

text_freqs

text_freqs(text_len: int, *, device, dtype) -> Tensor

[text_len, rot_dim] text frequencies (one axis repeated rope_dim times).

window_freqs

window_freqs(
    window_shape: tuple[int, int, int],
    text_len: int,
    *,
    device,
    dtype,
) -> Tensor

[f * h * w, rot_dim] frequencies of one window, window-local.

window_freqs_batch

window_freqs_batch(
    window_shapes: Tensor | list[tuple[int, int, int]],
    text_len: int,
    *,
    device,
    dtype,
) -> Tensor

Concatenated per-window video frequencies for a window batch.

NaMMSRTransformerBlock

Bases: Module

Reference NaMMSRTransformerBlock: modulated attention + SwiGLU MLP.

ada instance-attribute

ada = MMModule(
    lambda d: AdaSingle(
        int(d), emb_dim, layers=["attn", "mlp"]
    ),
    dims,
    shared_weights=shared_weights,
    vid_only=is_last_layer,
)

attn instance-attribute

attn = NaSwinAttention(
    vid_dim=vid_dim,
    txt_dim=txt_dim,
    heads=heads,
    head_dim=head_dim,
    qk_bias=qk_bias,
    qk_norm_eps=norm_eps,
    rope_dim=rope_dim,
    shared_weights=shared_weights,
    use_varlen_kernel=use_varlen_kernel,
)

attn_norm instance-attribute

attn_norm = MMModule(
    lambda d: RMSNorm(
        int(d), eps=norm_eps, elementwise_affine=False
    ),
    dims,
    shared_weights=shared_weights,
)

is_last_layer instance-attribute

is_last_layer = bool(is_last_layer)

mlp instance-attribute

mlp = MMModule(
    lambda d: SwiGLUMLP(int(d), expand_ratio),
    dims,
    shared_weights=shared_weights,
    vid_only=is_last_layer,
)

mlp_norm instance-attribute

mlp_norm = MMModule(
    lambda d: RMSNorm(
        int(d), eps=norm_eps, elementwise_affine=False
    ),
    dims,
    shared_weights=shared_weights,
    vid_only=is_last_layer,
)

forward

forward(
    vid: Tensor,
    txt: Tensor,
    emb: Tensor,
    ctx: LocalWindowContext,
    runtime: SeedVR2WindowRuntime | None,
) -> tuple[Tensor, Tensor]

NaSwinAttention

Bases: Module

Joint video+text window attention with a globally averaged text output.

attention instance-attribute

attention = Attention(
    num_heads=heads,
    head_size=head_dim,
    causal=False,
    softmax_scale=self.softmax_scale,
    role="self",
    skip_sequence_parallel=True,
)

attention_backend_name instance-attribute

attention_backend_name = (
    backend.get_name() if backend is not None else None
)

attention_path instance-attribute

attention_path = (
    "packed_varlen"
    if self.use_varlen_kernel
    else "grouped_sdpa"
)

attention_stats instance-attribute

attention_stats = {
    "packed_varlen_calls": 0,
    "grouped_sdpa_calls": 0,
    "no_local_windows_calls": 0,
}

head_dim instance-attribute

head_dim = int(head_dim)

heads instance-attribute

heads = int(heads)

inner_dim instance-attribute

inner_dim = int(inner_dim)

norm_k instance-attribute

norm_k = MMModule(
    lambda d: RMSNorm(
        int(d), eps=qk_norm_eps, elementwise_affine=True
    ),
    MMArg(head_dim, head_dim),
    shared_weights=shared_weights,
)

norm_q instance-attribute

norm_q = MMModule(
    lambda d: RMSNorm(
        int(d), eps=qk_norm_eps, elementwise_affine=True
    ),
    MMArg(head_dim, head_dim),
    shared_weights=shared_weights,
)

proj_out instance-attribute

proj_out = MMModule(
    lambda d: nn.Linear(inner_dim, int(d)),
    dims,
    shared_weights=shared_weights,
)

proj_qkv instance-attribute

proj_qkv = MMModule(
    lambda d: nn.Linear(
        int(d), 3 * inner_dim, bias=qk_bias
    ),
    dims,
    shared_weights=shared_weights,
)

rope instance-attribute

rope = NaMMRotaryEmbedding3d(
    rotary_dim=rope_dim, num_axes=3
)

softmax_scale instance-attribute

softmax_scale = 1.0 / math.sqrt(head_dim)

use_varlen_kernel instance-attribute

use_varlen_kernel = requested_varlen and supports_varlen

varlen_fallback_reason instance-attribute

varlen_fallback_reason: str | None = None

forward

forward(
    vid: Tensor,
    txt: Tensor,
    ctx: LocalWindowContext,
    runtime: SeedVR2WindowRuntime | None,
) -> tuple[Tensor, Tensor]

OutAda

Bases: Module

The released vid_out_ada parameters (see the module docstring).

out_scale instance-attribute

out_scale = nn.Parameter(torch.ones(dim))

out_shift instance-attribute

out_shift = nn.Parameter(torch.zeros(dim))

RMSNorm

Bases: Module

Reference CustomRMSNorm: normalise in the input dtype, optional affine.

eps instance-attribute

eps = float(eps)

weight instance-attribute

weight = nn.Parameter(torch.ones(dim))

forward

forward(x: Tensor) -> Tensor

RankWindowPlan dataclass

Rank-local view of a layout: what this rank owns and how it is packed.

global_token_ids instance-attribute

global_token_ids: Tensor

layout_key instance-attribute

layout_key: WindowLayoutKey

num_local_tokens property

num_local_tokens: int

num_local_windows property

num_local_windows: int

rank instance-attribute

rank: int

video_cu_seqlens instance-attribute

video_cu_seqlens: Tensor

window_ids instance-attribute

window_ids: Tensor

window_shapes instance-attribute

window_shapes: Tensor

SeedVR2NaDiT

Bases: Module

SeedVR2 3B NaDiT transformer with window-aligned SP support.

blocks instance-attribute

blocks = nn.ModuleList(
    [
        NaMMSRTransformerBlock(
            vid_dim=vid_dim,
            txt_dim=txt_dim,
            emb_dim=emb_dim,
            heads=heads,
            head_dim=head_dim,
            expand_ratio=expand_ratio,
            norm_eps=norm_eps,
            qk_bias=False,
            mlp_type=mlp_type,
            shared_weights=not index < mm_layers,
            rope_dim=rope_dim,
            is_last_layer=index == num_layers - 1,
            use_varlen_kernel=use_varlen_kernel,
        )
        for index in range(num_layers)
    ]
)

emb_in instance-attribute

emb_in = TimeEmbedding(
    sinusoidal_dim=256,
    hidden_dim=max(vid_dim, txt_dim),
    output_dim=emb_dim,
)

num_layers instance-attribute

num_layers = int(num_layers)

patch_size instance-attribute

patch_size = tuple(int(v) for v in patch_size)

txt_dim instance-attribute

txt_dim = int(txt_dim)

txt_in instance-attribute

txt_in = (
    nn.Linear(txt_in_dim, txt_dim)
    if txt_in_dim != txt_dim
    else nn.Identity()
)

vid_dim instance-attribute

vid_dim = int(vid_dim)

vid_in instance-attribute

vid_in = _NaPatchIn(
    in_channels=vid_in_channels,
    patch_size=self.patch_size,
    dim=vid_dim,
)

vid_out instance-attribute

vid_out = _NaPatchOut(
    out_channels=vid_out_channels,
    patch_size=self.patch_size,
    dim=vid_dim,
)

vid_out_ada instance-attribute

vid_out_ada = OutAda(vid_dim)

vid_out_norm instance-attribute

vid_out_norm = (
    RMSNorm(vid_dim, eps=norm_eps, elementwise_affine=True)
    if vid_out_norm
    else None
)

window instance-attribute

window = tuple(int(v) for v in window)

window_method instance-attribute

window_method = tuple(window_method)

attention_path_summary

attention_path_summary() -> dict

Requested vs resolved attention path, backends and per-path call counts.

build_runtime

build_runtime(
    token_grid: tuple[int, int, int],
    *,
    text_len: int,
    group=None,
    world_size: int = 1,
    rank: int = 0,
    parallel_config: DiffusionParallelConfig | None = None,
) -> SeedVR2WindowRuntime

Create the window-SP driver for one request (SP=1 included).

forward

forward(
    vid: Tensor,
    txt: Tensor,
    vid_shape: Tensor,
    txt_shape: Tensor,
    timestep,
    runtime: SeedVR2WindowRuntime | None = None,
) -> NaDiTOutput

Run the transformer.

vid is the flattened pre-patchify latent [T*H*W, C] with vid_shape holding the raw latent (T, H, W); txt is [L, txt_in_dim]. With runtime set, only this rank's windows are ever materialised; without it the model runs the SP=1 path through the same planner, so a regular -> shifted transition still permutes rows.

reset_attention_stats

reset_attention_stats() -> None

token_grid_for

token_grid_for(vid_shape: Tensor) -> tuple[int, int, int]

SeedVR2WindowRuntime

Owns the window-SP plan for one request and drives the layer schedule.

group instance-attribute

group = group

manager instance-attribute

manager = WindowLayoutManager(
    token_grid,
    group=group,
    world_size=self.world_size,
    rank=self.rank,
    window=window,
    methods=methods,
    num_layers=num_layers,
    planner_version=planner_version,
)

num_layers instance-attribute

num_layers = int(num_layers)

rank instance-attribute

rank = int(rank)

stats instance-attribute

stats: dict[str, int] = {
    "layout_transitions": 0,
    "network_transitions": 0,
    "local_reorders": 0,
    "remote_video_rows": 0,
    "text_all_reduces": 0,
    "fused_text_mean_calls": 0,
}

text_len instance-attribute

text_len = int(text_len)

transitions property

transitions: int

world_size instance-attribute

world_size = int(world_size)

collective_group

collective_group() -> ProcessGroup | None

Group for the text reduction, or None for a local mean.

world_size == 1 is the local case. A multi-rank runtime without a usable group is rejected rather than silently degrading to a rank-local mean, which would produce a different text state.

context

context(
    layout: WindowLayout, device: device | str
) -> LocalWindowContext

current_device

current_device() -> device | None

ensure_layout

ensure_layout(
    hidden: Tensor,
    current: WindowLayoutKey,
    required: WindowLayoutKey,
) -> Tensor

layout_for_layer

layout_for_layer(layer_index: int) -> WindowLayout

local_rows_for

local_rows_for(
    canonical_rows: Tensor, layout: WindowLayout
) -> Tensor

Select this rank's rows of a canonical [N, C] tensor.

reduce_text

reduce_text(
    local_window_sum: Tensor, global_windows: int
) -> Tensor

Global window mean of the per-window text attention outputs.

Shape adaptation and the runtime statistics live here; the arithmetic and the collective live in :func:global_window_mean, which the tests and the single-rank attention path use as well.

remote_video_bytes

remote_video_bytes(
    hidden_channels: int, element_size: int
) -> int

Logical source->destination payload of the transitions seen so far.

to_canonical_rows

to_canonical_rows(
    local_rows: Tensor, layout: WindowLayout
) -> Tensor

Gather every rank's rows back into canonical token order (all ranks).

SwiGLUMLP

Bases: Module

Reference SwiGLU MLP (multiple_of=256, no biases).

proj_in instance-attribute

proj_in = nn.Linear(dim, hidden, bias=False)

proj_in_gate instance-attribute

proj_in_gate = nn.Linear(dim, hidden, bias=False)

proj_out instance-attribute

proj_out = nn.Linear(hidden, dim, bias=False)

forward

forward(x: Tensor) -> Tensor

TimeEmbedding

Bases: Module

Sinusoidal timestep embedding followed by a two-layer SiLU MLP.

act instance-attribute

act = nn.SiLU()

proj_hid instance-attribute

proj_hid = nn.Linear(hidden_dim, hidden_dim)

proj_in instance-attribute

proj_in = nn.Linear(sinusoidal_dim, hidden_dim)

proj_out instance-attribute

proj_out = nn.Linear(hidden_dim, output_dim)

sinusoidal_dim instance-attribute

sinusoidal_dim = int(sinusoidal_dim)

forward

forward(timestep, device: device, dtype: dtype) -> Tensor

WindowAssignment dataclass

Which rank owns which window (a whole window belongs to one rank).

ideal_lower_bound property

ideal_lower_bound: int

Largest window, or perfectly balanced load, whichever is bigger.

layout_key instance-attribute

layout_key: WindowLayoutKey

max_rank_tokens property

max_rank_tokens: int

max_window_tokens class-attribute instance-attribute

max_window_tokens: int = 0

owner_rank instance-attribute

owner_rank: Tensor

planner_version instance-attribute

planner_version: int

rank_token_counts instance-attribute

rank_token_counts: tuple[int, ...]

rank_window_counts instance-attribute

rank_window_counts: tuple[int, ...]

rank_window_ids instance-attribute

rank_window_ids: tuple[Tensor, ...]

world_size instance-attribute

world_size: int

load_imbalance

load_imbalance() -> float

WindowLayout dataclass

A partition of the video tokens into windows, in window-packed order.

window_token_ids holds canonical token ids ((t * H + h) * W + w) of every packed row; window_offsets delimits the windows inside it.

key instance-attribute

num_tokens instance-attribute

num_tokens: int

num_windows property

num_windows: int

window_offsets instance-attribute

window_offsets: Tensor

window_shapes instance-attribute

window_shapes: Tensor

window_token_ids instance-attribute

window_token_ids: Tensor

window_lengths

window_lengths() -> Tensor

window_token_ids_of

window_token_ids_of(window_ids: Tensor) -> Tensor

Canonical ids of the given windows, concatenated in that order.

WindowLayoutKey dataclass

Identity of one window layout (shape + geometry, not ownership).

geometry_fingerprint instance-attribute

geometry_fingerprint: str

geometry_version instance-attribute

geometry_version: int

token_grid instance-attribute

token_grid: tuple[int, int, int]

window_method instance-attribute

window_method: str

window_shape instance-attribute

window_shape: tuple[int, int, int]

WindowLayoutManager

Drives ensure_layout across a model's per-layer window schedule.

One manager per request / per model instance. It owns the device plans of the current layout pair and is the only place that decides whether a transition is needed and in which direction. The CPU plan cache is shared process-wide, so a second request with the same geometry does not rebuild layouts and routing plans that cost seconds on a large token grid.

cache instance-attribute

cache = shared_plan_cache() if cache is None else cache

device property

device: device

group instance-attribute

group = group

methods instance-attribute

methods = tuple(methods)

num_layers instance-attribute

num_layers = int(num_layers)

planner_version instance-attribute

planner_version = int(planner_version)

rank instance-attribute

rank = int(rank)

token_grid instance-attribute

token_grid = (
    int(token_grid[0]),
    int(token_grid[1]),
    int(token_grid[2]),
)

transitions instance-attribute

transitions = 0

window instance-attribute

window = tuple(window)

world_size instance-attribute

world_size = int(world_size)

assignment

assignment(layout: WindowLayout) -> WindowAssignment

device_plan

device_plan(
    src_layout: WindowLayout, dst_layout: WindowLayout
) -> DeviceWindowRedistributionPlan

ensure_layout

ensure_layout(
    hidden: Tensor,
    src_layout_key: WindowLayoutKey,
    dst_layout_key: WindowLayoutKey,
) -> Tensor

Return hidden re-sharded into the destination layout (or as-is).

layer_layout

layer_layout(layer_index: int) -> WindowLayout

layer_method

layer_method(layer_index: int) -> str

layout

layout(window_method: str) -> WindowLayout

layout_for_key

layout_for_key(key: WindowLayoutKey) -> WindowLayout

rank_plan

rank_plan(
    layout: WindowLayout,
    assignment: WindowAssignment | None = None,
) -> RankWindowPlan

rank_plan_for

rank_plan_for(
    layout: WindowLayout,
    rank: int,
    assignment: WindowAssignment | None = None,
) -> RankWindowPlan

set_device

set_device(device: device | str) -> None

WindowPlanCache

Small LRU cache for layouts, assignments, rank plans and routes.

The assignment / routing identity includes the SP world size, the planner version and the geometry fingerprint, so a change in any of them cannot hit a stale entry. Runtime (device) objects are intentionally not cached here: they are owned by the caller and must be released when the process group or device changes.

hits instance-attribute

hits = 0

misses instance-attribute

misses = 0

assignment

assignment(
    layout: WindowLayout,
    world_size: int,
    planner_version: int = PLANNER_VERSION,
)

clear

clear() -> None

layout

layout(
    token_grid,
    window_method,
    window=DEFAULT_WINDOW,
    geometry_version=GEOMETRY_VERSION,
)

rank_plan

rank_plan(
    layout: WindowLayout,
    assignment: WindowAssignment,
    rank: int,
)

route

route(
    src_layout: WindowLayout,
    src_assignment: WindowAssignment,
    dst_layout: WindowLayout,
    dst_assignment: WindowAssignment,
)

WindowRedistributionPlan dataclass

Plan A routing contract for one rank of one layout transition.

input_split_sizes[d] is the number of rows this rank sends to rank d and output_split_sizes[s] the number of rows it receives from rank s (feature rows, not bytes). send_indices selects rows of the local source tensor in send order; recv_to_dst_indices permutes the received rows into destination-local order.

dst_key instance-attribute

dst_key: WindowLayoutKey

input_split_sizes instance-attribute

input_split_sizes: tuple[int, ...]

network_exchange_required instance-attribute

network_exchange_required: bool

num_recv_rows property

num_recv_rows: int

num_send_rows property

num_send_rows: int

output_split_sizes instance-attribute

output_split_sizes: tuple[int, ...]

planner_version class-attribute instance-attribute

planner_version: int = PLANNER_VERSION

rank instance-attribute

rank: int

recv_to_dst_indices instance-attribute

recv_to_dst_indices: Tensor

send_indices instance-attribute

send_indices: Tensor

src_key instance-attribute

src_key: WindowLayoutKey

world_size instance-attribute

world_size: int

apply_rotary_emb

apply_rotary_emb(freqs: Tensor, t: Tensor) -> Tensor

Rotate t with precomputed freqs (fp32 math, original dtype out).

Mirrors rotary_embedding_torch.apply_rotary_emb for the 2-D frequency tables and 3-D (heads, seq, dim) inputs used here: the rotated span is freqs.shape[-1] channels and any remaining channels are untouched.

build_local_window_context

build_local_window_context(
    rank_plan: RankWindowPlan,
    *,
    text_len: int,
    global_windows: int,
    device: device | str,
) -> LocalWindowContext

Build the device-local attention metadata for one rank and one layout.

build_rank_window_plan

build_rank_window_plan(
    layout: WindowLayout,
    assignment: WindowAssignment,
    rank: int,
) -> RankWindowPlan

Rank-local metadata for one layout (used for windows/attention packing).

build_redistribution_plans

build_redistribution_plans(
    src_layout: WindowLayout,
    src_assignment: WindowAssignment,
    dst_layout: WindowLayout,
    dst_assignment: WindowAssignment,
    *,
    planner_version: int = PLANNER_VERSION,
) -> tuple[WindowRedistributionPlan, ...]

Build the bidirectional Plan A routing plans for every rank.

Returns one :class:WindowRedistributionPlan per rank, indexed by SP-local rank. Both directions (regular -> shifted and shifted -> regular) are plain calls to this function; neither is derived from the other by transposing.

build_window_assignment

build_window_assignment(
    layout: WindowLayout,
    world_size: int,
    *,
    planner_version: int = PLANNER_VERSION,
) -> WindowAssignment

Deterministic token-balanced window -> rank assignment (LPT).

Windows are considered in (-video_token_count, window_id) order and each is placed on the rank with the smallest (assigned_tokens, rank_id). The result is independent of container iteration order.

build_window_layout

build_window_layout(
    token_grid: tuple[int, int, int],
    window_method: str,
    *,
    window: tuple[int, int, int] = DEFAULT_WINDOW,
    geometry_version: int = GEOMETRY_VERSION,
) -> WindowLayout

Build the CPU layout of one layer's window partition.

geometry_fingerprint

geometry_fingerprint(
    size: Size3,
    window_method: str,
    window: Size3,
    geometry_version: int = GEOMETRY_VERSION,
) -> str

Deterministic fingerprint of a layout geometry.

Uses SHA-1 over the parameters and the materialised window boundaries so the value is stable across processes (unlike Python's randomised hash()).

get_window_op

get_window_op(
    name: str,
) -> Callable[[Size3, Size3], WindowSlices]

Return the window function for a reference window_method name.

global_window_mean

global_window_mean(
    local_window_sum: Tensor,
    global_window_count: int,
    *,
    group: ProcessGroup | None = None,
    dtype: dtype | None = None,
) -> Tensor

Reference text reduction: (sum over all windows) / global_window_count.

Every rank contributes the sum of its own windows' text attention output (a zero tensor on ranks without windows) and all ranks divide by the global window count. This is neither a token-weighted mean nor a mean of rank means.

group=None means local reduction: no collective is issued, and None is never forwarded to all_reduce (which would silently use the default process group). fp16/bf16 inputs are accumulated in fp32, fp32/fp64 keep their own precision, the caller's tensor is not modified, and dtype is applied only after the reduction and the division.

grouped_window_sdpa

grouped_window_sdpa(
    q: Tensor,
    k: Tensor,
    v: Tensor,
    ctx: LocalWindowContext,
    *,
    softmax_scale: float,
) -> Tensor

Exact window-local attention without a varlen kernel.

Windows are packed back-to-back in the joint sequence, so this groups the equal-length windows and runs one batched SDPA per group (ragged window sizes keep the group count small). It is the portable path used when the selected backend has no packed-varlen entry point.

joint_cu_seqlens

joint_cu_seqlens(
    video_cu_seqlens: Tensor, text_len: int
) -> Tensor

[W + 1] int32 joint (video + replicated text) offsets for local attention.

video_cu_seqlens describes video packing only and must not be used directly as joint Q/K offsets. The offsets are accumulated in int64 on the input's device and converted to int32 only after validation, so CUDA metadata never round-trips through the host and an overflow is caught before a wrapped offset can be produced. text_len may be 0; an input of [0] (a rank without windows) yields [0].

lang_freqs

lang_freqs(
    axis_dim: int, theta: float = LANG_THETA
) -> Tensor

1 / theta ** (arange(0, axis_dim, 2) / axis_dim) (rotary_embedding_torch).

make_720p_shifted_windows

make_720p_shifted_windows(
    size: Size3, num_windows: Size3
) -> WindowSlices

Shifted window slices with boundary clipping, reference-faithful.

make_720p_windows

make_720p_windows(
    size: Size3, num_windows: Size3
) -> WindowSlices

Regular (non-shifted) window slices, reference-faithful.

materialize_redistribution_plan

materialize_redistribution_plan(
    plan: WindowRedistributionPlan, *, device: device | str
) -> DeviceWindowRedistributionPlan

Move one CPU redistribution plan to the device (index tensors only).

pack_joint_windows

pack_joint_windows(
    video: Tensor, text: Tensor, ctx: LocalWindowContext
) -> Tensor

Interleave [video window, replicated text] pairs for varlen attention.

patchify

patchify(
    vid: Tensor,
    shape: tuple[int, int, int],
    patch_size: Sequence[int],
)

Patch embedding layout of the reference (token-major rows).

redistribute_window_rows

redistribute_window_rows(
    hidden: Tensor,
    plan: DeviceWindowRedistributionPlan,
    *,
    group: ProcessGroup,
) -> Tensor

Re-shard local video rows from one window layout to another (Plan A).

hidden is [N_src, C] (contiguous trailing feature dimension) in the source layout's rank-local row order; the result is [N_dst, C] in the destination layout's rank-local row order. Only row placement and order change -- values are moved bit-exactly.

The first version is deliberately synchronous: no extra streams, no cuda.synchronize() / barrier "fixups", and no hidden catch-up all-gather.

reduction_dtype

reduction_dtype(dtype: dtype) -> dtype

Accumulator dtype for a window reduction (fp16/bf16 reduce in fp32).

rotate_half

rotate_half(x: Tensor) -> Tensor

schedule_methods

schedule_methods(
    num_layers: int,
    methods: Iterable[str] = DEFAULT_WINDOW_METHODS,
) -> tuple[str, ...]

shared_plan_cache

shared_plan_cache() -> WindowPlanCache

The process-wide plan cache that requests reuse across identical geometry.

to_cu_seqlens

to_cu_seqlens(
    offsets: Tensor, *, context: str = ""
) -> Tensor

Convert int64 offsets to int32 cu_seqlens with an explicit overflow check.

transition_count

transition_count(
    methods: Sequence[str], num_layers: int | None = None
) -> int

Number of layout transitions for a layer schedule (diagnostic only).

The 3B schedule alternates, so every consecutive layer pair is a transition; the entry distribution and the final reconstruction are counted separately by the caller.

unpack_joint_windows

unpack_joint_windows(
    joint_out: Tensor, ctx: LocalWindowContext
) -> tuple[Tensor, Tensor]

Split the attention output into packed video rows and per-window text rows.

unpatchify

unpatchify(
    rows: Tensor,
    token_grid: tuple[int, int, int],
    patch_size: Sequence[int],
) -> Tensor

Inverse of :func:patchify: [N, t*h*w*c] -> [T*H*W, c] (reference order).

validate_seedvr2_parallel_config

validate_seedvr2_parallel_config(
    parallel: DiffusionParallelConfig,
) -> None

window_layout_geometry

window_layout_geometry(
    size: Size3,
    window_method: str,
    window: Size3 = DEFAULT_WINDOW,
)

Build the CPU geometry (offsets, canonical ids, shapes) of one layout.

Returns (window_offsets, window_token_ids, window_shapes) where

  • window_offsets is int64 [W + 1] (start row of every window in window-packed order),
  • window_token_ids is int64 [N] with the canonical token id of every packed row, and
  • window_shapes is int64 [W, 3] with the (t, h, w) extent of every window.

The ids satisfy window_token_ids[offsets[i]:offsets[i + 1]] = canonical ids of window i, in reference window order and reference in-window flatten order.