Skip to content

vllm_omni.diffusion.attention.backends.utils.piecewise_attn

Piecewise attention for mixed causal / full (bidirectional) masks.

Dispatches each segment as a separate attention call whose causal flag follows FlashAttention's bottom-right convention (K[:e] is attended by Q[s:e], with causal alignment anchored at the bottom-right corner).

Per segment
  • causal segment [s, e): attn(Q[:, s:e], K[:, :e], V[:, :e], causal=True)
  • full-attn span [a, b) intersecting the query range at [s, e): attn(Q[:, s:e], K[:, :b], V[:, :b], causal=False)

PagedPiecewiseRunner module-attribute

PagedPiecewiseRunner = Callable[
    [
        torch.Tensor,
        torch.Tensor | None,
        torch.Tensor | None,
        object,
        torch.Tensor | None,
    ],
    torch.Tensor,
]

PagedPiecewisePlan dataclass

Piecewise segment batches for one packed paged-attention batch.

homogeneous_batch_shape class-attribute instance-attribute

homogeneous_batch_shape: tuple[int, int] | None = None

num_query_tokens instance-attribute

num_query_tokens: int

segments instance-attribute

segments: tuple[PagedPiecewiseSegment, ...]

segments_cover_query_contiguously instance-attribute

segments_cover_query_contiguously: bool

spans instance-attribute

spans: tuple[tuple[tuple[int, int], ...], ...]

PagedPiecewiseSegment dataclass

One aligned segment across rows in a packed paged batch.

local_query_range class-attribute instance-attribute

local_query_range: tuple[int, int] | None = None

query_indices instance-attribute

query_indices: Tensor

query_range instance-attribute

query_range: tuple[int, int] | None

row_segments instance-attribute

row_segments: tuple[Segment, ...]

Segment

Bases: NamedTuple

kv_end instance-attribute

kv_end: int

mode instance-attribute

mode: Literal['causal', 'full']

q_end instance-attribute

q_end: int

q_start instance-attribute

q_start: int

build_paged_piecewise_plan

build_paged_piecewise_plan(
    full_attn_spans: Sequence[Sequence[tuple[int, int]]],
    query_offsets: Sequence[int],
    query_lens: Sequence[int],
    seq_lens: Sequence[int],
    *,
    device: device | str | None = None,
) -> PagedPiecewisePlan

Build packed indices for corresponding piecewise segments in each row.

build_segments

build_segments(full_attn_spans, query_offset, query_len)

full_attn_spans: list of (start, end) half-open spans in global coordinates query_offset: starting position of query in the global sequence query_len: length of the query

return

List[Segment] in global coordinates, clipped to [query_offset, query_offset + query_len). Full-attention segments retain the original span end as kv_end so a local query shard can attend past its own boundary.

piecewise_attn

piecewise_attn(
    query,
    key,
    value,
    full_attn_spans: list[list[tuple[int, int]]],
    softmax_scale: float,
    attn_func,
    query_ranges: tuple[QueryRange, ...] | None = None,
)

run_paged_piecewise_plan

run_paged_piecewise_plan(
    query: Tensor,
    key: Tensor | None,
    value: Tensor | None,
    plan: PagedPiecewisePlan,
    segment_metadata: Sequence[object],
    segment_runner: PagedPiecewiseRunner,
    output_buffer: Tensor | None = None,
    *,
    use_homogeneous_batch: bool = False,
) -> Tensor

Run piecewise attention with view fast paths for contiguous segments.

use_homogeneous_batch is an opt-in for native backends that can keep identical rows in the legacy batch layout. The default remains the indexed packed path so existing GPU/heterogeneous contracts are unchanged.