Skip to content
28 changes: 22 additions & 6 deletions fastvideo/hooks/layerwise_offload.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@ def __init__(
async_copy_stream: torch.cuda.Stream,
device: torch.device,
next_state: "LayerwiseOffloadState | None" = None,
pin_cpu_memory: bool = True,
) -> None:
self.async_copy_stream = async_copy_stream
self.next_state = next_state
Expand All @@ -55,6 +56,7 @@ def __init__(
self.module_ref: nn.Module = None # type: ignore
self.device: torch.device = device
self.cpu_arena: PinnedTensorArena | None = None
self.pin_cpu_memory = pin_cpu_memory

def _will_offload(self, name: str) -> bool:
return True
Expand All @@ -63,12 +65,18 @@ def _will_offload(self, name: str) -> bool:
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))
if self.pin_cpu_memory:
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):
host = self.cpu_arena.empty_like(name, param)
host.copy_(param.data.detach())
if self.cpu_arena is not None:
host = self.cpu_arena.empty_like(name, param)
host.copy_(param.data.detach())
else:
# Retain checkpoint-backed CPU storage so the OS can reclaim
# inactive file pages instead of holding an anonymous pinned copy.
host = param.data.detach().to("cpu")
self.cpu_named_parameters[name] = host
param.data = _tensor_placeholder(param.data, self.device)

