Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
a216768
[wip]: exact-size pinned arenas for H3 offload and 4090 benchmarks
aryan5v Oct 3, 2026
a491f15
[wip]: document 4090 setup validation and reproducible baselines
aryan5v Oct 3, 2026
4ae791c
[wip]: record checkpoint revision and GPU driver in 4090 benchmark re…
aryan5v Oct 3, 2026
21f9899
[wip]: record completed 480p baseline and stage summary tooling
aryan5v Oct 3, 2026
c716408
[perf]: reduce H3 attention activation copies and share FP8 input qua…
aryan5v Oct 3, 2026
bd713dd
[wip]: prototype tile-64 INT8 QK and FP8 PV attention on sm89
aryan5v Oct 3, 2026
5bf9804
[perf]: retain offloaded H3 VAEs on host until their consuming stages
aryan5v Oct 3, 2026
59e6946
[perf]: add opt-in sm89 tile-64 BF16 and INT8 QK attention
aryan5v Oct 3, 2026
b3ab6e1
[docs]: record cached 4090 timings and sm89 precision validation
aryan5v Oct 3, 2026
994c220
[wip]: validate tilewise FP8 values and dynamic probability scales
aryan5v Oct 3, 2026
6b18535
[docs]: record resident-block timings and encoder constraints
aryan5v Oct 3, 2026
92fc26c
[docs]: record five-second 4090 timing and FP8 encoder footprint
aryan5v Oct 3, 2026
d986589
[feat]: stream the H3 encoder for consumer VRAM limits
aryan5v Oct 3, 2026
887deaa
[perf]: fuse serialized NVFP4 encoder weight dequantization
aryan5v Oct 3, 2026
2aa19c4
[perf]: release H3 fine-attention copies before the gated merge
aryan5v Oct 3, 2026
9c9f1ed
[perf]: share VAE INT8 input preparation and avoid weight copies
aryan5v Oct 3, 2026
e457b68
[docs]: record 12 GiB and streamed 4090 results
aryan5v Oct 3, 2026
7a0d7d3
[perf]: read H3 INT8 sparse attention from existing layouts
aryan5v Oct 3, 2026
3c0668f
[perf]: fuse H3 INT8 VAE scaling and bias without intermediates
aryan5v Oct 3, 2026
a7f7ede
[bugfix]: retain consumer attention environment parsing after core re…
aryan5v Oct 3, 2026
87b22a5
[test]: adapt consumer VSA validation to packed segment metadata
aryan5v Oct 3, 2026
fb92af1
[bench]: sample total GPU memory for consumer budget checks
aryan5v Oct 3, 2026
5ac85ad
[docs]: record 42-second 4090 clips and consumer memory validation
aryan5v Oct 3, 2026
98cdc7d
[docs]: record Track B PR and strict 8 GiB budget limit
aryan5v Oct 3, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
189 changes: 189 additions & 0 deletions fastvideo/attention/backends/minimax_h3_sparse_int8.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,189 @@
# SPDX-License-Identifier: Apache-2.0
"""Experimental sm89 tile-64 VSA with BF16 or INT8 QK and FP32 accumulation.

Retains each query tile's original key selection and masks partial key tiles.
Unlike a 128-query adapter, it adds no attention blocks. Q/K use per-token
scales; K centering is a softmax-invariant shift. V uses one scale per head
and channel, so its dequantization can be applied once in the epilogue.
BF16 PV is the default: FP8 PV had excessive error on real H3 inputs.
Numerical validation and same-seed clip review are required before enabling.
"""
from __future__ import annotations

import torch
import triton
import triton.language as tl