Expand Down Expand Up @@ -176,7 +184,13 @@ def enable_layerwise_offload(model: nn.Module,
is_replace: bool = False,
*,
resident_blocks: int | None = None,
cyclic: bool = True):
cyclic: bool = True,
pin_cpu_memory: bool = True):
"""Stream blocks, optionally retaining their existing pageable CPU storage.

Disabling pinning avoids a private copy of file-backed checkpoint tensors.
Transfers can be slower, but inactive checkpoint pages remain reclaimable.
"""
if torch.cuda.is_available():
device = torch.device("cuda", torch.cuda.current_device())
else:
Expand All @@ -199,7 +213,9 @@ def enable_layerwise_offload(model: nn.Module,
for idx, module_entry in enumerate(submodule):
if idx < resident:
continue
state = LayerwiseOffloadState(async_copy_stream=async_stream, device=device)
state = LayerwiseOffloadState(async_copy_stream=async_stream,
device=device,
pin_cpu_memory=pin_cpu_memory)
state_list.append(state)
hook_mgr = ModuleHookManager.get_from_or_default(module_entry)
hook = LayerwiseOffloadHook(state)
Expand Down
1 change: 1 addition & 0 deletions fastvideo/models/encoders/minimax_h3_checkpoint_nvfp4.py
Original file line number Diff line number Diff line change
Expand Up @@ -346,6 +346,7 @@ def load(param: torch.Tensor, loaded_weight: torch.Tensor, *args: Any, **kwargs:
f"got {loaded_weight.dtype} for a {param.dtype} parameter of shape {tuple(param.shape)}")
return base_loader(param, loaded_weight, *args, **kwargs)

load._h3_base_loader = base_loader
return load


Expand Down
31 changes: 28 additions & 3 deletions fastvideo/models/encoders/minimax_h3_qwen3_vl.py
Original file line number Diff line number Diff line change
Expand Up @@ -755,7 +755,7 @@ def encode_ids(
raise RuntimeError(f"MiniMax-H3 language model returned unexpected shape={tuple(hidden_states.shape)}")
return hidden_states[0]

def prepare_layerwise_offload(self, device: torch.device) -> None:
def prepare_layerwise_offload(self, device: torch.device, *, pin_cpu_memory: bool = True) -> None:
"""Stream language layers for text-only CUDA inference, retaining embeddings on CPU."""
if getattr(self, "_h3_encoder_layerwise_device", None) is not None:
return
Expand All @@ -770,7 +770,8 @@ def prepare_layerwise_offload(self, device: torch.device) -> None:
self.language_model.rotary_emb.to(device)
if self.language_model.norm is not None:
self.language_model.norm.to(device)
enable_layerwise_offload(self.language_model, resident_blocks=0, cyclic=False)
enable_layerwise_offload(self.language_model, resident_blocks=0, cyclic=False,
pin_cpu_memory=pin_cpu_memory)
self._h3_encoder_layerwise_device = device

def forward(
Expand Down Expand Up @@ -804,10 +805,34 @@ def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
raise ValueError(f"Unexpected MiniMax-H3 Qwen3-VL checkpoint key: {source_name}")
parameter = parameters[name]
loader = getattr(parameter, "weight_loader", default_weight_loader)
loader(parameter, tensor)
base_loader = getattr(loader, "_h3_base_loader", loader)
copy_only = (base_loader is default_weight_loader or getattr(base_loader, "__func__", None) in (
ColumnParallelLinear.weight_loader, RowParallelLinear.weight_loader,
VocabParallelEmbedding.weight_loader))
if getattr(base_loader, "__func__", None) is VocabParallelEmbedding.weight_loader:
output_dim = getattr(parameter, "output_dim", None)
copy_only = (not getattr(parameter, "is_gguf_weight_type", False)
and getattr(parameter, "packed_dim", None) is None
and (output_dim is None or (tensor.ndim > output_dim
and tensor.shape[output_dim] == base_loader.__self__.org_vocab_size)))
if (getattr(self, "_h3_checkpoint_backed_cpu", False) and copy_only
and parameter.device.type == tensor.device.type == "cpu"
and parameter.shape == tensor.shape and parameter.dtype == tensor.dtype):
# TP=1 loaders only copy already matching tensors. Keep the mapping
# so the OS can reclaim encoder checkpoint pages during denoising.
# Preserve the Parameter and its loader/quantization attributes.
parameter.data = tensor.detach()
else:
loader(parameter, tensor)
loaded.add(name)
return loaded

def enable_checkpoint_backed_cpu_load(self) -> None:
"""Retain immutable CPU checkpoint storage for single-GPU streamed inference."""
if get_tp_world_size() != 1:
raise ValueError("Checkpoint-backed H3 encoder requires tensor parallel size 1")
self._h3_checkpoint_backed_cpu = True

def _is_omitted_checkpoint_key(self, name: str) -> bool:
"""Return whether a valid checkpoint key belongs to an unbuilt layer."""
language_model = self.language_model
Expand Down
16 changes: 13 additions & 3 deletions fastvideo/models/loader/component_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -422,6 +422,12 @@ def load_model(
with target_device:
model = model_cls(model_config) # type: ignore

retain_checkpoint = getattr(model, "enable_checkpoint_backed_cpu_load", None)
checkpoint_backed_cpu = (target_device.type == "cpu" and envs.FASTVIDEO_H3_ENCODER_LAYERWISE.get()
and callable(retain_checkpoint) and not fastvideo_args.pin_cpu_memory)
if checkpoint_backed_cpu:
retain_checkpoint()

weights_to_load = {name for name, _ in model.named_parameters()}
if (use_text_encoder_override and fastvideo_args.override_text_encoder_safetensors is not None):
if os.path.isdir(checkpoint_path):
Expand Down Expand Up @@ -462,7 +468,11 @@ def load_model(
f"checkpoint: {weights_not_loaded}")

if checkpoint_quant_config is not None:
processed_linears = _process_quantized_text_encoder_weights(model, runtime_device)
# NVFP4 validation and scalar derivation work on the host. Moving
# packed layers to CUDA and back would discard checkpoint mappings.
process_device = (target_device if checkpoint_backed_cpu and checkpoint_quant_config.get_name() == "nvfp4"
else runtime_device)
processed_linears = _process_quantized_text_encoder_weights(model, process_device)
logger.info("Validated %d serialized %s text-encoder linears", processed_linears,
checkpoint_quant_config.get_name())

Expand All @@ -473,7 +483,7 @@ def load_model(
if envs.FASTVIDEO_H3_ENCODER_LAYERWISE.get() and callable(prepare_layerwise):
if target_device.type != "cpu":
raise ValueError("Layerwise H3 encoder requires text_encoder_cpu_offload=True")
prepare_layerwise(runtime_device)
prepare_layerwise(runtime_device, pin_cpu_memory=fastvideo_args.pin_cpu_memory)
use_cpu_offload = False
logger.info("Enabled text-only layerwise H3 encoder with CPU token embeddings")

Expand Down Expand Up @@ -1211,7 +1221,7 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs):
# Check if model has nn.ModuleList for layerwise offload compatibility
has_module_list = any(isinstance(m, nn.ModuleList) for m in model.children())
if has_module_list:
enable_layerwise_offload(model)
enable_layerwise_offload(model, pin_cpu_memory=fastvideo_args.pin_cpu_memory)
# Blocks now hold placeholders; the remaining (non-block) weights and buffers belong on the GPU.
model = model.to(get_local_torch_device())
else:
Expand Down
4 changes: 3 additions & 1 deletion fastvideo/models/loader/fsdp_load.py
Original file line number Diff line number Diff line change
Expand Up @@ -360,7 +360,9 @@ def maybe_load_fsdp_model(
packed_nvfp4_export,
)
packed_nvfp4_export = None
load_weights_to_cpu = cpu_offload or packed_nvfp4_export is not None
# CPU-targeted layerwise/table loads must not stage the full checkpoint on
# the GPU before offload hooks or skipped projection weights are applied.
load_weights_to_cpu = device.type == "cpu" or cpu_offload or packed_nvfp4_export is not None
weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=load_weights_to_cpu)
logger.info("Loading transformer weights with to_cpu=%s", load_weights_to_cpu)
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
Expand Down
31 changes: 27 additions & 4 deletions fastvideo/tests/encoders/test_minimax_h3_encoder_layerwise.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,8 +20,9 @@

@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for encoder streaming")
@pytest.mark.parametrize("quantized,fused", [(False, False), (True, False), (True, True)])
@pytest.mark.parametrize("pin_cpu_memory,checkpoint_backed", [(True, False), (False, False), (False, True)])
def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup, monkeypatch, env_overrides, quantized,
fused):
fused, pin_cpu_memory, checkpoint_backed, tmp_path):
# DiT residency must not accidentally keep encoder layers resident too.
env_overrides.enter_context(envs.FASTVIDEO_LAYERWISE_RESIDENT_BLOCKS.override(6))
env_overrides.enter_context(envs.FASTVIDEO_H3_ENCODER_FUSED_DEQUANT.override(False))
Expand Down Expand Up @@ -54,6 +55,11 @@ def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup
with patch.object(torch.Tensor, "item", side_effect=AssertionError("Unexpected device scalar read")):
actual_linear = linear(x)[0]
torch.testing.assert_close(actual_linear, expected_linear, rtol=0, atol=0)
checkpoint = tmp_path / "encoder.safetensors"
if checkpoint_backed:
from safetensors.torch import save_file
save_file({name: parameter.detach().cpu().contiguous() for name, parameter in model.named_parameters()},
checkpoint)
ids = torch.tensor([1, 7, 4, 21, 5, 31, 18], device="cuda")
model.to("cuda")
expected = model.encode_ids(ids)
Expand All @@ -63,8 +69,18 @@ def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup
if hasattr(layer, "_nvfp4_fused_dequant"):
layer._nvfp4_fused_dequant = True
model.to("cpu")
model.prepare_layerwise_offload(torch.device("cuda"))
model.prepare_layerwise_offload(torch.device("cuda")) # repeated setup is harmless
mapped = {}
if checkpoint_backed:
from safetensors.torch import load_file
mapped = load_file(checkpoint, device="cpu")
model.enable_checkpoint_backed_cpu_load()
model.load_weights(mapped.items())
if quantized:
_process_quantized_text_encoder_weights(model, torch.device("cpu"))
for name, parameter in model.named_parameters():
assert parameter.data_ptr() == mapped[name].data_ptr()
model.prepare_layerwise_offload(torch.device("cuda"), pin_cpu_memory=pin_cpu_memory)
model.prepare_layerwise_offload(torch.device("cuda"), pin_cpu_memory=pin_cpu_memory) # repeated setup is harmless
assert model.language_model.embed_tokens.weight.device.type == "cpu"
assert next(model.visual.parameters()).device.type == "cpu"
for _ in range(2):
Expand All @@ -74,7 +90,14 @@ def test_streamed_encoder_matches_resident_and_releases_layers(distributed_setup
assert all(parameter.numel() == 0 for parameter in layer.parameters())
manager = ModuleHookManager.get_from(layer)
assert manager is not None
assert not manager.forward_hooks["LayerwiseOffloadHook"].state.gpu_named_parameters
state = manager.forward_hooks["LayerwiseOffloadHook"].state
assert not state.gpu_named_parameters
assert state.pin_cpu_memory == pin_cpu_memory
assert (state.cpu_arena is not None) == pin_cpu_memory
if checkpoint_backed:
for name, tensor in state.cpu_named_parameters.items():
layer_name = next(key for key, value in model.named_modules() if value is layer)
assert tensor.data_ptr() == mapped[f"{layer_name}.{name}"].data_ptr()
with pytest.raises(ValueError, match="text-only"):
model.encode_ids(ids, pixel_values=torch.zeros(1, device="cuda"),
image_grid_thw=torch.ones(1, 3, device="cuda", dtype=torch.int64))
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -454,10 +454,13 @@ def _tiny_conditioner_config(keep_bf16: tuple[str, ...] = ("mlp.down_proj", )) -
return config


def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup) -> None:
@pytest.mark.parametrize("checkpoint_backed", [False, True])
def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup, checkpoint_backed) -> None:
"""The real chain: a conditioner built with the NVFP4 config, checkpoint keys spelled the way the
converter writes them, ``load_weights``, the strict missing-tensor check, then the post-load hook."""
model = MiniMaxH3Qwen3VLConditioner(_tiny_conditioner_config())
if checkpoint_backed:
model.enable_checkpoint_backed_cpu_load()
layer = model.language_model.layers[0]
assert isinstance(layer.self_attn.q_proj.quant_method, MiniMaxH3SerializedNVFP4LinearMethod)
assert isinstance(layer.self_attn.o_proj.quant_method, MiniMaxH3SerializedNVFP4LinearMethod)
Expand All @@ -482,6 +485,9 @@ def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup)
loaded = model.load_weights(iter(checkpoint.items()))
assert loaded == expected
assert _process_quantized_text_encoder_weights(model, torch.device("cpu")) == 6
if checkpoint_backed:
for name, parameter in model.named_parameters():
assert parameter.data_ptr() == checkpoint[f"model.{name}"].data_ptr()
assert layer.self_attn.q_proj._nvfp4_alpha.item() == pytest.approx(0.5)
assert layer.mlp.up_proj._nvfp4_alpha.item() == pytest.approx(0.5)

Expand All @@ -504,3 +510,32 @@ def test_conditioner_loads_converter_named_tensors_end_to_end(distributed_setup)
stray["model.language_model.layers.0.self_attn.q_proj.weight"] = torch.zeros(128, 128)
with pytest.raises(ValueError, match="Unexpected"):
model.load_weights(iter(stray.items()))


def test_checkpoint_backed_load_keeps_custom_loader_semantics(distributed_setup) -> None:
model = MiniMaxH3Qwen3VLConditioner(_tiny_conditioner_config())
model.enable_checkpoint_backed_cpu_load()
name = "language_model.layers.0.input_layernorm.weight"
parameter = dict(model.named_parameters())[name]
source = torch.ones_like(parameter)

def custom_loader(target, value):
target.data.copy_(value + 1)

parameter.weight_loader = custom_loader
model.load_weights([(name, source)])
torch.testing.assert_close(parameter, source + 1, rtol=0, atol=0)
assert parameter.data_ptr() != source.data_ptr()


def test_checkpoint_backed_embedding_rejects_padded_checkpoint_rows(distributed_setup) -> None:
config = _tiny_conditioner_config()
config.arch_config.vocab_size = 63
model = MiniMaxH3Qwen3VLConditioner(config)
model.enable_checkpoint_backed_cpu_load()
weight = model.language_model.embed_tokens.weight
assert weight.shape[0] == 64
# Matching the allocated padded shape must not bypass the loader's check
# that checkpoint rows match the original unpadded vocabulary.
with pytest.raises(AssertionError):
model.load_weights([("language_model.embed_tokens.weight", torch.ones_like(weight))])
30 changes: 30 additions & 0 deletions fastvideo/tests/hooks/test_layerwise_offload.py
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,36 @@ def forward(self, x: torch.Tensor) -> torch.Tensor:
return x