@triton.jit
def _quantize_qk(X, Mean, VBS, Y, Scale, L: tl.constexpr, D: tl.constexpr, H: tl.constexpr, XB: tl.constexpr,
XH: tl.constexpr, XS: tl.constexpr, XD: tl.constexpr, CENTER: tl.constexpr, ROWS: tl.constexpr):
hz = tl.program_id(1)
rows = tl.program_id(0) * ROWS + tl.arange(0, ROWS)
cols = tl.arange(0, D)
offset = (hz // H) * XB + (hz % H) * XH + rows[:, None] * XS + cols[None, :] * XD
x = tl.load(X + offset, rows[:, None] < L, 0).to(tl.float32)
if CENTER:
mean = tl.load(Mean + hz * D + cols)
valid_size = tl.load(VBS + rows // 64, rows < L, 0)
x = tl.where((rows % 64 < valid_size)[:, None], x - mean[None, :], 0.0)
scale = tl.maximum(tl.max(tl.abs(x), 1) / 127.0, 1e-8)
y = tl.floor(x / scale[:, None] + 0.5).to(tl.int8)
tl.store(Y + (hz * L + rows[:, None]) * D + cols[None, :], y, rows[:, None] < L)
tl.store(Scale + hz * L + rows, scale, rows < L)


@triton.jit
def _quantize_v(X, Scale, Y, L: tl.constexpr, D: tl.constexpr, ROWS: tl.constexpr):
hz = tl.program_id(1)
rows = tl.program_id(0) * ROWS + tl.arange(0, ROWS)
cols = tl.arange(0, D)
scale = tl.load(Scale + hz * D + cols)
x = tl.load(X + (hz * L + rows[:, None]) * D + cols[None, :], rows[:, None] < L, 0).to(tl.float32)
tl.store(Y + (hz * L + rows[:, None]) * D + cols[None, :], (x / scale[None, :]).to(tl.float8e4nv), rows[:, None]
< L)


@triton.jit
def _quantize_v_tiles(X, Y, Scale, L: tl.constexpr, D: tl.constexpr):
tile, hz = tl.program_id(0), tl.program_id(1)
rows = tile * 64 + tl.arange(0, 64)
cols = tl.arange(0, D)
x = tl.load(X + (hz * L + rows[:, None]) * D + cols[None, :]).to(tl.float32)
scale = tl.maximum(tl.max(tl.abs(x), 0) / 448.0, 1e-8)
tl.store(Scale + (hz * (L // 64) + tile) * D + cols, scale)
tl.store(Y + (hz * L + rows[:, None]) * D + cols[None, :], (x / scale[None, :]).to(tl.float8e4nv))


@triton.autotune(
configs=[triton.Config({}, num_warps=w, num_stages=s) for w, s in ((4, 2), (4, 3), (4, 4), (8, 2), (8, 3))],
key=["L", "D"])
@triton.jit
def _sparse_int8_fp8(Q, K, V, QS, KS, VS, Index, Count, VBS, Out, L: tl.constexpr, D: tl.constexpr, H: tl.constexpr,
VB: tl.constexpr, VH: tl.constexpr, VS_ROW: tl.constexpr, VD: tl.constexpr, INT8_QK: tl.constexpr,
FP8_PV: tl.constexpr, V_TILE: tl.constexpr, P_DYNAMIC: tl.constexpr):
tile, hz = tl.program_id(0), tl.program_id(1)
nt: tl.constexpr = L // 64
rows = tile * 64 + tl.arange(0, 64)
cols = tl.arange(0, D)
q = tl.load(Q + (hz * L + rows[:, None]) * D + cols[None, :])
if INT8_QK:
qs = tl.load(QS + hz * L + rows)
nblocks = tl.load(Count + hz * nt + tile)
m = tl.full((64, ), -float("inf"), tl.float32)
den = tl.zeros((64, ), tl.float32)
acc = tl.zeros((64, D), tl.float32)
for block in range(nblocks):
kv = tl.load(Index + (hz * nt + tile) * nt + block)
key_rows = kv * 64 + tl.arange(0, 64)
k = tl.load(K + (hz * L + key_rows[None, :]) * D + cols[:, None])
if INT8_QK:
ks = tl.load(KS + hz * L + key_rows)
valid = tl.load(VBS + kv)
if valid > 0:
logits = tl.dot(q, k).to(tl.float32)
if INT8_QK:
logits = logits * qs[:, None] * ks[None, :]
logits = logits * (1.4426950408889634 / D**0.5)
logits = tl.where((tl.arange(0, 64) < valid)[None, :], logits, -float("inf"))
block_max = tl.max(logits, 1)
new_m = tl.maximum(m, block_max)
p = tl.exp2(logits - new_m[:, None])
alpha = tl.exp2(m - new_m)
den = den * alpha + tl.sum(p, 1)
acc = acc * alpha[:, None]
v = tl.load(V + (hz // H) * VB + (hz % H) * VH + key_rows[:, None] * VS_ROW + cols[None, :] * VD)
if FP8_PV:
if P_DYNAMIC:
pscale = tl.maximum(tl.exp2(block_max - new_m) / 448.0, 1e-30)
pv = tl.dot((p / pscale[:, None]).to(tl.float8e4nv), v, out_dtype=tl.float32)
pv = pv * (pscale[:, None] * 448.0)
else:
pv = tl.dot((p * 448.0).to(tl.float8e4nv), v, out_dtype=tl.float32)
if V_TILE:
scale = tl.load(VS + (hz * nt + kv) * D + cols)
acc += pv * (scale[None, :] / 448.0)
else:
acc += pv
else:
acc += tl.dot(p.to(tl.bfloat16), v, out_dtype=tl.float32)
m = new_m
result = acc / den[:, None]
if FP8_PV and not V_TILE:
vs = tl.load(VS + hz * D + cols)
result = result * (vs[None, :] / 448.0)
result = tl.where(den[:, None] > 0, result, 0.0)
tl.store(Out + (hz * L + rows[:, None]) * D + cols[None, :], result.to(Out.dtype.element_ty))


def sparse_sm89_attention(q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor,
vbs: torch.Tensor,
*,
int8_qk: bool = True,
fp8_pv: bool = False,
fp8_v_tiles: bool = False,
fp8_dynamic_p: bool = False) -> torch.Tensor:
"""Forward-only ``[B,H,S,128]`` BF16 attention on sm89, with 64-token tiles."""
if torch.is_grad_enabled():
raise ValueError("Sparse INT8/FP8 attention is inference-only")
if not q.is_cuda or torch.cuda.get_device_capability(q.device) != (8, 9):
raise ValueError("Sparse INT8/FP8 attention requires sm89 CUDA")
if q.dtype != torch.bfloat16 or q.shape[-1] != 128 or q.shape != k.shape or q.shape != v.shape:
raise ValueError("Sparse INT8/FP8 attention requires matching BF16 Q/K/V with head dimension 128")
b, h, length, dim = q.shape
if length != vbs.numel() * 64 or mask.shape != (b, h, length // 64, length // 64):
raise ValueError("Sparse INT8/FP8 attention requires a tile-64 mask and validity vector")
from fastvideo_kernel.triton_kernels.index import map_to_index

# The production INT8-QK/BF16-PV route reads BSHD-backed views directly.
# Quantized Q/K and the output remain contiguous BHSD. Other ablations
# retain their established layout and arithmetic.
if not int8_qk or fp8_pv:
q, k, v = q.contiguous(), k.contiguous(), v.contiguous()
vbs = vbs.to(device=q.device, dtype=torch.int32).contiguous()
grid = (triton.cdiv(length, 16), b * h)
qi, ki, vf = q, k, v
qs, ks, vs = q, k, v # unused pointers in BF16 ablations
if int8_qk:
# Tile pads are zero by contract; avoid a full FP32 copy for the reduction.
# Preserve the exact reduction used by the old contiguous adapter;
# its temporary copy dies before Q/K quantization and fine attention.
mean = k.contiguous().sum(dim=2, dtype=torch.float32) / vbs.sum().clamp_min(1)
qi = torch.empty(q.shape, device=q.device, dtype=torch.int8)
ki = torch.empty(k.shape, device=k.device, dtype=torch.int8)
qs = torch.empty((b, h, length), device=q.device, dtype=torch.float32)
ks = torch.empty_like(qs)
_quantize_qk[grid](q, mean, vbs, qi, qs, length, dim, h, *q.stride(), CENTER=False, ROWS=16, num_warps=4)
_quantize_qk[grid](k, mean, vbs, ki, ks, length, dim, h, *k.stride(), CENTER=True, ROWS=16, num_warps=4)
if fp8_pv:
vf = torch.empty_like(v, dtype=torch.float8_e4m3fn)
if fp8_v_tiles:
vs = torch.empty((b, h, length // 64, dim), device=q.device, dtype=torch.float32)
_quantize_v_tiles[(length // 64, b * h)](v, vf, vs, length, dim, num_warps=8)
else:
vs = (v.abs().amax(dim=2).float() / 448).clamp_min(1e-8)
_quantize_v[grid](v, vs, vf, length, dim, ROWS=16, num_warps=4)
index, count = map_to_index(mask.contiguous())
out = torch.empty(q.shape, device=q.device, dtype=q.dtype)
_sparse_int8_fp8[(length // 64, b * h)](qi,
ki,
vf,
qs,
ks,
vs,
index,
count,
vbs,
out,
length,
dim,
h,
*vf.stride(),
INT8_QK=int8_qk,
FP8_PV=fp8_pv,
V_TILE=fp8_v_tiles,
P_DYNAMIC=fp8_dynamic_p)
return out
53 changes: 40 additions & 13 deletions fastvideo/attention/backends/video_sparse_attn_h3.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,7 @@

import functools
import math
import os
from dataclasses import dataclass
from typing import Any

Expand Down Expand Up @@ -547,6 +548,9 @@ def __init__(
self.prefix = prefix
self.layer_idx = layer_idx_from_prefix(prefix, default=-1)
self.head_size = head_size
self._sm89_kernel = os.environ.get("FASTVIDEO_H3_VSA_SM89_KERNEL", "original")
if self._sm89_kernel not in {"original", "bf16", "int8"}:
raise ValueError("FASTVIDEO_H3_VSA_SM89_KERNEL must be original, bf16, or int8")
# Generic torch.compile must not specialize the shared VSA forward on
# the Python ``layer_idx`` value of each of H3's 50 blocks. This
# tensor is prepared after weights load and drives only the compiled
Expand Down Expand Up @@ -778,9 +782,14 @@ def forward( # type: ignore[override]
# kernels' granularity. These entries take BHSD ([B, H, S_pad, D]);
# mirror block_sparse_attn_256_bshd's Triton branch and transpose
# around the call.
q_bhsd = query.transpose(1, 2).contiguous()
k_bhsd = key.transpose(1, 2).contiguous()
v_bhsd = value.transpose(1, 2).contiguous()
sm89_strided = (self._sm89_kernel == "int8" and not torch.is_grad_enabled() and not compiling
and query.dtype == torch.bfloat16 and query.shape[-1] == 128
and torch.cuda.get_device_capability(query.device) == (8, 9))
q_bhsd = query.transpose(1, 2)
k_bhsd = key.transpose(1, 2)
v_bhsd = value.transpose(1, 2)
if not sm89_strided:
q_bhsd, k_bhsd, v_bhsd = (t.contiguous() for t in (q_bhsd, k_bhsd, v_bhsd))

sm100a_mask = mask
sm100a_variable_block_sizes = attn_metadata.variable_block_sizes
Expand Down Expand Up @@ -877,19 +886,37 @@ def forward( # type: ignore[override]
)
else:
if has_sm100a_pair:
q_bhsd = q_bhsd[:, :, :logical_seq_len].contiguous()
k_bhsd = k_bhsd[:, :, :logical_seq_len].contiguous()
v_bhsd = v_bhsd[:, :, :logical_seq_len].contiguous()
out_bhsd, _ = block_sparse_attn_64_bhsd(
q_bhsd,
k_bhsd,
v_bhsd,
mask,
attn_metadata.variable_block_sizes,
)
q_bhsd = q_bhsd[:, :, :logical_seq_len]
k_bhsd = k_bhsd[:, :, :logical_seq_len]
v_bhsd = v_bhsd[:, :, :logical_seq_len]
if not sm89_strided:
q_bhsd, k_bhsd, v_bhsd = (t.contiguous() for t in (q_bhsd, k_bhsd, v_bhsd))
if (self._sm89_kernel != "original" and not torch.is_grad_enabled() and not compiling
and q_bhsd.dtype == torch.bfloat16 and q_bhsd.shape[-1] == 128
and torch.cuda.get_device_capability(q_bhsd.device) == (8, 9)):
from fastvideo.attention.backends.minimax_h3_sparse_int8 import sparse_sm89_attention
logger.info_once(f"MiniMax-H3 VSA tile-64 forward: sm89 {self._sm89_kernel} QK / BF16 PV")
out_bhsd = sparse_sm89_attention(q_bhsd,
k_bhsd,
v_bhsd,
mask,
attn_metadata.variable_block_sizes,
int8_qk=self._sm89_kernel == "int8",
fp8_pv=False)
else:
out_bhsd, _ = block_sparse_attn_64_bhsd(
q_bhsd,
k_bhsd,
v_bhsd,
mask,
attn_metadata.variable_block_sizes,
)
if has_sm100a_pair and use_sm100a:
out_bhsd = out_bhsd[:, :, :logical_seq_len]
out = out_bhsd.transpose(1, 2).contiguous()
# Fine attention is complete. Release its layout copies before the
# gated compression merge creates full-sequence temporaries.
del q_bhsd, k_bhsd, v_bhsd, out_bhsd
else:
out, _ = block_sparse_attn_256_bshd(
logical_query,
Expand Down
59 changes: 40 additions & 19 deletions fastvideo/hooks/layerwise_offload.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
import torch
from torch import nn
from fastvideo.hooks.hooks import ForwardHook, ModuleHookManager
from fastvideo.hooks.pinned_memory import PinnedTensorArena
from fastvideo.logger import init_logger

logger = init_logger(__name__)
Expand Down Expand Up @@ -53,18 +54,30 @@ def __init__(
self.cpu_named_parameters: dict[str, torch.Tensor] = {}
self.module_ref: nn.Module = None # type: ignore
self.device: torch.device = device
self.cpu_arena: PinnedTensorArena | None = None

def _will_offload(self, name: str) -> bool:
return True

@torch.compiler.disable
def on_init(self, module: nn.Module):
self.module_ref = module
self.clear_cpu_storage()
self.cpu_arena = PinnedTensorArena(
(name, param) for name, param in _offload_tensors(module) if self._will_offload(name))
for name, param in _offload_tensors(self.module_ref):
if self._will_offload(name):
self.cpu_named_parameters[name] = (param.data.detach().to("cpu").pin_memory())
host = self.cpu_arena.empty_like(name, param)
host.copy_(param.data.detach())
self.cpu_named_parameters[name] = host
param.data = _tensor_placeholder(param.data, self.device)

def clear_cpu_storage(self) -> None:
self.cpu_named_parameters.clear()
if self.cpu_arena is not None:
self.cpu_arena.close()
self.cpu_arena = None

@torch.compiler.disable
def wait_and_replace_params(self):
torch.cuda.current_stream().wait_stream(self.async_copy_stream)
Expand Down Expand Up @@ -108,16 +121,17 @@ def on_attach(self, module: nn.Module):
self.state.on_init(module) # pyright: ignore

def on_detach(self, module: nn.Module):
self.state.async_copy_stream.synchronize()
named_parameters = dict(_offload_tensors(module, self.state.cpu_named_parameters))
for name, cpu_tensor in self.state.cpu_named_parameters.items():
if name not in self.state.gpu_named_parameters:
if name in named_parameters:
named_parameters[name].data = cpu_tensor.to(device=self.state.device)
else:
logger.warning(
"Parameter {} not found in module during detachment.",
name,
)
if name in named_parameters:
gpu_tensor = self.state.gpu_named_parameters.get(name)
named_parameters[name].data = gpu_tensor if gpu_tensor is not None else cpu_tensor.to(self.state.device)
else:
logger.warning("Parameter %s not found in module during detachment.", name)
self.state.gpu_named_parameters.clear()
self.state.clear_cpu_storage()
self.state.next_state = None

@classmethod
def name(cls) -> str:
Expand Down Expand Up @@ -154,12 +168,15 @@ def mutate_params_scope(self):
yield
finally:
# instead of releasing, we should overwrite the original params since they have been modified
self.state.cpu_named_parameters.clear()
self.state.gpu_named_parameters.clear()
self.state.on_init(self.state.module_ref) # pyright: ignore


def enable_layerwise_offload(model: nn.Module, is_replace: bool = False):
def enable_layerwise_offload(model: nn.Module,
is_replace: bool = False,
*,
resident_blocks: int | None = None,
cyclic: bool = True):
if torch.cuda.is_available():
device = torch.device("cuda", torch.cuda.current_device())
else:
Expand All @@ -170,12 +187,15 @@ def enable_layerwise_offload(model: nn.Module, is_replace: bool = False):
# The first N entries skip offloading and stay wherever the model is placed (normally the
# GPU), so a GPU with spare memory streams only the remainder over PCIe.
import os
try:
resident = max(0, int(os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS", "0")))
except ValueError:
logger.warning("Ignoring malformed FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=%r",
os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS"))
resident = 0
if resident_blocks is not None:
resident = max(0, resident_blocks)
else:
try:
resident = max(0, int(os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS", "0")))
except ValueError:
logger.warning("Ignoring malformed FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS=%r",
os.environ.get("FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS"))
resident = 0
for name, submodule in model.named_children():
if isinstance(submodule, nn.ModuleList):
for idx, module_entry in enumerate(submodule):
Expand All @@ -201,6 +221,7 @@ def enable_layerwise_offload(model: nn.Module, is_replace: bool = False):
return
raise ValueError("No nn.ModuleList found in the model for layerwise offloading.")

# circular linking of states
# Repeated DiT steps prefetch the first block after the last. A once-per-request
# encoder can skip that unused copy and release every layer after its forward.
for i in range(len(state_list)):
state_list[i].next_state = state_list[(i + 1) % len(state_list)]
state_list[i].next_state = state_list[(i + 1) % len(state_list)] if cyclic or i + 1 < len(state_list) else None
Loading
Loading