@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_pageable_offload_retains_mapped_storage_and_matches_repeated_forwards(tmp_path):
"""File-backed weights stay reclaimable; GPU prefetch preserves their values."""
model = SimpleModelWithModuleList(num_blocks=3, hidden_size=32)
pointers = {}
for block_index, block in enumerate(model.blocks):
for name, parameter in block.named_parameters():
path = tmp_path / f"{block_index}-{name}.bin"
mapped = torch.from_file(str(path), shared=True, size=parameter.numel(), dtype=parameter.dtype)
mapped.copy_(parameter.detach().reshape(-1))
parameter.data = mapped.view_as(parameter)
pointers[block_index, name] = parameter.data_ptr()
reference = SimpleModelWithModuleList(num_blocks=3, hidden_size=32).cuda()
reference.load_state_dict(model.state_dict())
x = torch.randn(2, 9, 32, device="cuda")
with torch.inference_mode():
expected = reference(x)
enable_layerwise_offload(model, pin_cpu_memory=False)
for block_index, block in enumerate(model.blocks):
manager = ModuleHookManager.get_from(block)
state = manager.forward_hooks["LayerwiseOffloadHook"].state
assert state.cpu_arena is None
for name, host in state.cpu_named_parameters.items():
assert not host.is_pinned()
assert host.data_ptr() == pointers[block_index, name]
with torch.inference_mode():
for _ in range(3):
torch.testing.assert_close(model(x), expected, rtol=0, atol=0)


@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_layerwise_offload_basic():
"""Test basic functionality of layerwise offloading."""
Expand Down
Loading
Loading