diff --git a/.buildkite/scripts/lanes/lora_extraction.sh b/.buildkite/scripts/lanes/lora_extraction.sh index 46d7e0ec22..4771847dc4 100755 --- a/.buildkite/scripts/lanes/lora_extraction.sh +++ b/.buildkite/scripts/lanes/lora_extraction.sh @@ -2,4 +2,4 @@ # Canonical Slurm CI selection for the LoRA-extraction lane. set -euo pipefail -exec pytest ./fastvideo/tests/lora_extraction/test_lora_extraction.py -vs +exec pytest ./fastvideo/tests/lora_extraction/ -vs diff --git a/docs/training/finetune.md b/docs/training/finetune.md index 7389242bf6..722ab8feb1 100644 --- a/docs/training/finetune.md +++ b/docs/training/finetune.md @@ -110,7 +110,7 @@ Key differences from full finetune: ## LoRA Extraction and Merging -FastVideo provides tools to extract LoRA adapters from finetuned models and merge them back. +FastVideo provides generic runtime adapter extraction and retains a legacy merger for its previously supported adapter layouts. ### Extract LoRA Adapter @@ -121,25 +121,41 @@ python scripts/lora_extraction/extract_lora.py \ --base Wan-AI/Wan2.1-T2V-1.3B-Diffusers \ --ft path/to/your/finetuned_model \ --out adapter_r32.safetensors \ - --rank 32 + --rank 32 \ + --exact-tensor-pattern '^condition_embedder\.' \ + --exact-tensor-pattern '^proj_out\.weight$' ``` | Argument | Description | |----------|-------------| -| `--base` | Base model (HuggingFace ID or local path) | -| `--ft` | Finetuned model path | -| `--out` | Output adapter file (.safetensors) | -| `--rank` | LoRA rank (16, 32, 64, 128) | +| `--base` | Base model (Hugging Face ID or local path) | +| `--ft` | Fine-tuned model (Hugging Face ID or local path) | +| `--out` | Output adapter file (`.safetensors`) | +| `--rank` | Requested LoRA rank (for example, 16, 32, 64, or 128) | | `--full-rank` | Extract full-rank adapter (optional) | +| `--min-delta` | Omit tensors whose maximum absolute FP32 delta is at or below the threshold (default: `1e-8`) | +| `--load-mode` | `auto` (indexed, then pipeline fallback), `indexed`, or `pipeline` | +| `--device` | SVD device, such as `cpu` or `cuda:0` | +| `--svd-method` | Exact or randomized SVD | +| `--factor-dtype` | Storage dtype for the low-rank factors | +| `--dense-dtype` | Storage dtype for exact `.diff`/`.diff_b`/`.diff_param` payloads (default: `float32`) | +| `--exact-tensor-pattern` | Repeatable regex for matrices the target runtime cannot load as LoRA factors | -### Merge LoRA Adapter +For large checkpoints, indexed loading streams one transformer tensor pair at a time and downloads only `transformer/*`. The extractor also preserves changed norms, biases, and standalone parameters as exact deltas, and fine-tuned-only parameters as `.set_weight` or `.set_param`. Matrix selection is runtime-agnostic, so full-finetune extraction must keep runtime-unsupported matrices exact, as the Wan example does. See the [LoRA utilities](../utilities/lora.md) for the GPU/randomized-SVD command, resume options, and accuracy controls. -Merge an adapter back into a base model: +Mixed low-rank/dense adapters from the generic extractor must be supplied when constructing FastVideo through +`ComponentConfig(lora_path=...)`; their dense payload cannot be swapped later with `set_lora_adapter`. The legacy +offline merger below is not part of this extraction workflow. + +### Legacy Merge LoRA Adapter + +The command below documents the pre-existing merger for adapters it already supports. Do not pass a mixed adapter from +the generic extractor to it: the legacy merger does not apply exact dense or replacement payloads. ```bash python scripts/lora_extraction/merge_lora.py \ --base Wan-AI/Wan2.1-T2V-1.3B-Diffusers \ - --adapter adapter_r32.safetensors \ + --adapter legacy_factor_only_adapter.safetensors \ --ft path/to/your/finetuned_model \ --output merged_model ``` diff --git a/docs/utilities/lora.md b/docs/utilities/lora.md index e9109db7d9..07daf0c90b 100644 --- a/docs/utilities/lora.md +++ b/docs/utilities/lora.md @@ -1,6 +1,6 @@ # LoRA Extraction and Merging -Tools for extracting and merging LoRA adapters for FastVideo models. +Generic runtime adapter extraction plus the existing legacy merge utilities for FastVideo models. ## Extract LoRA Adapter @@ -9,30 +9,75 @@ python scripts/lora_extraction/extract_lora.py \ --base Wan-AI/Wan2.2-TI2V-5B-Diffusers \ --ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \ --out adapter_r32.safetensors \ - --rank 32 + --rank 32 \ + --exact-tensor-pattern '^condition_embedder\.' \ + --exact-tensor-pattern '^proj_out\.weight$' ``` -**Options:** +The extractor is runtime-agnostic and cannot determine from checkpoint tensors whether the target runtime +wraps a given matrix as a LoRA layer. Use `--exact-tensor-pattern` for changed matrices that the runtime does not wrap; +the extractor preserves them as exact `.diff` tensors instead of emitting factors that the runtime cannot apply. The +Wan patterns above cover its excluded condition embedders and its unwrapped output projection. + +Exact CPU SVD remains the default. For a large transformer, stream its indexed safetensors and factorize on a GPU: + +```bash +python scripts/lora_extraction/extract_lora.py \ + --base \ + --ft \ + --out adapter_r64.safetensors \ + --rank 64 \ + --load-mode indexed \ + --device cuda:0 \ + --svd-method randomized \ + --randomized-q 320 \ + --niter 4 \ + --factor-dtype float16 \ + --dense-dtype float32 \ + --replacement-dtype source +``` + +`--load-mode indexed` downloads only `transformer/*` for a Hugging Face model and reads one base/fine-tuned tensor pair at a time. The default `auto` mode tries indexed loading first and falls back to the legacy FastVideo pipeline loader; `pipeline` selects the legacy loader directly. + +Important options: + +- `--base`, `--ft`: Hugging Face model IDs or local paths. +- `--rank`, `--full-rank`: truncated or full factorization rank. +- `--min-delta`: omit tensors whose maximum absolute FP32 delta is at or below this threshold (default: `1e-8`). +- `--device`: factorization device, such as `cpu` or `cuda:0`. +- `--svd-method`: `exact` or `randomized`. +- `--randomized-q`, `--niter`, `--seed`: randomized SVD accuracy and reproducibility. +- `--factor-dtype`, `--dense-dtype`, `--replacement-dtype`: adapter storage precision. Exact dense deltas default to `float32`. +- `--exact-tensor-pattern`: repeatable regex for a matrix that should remain an exact dense delta. +- `--base-revision`, `--ft-revision`: pin Hugging Face inputs in indexed mode; revisions are rejected for local paths and pipeline loading. +- `--work-dir`, `--resume`: resume an interrupted streaming extraction. Scratch is written to an + output-specific namespace under `fastvideo-lora-extract/`, and only that namespace is cleaned up. Resume requires indexed + safetensors and validates both checkpoints' index/shard fingerprints before reusing partial results. + +The adapter retains changes that do not fit a low-rank product: `.diff` and `.diff_b` hold exact additive weight/bias deltas, `.diff_param` handles standalone parameters such as `scale_shift_table`, and `.set_weight`/`.set_param` hold parameters absent from the base checkpoint. Bit-identical parameters are omitted. The extractor writes an adjacent `*.report.json` with tensor counts, settings, and reconstruction residuals. + +For the validated MiniMax-H3 rank-64 command, including its exact-boundary patterns, see [`scripts/lora_extraction/README.md`](https://github.com/hao-ai-lab/FastVideo/blob/main/scripts/lora_extraction/README.md). + +Mixed low-rank/dense adapters produced by the generic extractor must be supplied when constructing FastVideo through +`ComponentConfig(lora_path=...)`; their dense payload cannot be swapped later with `set_lora_adapter`. The legacy +offline merger below retains its existing scope and is not part of this extraction workflow. -- `--base`: Base model (HuggingFace ID or local path) -- `--ft`: Fine-tuned model (HuggingFace ID or local path) -- `--out`: Output adapter file -- `--rank`: LoRA rank (16, 32, 64, 128) -- `--full-rank`: Extract full-rank adapter (optional) +## Legacy Merge Adapter -## Merge Adapter +The command below documents the pre-existing merger for adapters it already supports. Do not pass a mixed adapter from +the generic extractor to it: the legacy merger does not apply the adapter's exact dense or replacement payloads. ```bash python scripts/lora_extraction/merge_lora.py \ --base Wan-AI/Wan2.2-TI2V-5B-Diffusers \ - --adapter adapter_r32.safetensors \ + --adapter legacy_factor_only_adapter.safetensors \ --ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \ --output merged_model ``` **Options:** -- `--base`: Base model (HuggingFace ID or local path) +- `--base`: Base model (Hugging Face ID or local path) - `--adapter`: LoRA adapter file (.safetensors) - `--ft`: Fine-tuned model (for configuration) - `--output`: Output directory diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index 05d8222cad..f5a6c156e6 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -241,11 +241,17 @@ def maybe_load_fsdp_model( weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=True) param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping) - dense_lora_patch = DenseLoRAPatch.from_adapter( - lora_path, - param_names_mapping_fn, - strength=lora_strength, - ) + dense_lora_patch = None + if lora_path: + # LoRA-specific mappings are optional and must not become a dependency of ordinary model loading. + lora_mapping = getattr(model, "lora_param_names_mapping", None) + lora_param_names_mapping_fn = get_param_names_mapping(lora_mapping) if lora_mapping else None + dense_lora_patch = DenseLoRAPatch.from_adapter( + lora_path, + param_names_mapping_fn, + lora_param_names_mapping=lora_param_names_mapping_fn, + strength=lora_strength, + ) if dense_lora_patch is not None: # H3's compression gate is created only by the VSA attention backend. Loading a # VSA student under dense attention would otherwise warn about 50 unmatched @@ -695,7 +701,7 @@ def load_model_from_full_model_state_dict( target_dtype = dtype_selector(new_param_name, param_dtype) if adapter_value is not None: if tuple(adapter_value.shape) != tuple(meta_sharded_param.shape): - raise ValueError(f"LoRA set_weight for {new_param_name} has shape {tuple(adapter_value.shape)}, " + raise ValueError(f"LoRA replacement for {new_param_name} has shape {tuple(adapter_value.shape)}, " f"but the parameter is {tuple(meta_sharded_param.shape)}") full_tensor = adapter_value.to(device=device, dtype=target_dtype) if not hasattr(meta_sharded_param, "device_mesh"): diff --git a/fastvideo/models/loader/lora_patch.py b/fastvideo/models/loader/lora_patch.py index 40925020ef..6154cca941 100644 --- a/fastvideo/models/loader/lora_patch.py +++ b/fastvideo/models/loader/lora_patch.py @@ -6,16 +6,17 @@ rank. Distilled video checkpoints break both often enough that dropping whatever does not fit silently loses real signal. -Two payload kinds cover the gap, named after the convention ComfyUI's loader already -reads so one file works in both places: +Two payload kinds cover the gap. The weight/bias spellings follow conventions used by +ComfyUI; the ``*_param`` spellings extend them to standalone FastVideo parameters: -``.diff`` / ``.diff_b`` +``.diff`` / ``.diff_b`` / ``.diff_param`` An exact additive delta for a parameter the base model has. Used where a rank-``r`` factorization buys nothing or cannot be formed at all -- RMSNorm vectors, biases, - and matrices whose smaller dimension is already at or below the rank that would be - chosen. Factoring a length-``n`` vector into rank ``r`` costs ``r(1 + n) > n``. + standalone parameters such as ``scale_shift_table``, and matrices whose smaller + dimension is already at or below the rank that would be chosen. Factoring a + length-``n`` vector into rank ``r`` costs ``r(1 + n) > n``. -``.set_weight`` +``.set_weight`` / ``.set_param`` A whole parameter the base model does not carry, so no delta is expressible. MiniMax H3's VSA ``to_gate_compress`` is the case that motivated this: it exists only under the sparse-attention backend, and :func:`load_model_from_full_model_state_dict` @@ -49,17 +50,11 @@ logger = init_logger(__name__) -# Suffix -> the parameter suffix it targets. ``.diff``/``.diff_b`` are additive, -# ``.set_weight`` replaces. Ordered longest-first so ``.diff_b`` is tested before -# ``.diff`` would match a truncated key. -ADDITIVE_SUFFIXES: dict[str, str] = {".diff_b": ".bias", ".diff": ".weight"} -REPLACEMENT_SUFFIXES: dict[str, str] = {".set_weight": ".weight"} - -# Recognized elsewhere in an adapter and deliberately not our business: the low-rank -# half, which ``LoRAPipeline`` merges through the wrapped-module path. -_LOW_RANK_MARKERS = (".lora_A", ".lora_B", ".lora_up", ".lora_down", ".lora_alpha", ".lora_rank", ".alpha", - ".dora_scale") - +# Suffix -> the parameter suffix it targets. An empty target suffix preserves the +# full parameter name for standalone nn.Parameters. Ordered longest-first so the +# generic spellings are tested before the shorter weight/bias spellings. +ADDITIVE_SUFFIXES: dict[str, str] = {".diff_param": "", ".diff_b": ".bias", ".diff": ".weight"} +REPLACEMENT_SUFFIXES: dict[str, str] = {".set_weight": ".weight", ".set_param": ""} # One low-rank pair has many spellings. PEFT writes ``.lora_A.weight``, and interposes # the adapter's name when it is not the default (``.lora_A.default.weight``); kohya and @@ -126,14 +121,14 @@ def from_adapter( lora_path: str | None, param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None = None, *, + lora_param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None = None, strength: float = 1.0, ) -> DenseLoRAPatch | None: """Build a patch from an adapter, or ``None`` when it carries no dense payload. - ``param_names_mapping`` is the same callable the checkpoint loader uses, so - adapter keys are resolved into the model's own parameter names by the identical - rules -- an adapter written against the published checkpoint layout needs no - separate conversion table. + ``lora_param_names_mapping`` first translates adapter-specific official names + into the published checkpoint layout. ``param_names_mapping`` then resolves that + layout into the model's parameter names, matching the low-rank loader's order. """ if not lora_path: return None @@ -150,7 +145,7 @@ def from_adapter( for path in files: with safe_open(path, framework="pt") as handle: for key in handle.keys(): - resolved = _resolve(key, param_names_mapping) + resolved = _resolve(key, lora_param_names_mapping, param_names_mapping) if resolved is None: continue target, kind = resolved @@ -163,7 +158,7 @@ def from_adapter( if not additive and not replacement: return None logger.info( - "LoRA adapter %s carries a dense payload: %d additive (.diff/.diff_b), %d replacement (.set_weight)", + "LoRA adapter %s carries a dense payload: %d additive, %d replacement parameters", lora_path, len(additive), len(replacement)) return cls(files, additive, replacement, strength) @@ -234,6 +229,7 @@ def _read(self, entry: tuple[str, str]) -> torch.Tensor: def _resolve( key: str, + lora_param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None, param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None, ) -> tuple[str, str] | None: """Map an adapter key to ``(model parameter name, "add" | "set")``. @@ -241,24 +237,41 @@ def _resolve( Returns ``None`` for anything that is not a dense payload key, which includes every low-rank factor -- those belong to ``LoRAPipeline``, not here. """ - if any(marker in key for marker in _LOW_RANK_MARKERS): - return None + # A terminal dense suffix is authoritative even when an ordinary module name + # contains text such as ``.alpha`` or ``.lora_A``. for suffix, param_suffix in ADDITIVE_SUFFIXES.items(): if key.endswith(suffix): - return _map_name(key[:-len(suffix)] + param_suffix, param_names_mapping, key), "add" + return _map_name( + key[:-len(suffix)] + param_suffix, + lora_param_names_mapping, + param_names_mapping, + key, + ), "add" for suffix, param_suffix in REPLACEMENT_SUFFIXES.items(): if key.endswith(suffix): - return _map_name(key[:-len(suffix)] + param_suffix, param_names_mapping, key), "set" + return _map_name( + key[:-len(suffix)] + param_suffix, + lora_param_names_mapping, + param_names_mapping, + key, + ), "set" return None def _map_name( param_name: str, + lora_param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None, param_names_mapping: Callable[[str], tuple[str, Any, Any]] | None, source_key: str, ) -> str: - """Run the checkpoint loader's own renaming rules over a resolved parameter name.""" + """Run the low-rank loader's two-stage renaming rules over a dense parameter.""" param_name = param_name.replace("diffusion_model.", "") + if lora_param_names_mapping is not None: + param_name, merge_index, _ = lora_param_names_mapping(param_name) + if merge_index is not None: + raise NotImplementedError(f"LoRA dense key {source_key} resolves to a fused parameter during the " + "adapter-specific mapping; whole-tensor payloads for fused parameters " + "are not supported") if param_names_mapping is None: return param_name mapped, merge_index, _ = param_names_mapping(param_name) diff --git a/fastvideo/pipelines/lora_pipeline.py b/fastvideo/pipelines/lora_pipeline.py index b291d6e535..351f77cc13 100644 --- a/fastvideo/pipelines/lora_pipeline.py +++ b/fastvideo/pipelines/lora_pipeline.py @@ -346,12 +346,12 @@ def set_lora_adapter(self, if not self._setting_constructor_adapter: if self._constructor_dense_lora_path is not None: raise RuntimeError( - "The active LoRA contains constructor-time .diff/.set_weight payload. " + "The active LoRA contains a constructor-time dense additive/replacement payload. " "Changing its adapter or strength at runtime would leave that dense payload stale; " "create a new VideoGenerator with ComponentConfig(lora_path=..., lora_strength=...).") if requested_path is not None and DenseLoRAPatch.from_adapter(requested_path) is not None: raise RuntimeError( - "Adapters containing .diff/.set_weight payload must be supplied when VideoGenerator is " + "Adapters containing dense additive/replacement payloads must be supplied when VideoGenerator is " "constructed with ComponentConfig(lora_path=..., lora_strength=...).") if lora_nickname not in self.lora_adapters and lora_path is None: diff --git a/fastvideo/tests/loader/test_lora_patch.py b/fastvideo/tests/loader/test_lora_patch.py index 55e33a38da..8fee777474 100644 --- a/fastvideo/tests/loader/test_lora_patch.py +++ b/fastvideo/tests/loader/test_lora_patch.py @@ -9,8 +9,10 @@ import torch from safetensors.torch import save_file +from fastvideo.configs.models.dits.wanvideo import WanVideoConfig +from fastvideo.models.loader import fsdp_load from fastvideo.models.loader.lora_patch import DenseLoRAPatch, normalize_lora_key - +from fastvideo.models.loader.utils import get_param_names_mapping def write_adapter(tmp_path, tensors, name="adapter_model.safetensors"): @@ -48,6 +50,8 @@ def test_normalize_lora_key_accepts_every_published_spelling(raw, expected): "blocks.0.norm1.diff", "blocks.0.attn.to_q.diff_b", "blocks.0.attn.to_gate_compress.set_weight", + "blocks.0.scale_shift_table.diff_param", + "blocks.0.extra_table.set_param", "blocks.0.attn.to_q.dora_scale", ]) def test_normalize_lora_key_disclaims_non_low_rank_keys(raw): @@ -82,13 +86,29 @@ def test_from_adapter_splits_additive_from_replacement(tmp_path): "blocks.0.attn.to_q.lora_A.weight": torch.zeros(4, 8), "blocks.0.norm1.diff": torch.zeros(8), "blocks.0.attn.to_q.diff_b": torch.zeros(8), + "blocks.0.scale_shift_table.diff_param": torch.zeros(8), "blocks.0.attn.to_gate_compress.set_weight": torch.zeros(8, 8), + "blocks.0.extra_table.set_param": torch.zeros(8), }) patch = DenseLoRAPatch.from_adapter(path) assert patch is not None - # `.diff` targets `.weight`; `.diff_b` targets `.bias`; `.set_weight` targets `.weight`. - assert set(patch._additive) == {"blocks.0.norm1.weight", "blocks.0.attn.to_q.bias"} - assert set(patch._replacement) == {"blocks.0.attn.to_gate_compress.weight"} + # Weight/bias suffixes restore their parameter suffix; generic parameter + # payloads preserve the complete name. + assert set(patch._additive) == { + "blocks.0.norm1.weight", + "blocks.0.attn.to_q.bias", + "blocks.0.scale_shift_table", + } + assert set(patch._replacement) == { + "blocks.0.attn.to_gate_compress.weight", + "blocks.0.extra_table", + } + + +def test_dense_suffix_is_not_hidden_by_an_alpha_like_parameter_name(tmp_path): + path = write_adapter(tmp_path, {"blocks.0.alpha.diff_param": torch.zeros(4)}) + patch = DenseLoRAPatch.from_adapter(path) + assert set(patch._additive) == {"blocks.0.alpha"} def test_param_names_mapping_is_applied_to_dense_keys(tmp_path): @@ -102,6 +122,126 @@ def mapping(name): assert set(patch._additive) == {"blocks.0.ff.fc_in.weight"} +def test_official_wan_mixed_adapter_maps_dense_key_like_factors(tmp_path): + path = write_adapter( + tmp_path, { + "blocks.0.self_attn.q.lora_A.weight": torch.zeros(2, 4), + "blocks.0.self_attn.q.lora_B.weight": torch.zeros(4, 2), + "blocks.0.self_attn.k.diff": torch.zeros(4), + }) + config = WanVideoConfig() + patch = DenseLoRAPatch.from_adapter( + path, + get_param_names_mapping(config.param_names_mapping), + lora_param_names_mapping=get_param_names_mapping(config.lora_param_names_mapping), + ) + assert set(patch._additive) == {"blocks.0.to_k.weight"} + + +def _stub_fsdp_loading(monkeypatch): + monkeypatch.setattr(fsdp_load, "set_mixed_precision_policy", lambda **_: None) + monkeypatch.setattr(fsdp_load, "safetensors_weights_iterator", lambda *_args, **_kwargs: iter(())) + monkeypatch.setattr(fsdp_load, "load_model_from_full_model_state_dict", lambda *_args, **_kwargs: None) + monkeypatch.setattr(fsdp_load, "_maybe_quantize_model", lambda _model: None) + + +def _load_tiny_model(model_cls, *, lora_path=None, lora_strength=1.0): + return fsdp_load.maybe_load_fsdp_model( + model_cls=model_cls, + init_params={}, + weight_dir_list=[], + device=torch.device("cpu"), + hsdp_replicate_dim=1, + hsdp_shard_dim=1, + default_dtype=torch.float32, + param_dtype=torch.float32, + reduce_dtype=torch.float32, + training_mode=False, + pin_cpu_memory=False, + lora_path=lora_path, + lora_strength=lora_strength, + ) + + +def test_fsdp_loader_forwards_lora_specific_mapping(monkeypatch): + captured = {} + + class TinyModel(torch.nn.Module): + param_names_mapping = {r"^hf\.(.*)$": r"custom.\1"} + lora_param_names_mapping = {r"^official\.(.*)$": r"hf.\1"} + + def capture_patch(cls, + lora_path, + param_names_mapping, + *, + lora_param_names_mapping=None, + strength=1.0): + captured["path"] = lora_path + captured["regular"] = param_names_mapping("hf.weight")[0] + captured["lora"] = lora_param_names_mapping("official.weight")[0] + captured["strength"] = strength + return None + + monkeypatch.setattr(DenseLoRAPatch, "from_adapter", classmethod(capture_patch)) + _stub_fsdp_loading(monkeypatch) + + _load_tiny_model(TinyModel, lora_path="adapter.safetensors", lora_strength=0.75) + + assert captured == { + "path": "adapter.safetensors", + "regular": "custom.weight", + "lora": "hf.weight", + "strength": 0.75, + } + + +def test_fsdp_loader_does_not_consult_lora_mapping_without_adapter(monkeypatch): + class TinyModel(torch.nn.Module): + param_names_mapping = {r"^hf\.(.*)$": r"custom.\1"} + + def unexpected_patch(*_args, **_kwargs): + pytest.fail("the no-LoRA load path must not inspect an adapter") + + monkeypatch.setattr(DenseLoRAPatch, "from_adapter", classmethod(unexpected_patch)) + _stub_fsdp_loading(monkeypatch) + + model = _load_tiny_model(TinyModel) + + assert isinstance(model, TinyModel) + assert not hasattr(model, "lora_param_names_mapping") + + +def test_fsdp_loader_allows_adapter_without_lora_specific_mapping(monkeypatch): + captured = {} + + class TinyModel(torch.nn.Module): + param_names_mapping = {r"^hf\.(.*)$": r"custom.\1"} + + def capture_patch(cls, + lora_path, + param_names_mapping, + *, + lora_param_names_mapping=None, + strength=1.0): + captured["path"] = lora_path + captured["regular"] = param_names_mapping("hf.weight")[0] + captured["lora"] = lora_param_names_mapping + captured["strength"] = strength + return None + + monkeypatch.setattr(DenseLoRAPatch, "from_adapter", classmethod(capture_patch)) + _stub_fsdp_loading(monkeypatch) + + _load_tiny_model(TinyModel, lora_path="adapter.safetensors") + + assert captured == { + "path": "adapter.safetensors", + "regular": "custom.weight", + "lora": None, + "strength": 1.0, + } + + def test_fused_target_is_refused_rather_than_guessed(tmp_path): path = write_adapter(tmp_path, {"blocks.0.attn.to_q.diff": torch.zeros(4)}) @@ -134,6 +274,13 @@ def test_apply_to_adds_the_delta(tmp_path): assert torch.allclose(out, torch.full((4, ), 1.25)) +def test_apply_to_supports_standalone_parameters(tmp_path): + path = write_adapter(tmp_path, {"blocks.0.scale_shift_table.diff_param": torch.full((4, ), 0.25)}) + patch = DenseLoRAPatch.from_adapter(path) + out = patch.apply_to("blocks.0.scale_shift_table", torch.ones(4)) + assert torch.allclose(out, torch.full((4, ), 1.25)) + + def test_apply_to_leaves_unrelated_parameters_untouched(tmp_path): path = write_adapter(tmp_path, {"blocks.0.norm1.diff": torch.full((4, ), 0.25)}) patch = DenseLoRAPatch.from_adapter(path) @@ -168,6 +315,13 @@ def test_apply_to_rejects_a_shape_mismatch(tmp_path): patch.apply_to("blocks.0.norm1.weight", torch.zeros(4)) +def test_replacement_for_supports_standalone_parameters(tmp_path): + table = torch.randn(6, 4) + path = write_adapter(tmp_path, {"blocks.0.extra_table.set_param": table}) + patch = DenseLoRAPatch.from_adapter(path) + assert torch.equal(patch.replacement_for("blocks.0.extra_table"), table) + + def test_replacement_for_returns_the_whole_tensor(tmp_path): gate = torch.randn(6, 4) path = write_adapter(tmp_path, {"blocks.0.attn.to_gate_compress.set_weight": gate}) diff --git a/fastvideo/tests/lora_extraction/test_lora_extraction.py b/fastvideo/tests/lora_extraction/test_lora_extraction.py index 51e35b4520..4ff39980a7 100644 --- a/fastvideo/tests/lora_extraction/test_lora_extraction.py +++ b/fastvideo/tests/lora_extraction/test_lora_extraction.py @@ -1,62 +1,139 @@ -"""Test LoRA extraction, merging, and verification pipeline.""" -import sys +"""Test extraction through the real FastVideo LoRA loading path.""" +import json from pathlib import Path +import sys +import tempfile + +import pytest +from safetensors import safe_open +import torch + +from fastvideo import VideoGenerator +from fastvideo.api import ComponentConfig, EngineConfig, GeneratorConfig, OffloadConfig, ParallelismConfig, PipelineSelection # Add scripts/lora_extraction to path for imports repo_root = Path(__file__).parents[3] lora_scripts = repo_root / "scripts" / "lora_extraction" sys.path.insert(0, str(lora_scripts)) -# Import the core functions -from extract_lora import extract_lora_adapter -from merge_lora import merge_lora -from verify_lora import main as verify_lora_main +from extract_lora import INDEX_FILENAME, _resolve_transformer_dir, extract_lora_adapter # noqa: E402 -def test_lora_extraction_pipeline(): - """Test end-to-end LoRA extraction workflow.""" - import tempfile +def _read_transformer_tensor(model: str, key: str, revision: str) -> torch.Tensor: + """Read one tensor without materializing or fingerprinting the full checkpoint.""" + transformer_dir = _resolve_transformer_dir(model, revision) + index_path = transformer_dir / INDEX_FILENAME + if index_path.is_file(): + index = json.loads(index_path.read_text(encoding="utf-8")) + shard_names = [index["weight_map"][key]] + else: + shard_names = [path.name for path in sorted(transformer_dir.glob("*.safetensors"))] + + for shard_name in shard_names: + with safe_open(transformer_dir / shard_name, framework="pt", device="cpu") as handle: + if key in handle.keys(): + return handle.get_tensor(key) + raise KeyError(f"{key} is absent from {transformer_dir}") - # Use temp directory for outputs to avoid polluting repo - with tempfile.TemporaryDirectory() as tmpdir: - tmpdir_path = Path(tmpdir) - adapter_path = tmpdir_path / "adapter_r16.safetensors" - merged_dir = tmpdir_path / "merged_r16" - # 1. Extract rank-16 adapter - print("\nExtracting rank-16 adapter") +def _collect_lora_application( + worker, + dense_param_name: str, + expected_dense_values: tuple[float, ...], +) -> dict[str, object]: + """Inspect the worker after the constructor applied its adapter.""" + pipeline = worker.pipeline + adapter = pipeline.lora_adapters[pipeline.cur_adapter_name] + available: set[str] = set() + adapted = 0 + for transformer_layers in pipeline.lora_layers.values(): + for _, layers in transformer_layers.lora_layers_by_block(): + for name, layer in layers.items(): + available.update((name + ".lora_A", name + ".lora_B", name + ".lora_alpha")) + if layer.lora_A is not None and layer.lora_B is not None and not layer.disable_lora: + adapted += 1 + unmatched = sorted(set(adapter) - available) + + transformer = pipeline.modules["transformer"] + dense_param = dict(transformer.named_parameters())[dense_param_name].detach().flatten() + expected_dense = torch.tensor(expected_dense_values, dtype=dense_param.dtype, device=dense_param.device) + dense_matches_finetuned = torch.equal(dense_param[:expected_dense.numel()], expected_dense) + return { + "adapted": adapted, + "dense_matches_finetuned": dense_matches_finetuned, + "pipeline": type(pipeline).__name__, + "unmatched": unmatched, + } + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="Wan2.2 integration requires a CUDA GPU") +def test_lora_extraction_pipeline() -> None: + """Extract Wan2.2 on a GPU and require every factor to reach the DMD pipeline.""" + base = "Wan-AI/Wan2.2-TI2V-5B-Diffusers" + base_revision = "b8fff7315c768468a5333511427288870b2e9635" + finetuned = "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers" + finetuned_revision = "3e187042a324f6f5fb68fd22110a78725253de8f" + dense_source_name = "condition_embedder.time_embedder.linear_1.bias" + dense_adapter_name = "condition_embedder.time_embedder.linear_1.diff_b" + dense_param_name = "condition_embedder.time_embedder.mlp.fc_in.bias" + with tempfile.TemporaryDirectory() as tmpdir: + adapter_path = Path(tmpdir) / "adapter_r16.safetensors" extract_lora_adapter( - base="Wan-AI/Wan2.2-TI2V-5B-Diffusers", - ft="FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers", + base=base, + ft=finetuned, out=str(adapter_path), rank=16, + base_revision=base_revision, + ft_revision=finetuned_revision, + load_mode="indexed", + device="cuda:0", + svd_method="exact", + exact_tensor_patterns=(r"^condition_embedder\.", r"^proj_out\.weight$"), ) - assert adapter_path.exists(), "Adapter file was not created" - - # 2. Merge adapter - print("\nMerging adapter") - merge_lora( - base="Wan-AI/Wan2.2-TI2V-5B-Diffusers", - adapter=str(adapter_path), - ft="FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers", - output=str(merged_dir), - ) - assert merged_dir.exists(), "Merged model directory was not created" - # 3. Verify numerical accuracy - print("\nVerifying merged model") - # verify_lora uses sys.argv, so we need to mock it - old_argv = sys.argv + with safe_open(adapter_path, framework="pt") as handle: + assert dense_adapter_name in handle.keys(), "precondition: the selected Wan parameter must be a dense delta" + expected_dense_values = tuple( + _read_transformer_tensor(finetuned, dense_source_name, finetuned_revision).flatten()[:16].tolist()) + + generator = VideoGenerator.from_config( + GeneratorConfig( + model_path=base, + revision=base_revision, + pipeline=PipelineSelection( + components=ComponentConfig( + lora_path=str(adapter_path), + override_pipeline_cls_name="WanDMDPipeline", + ), + experimental={ + "dmd_denoising_steps": [1000, 757, 522], + "flow_shift": 5.0, + }, + ), + engine=EngineConfig( + num_gpus=1, + use_fsdp_inference=False, + parallelism=ParallelismConfig(tp_size=1, sp_size=1), + offload=OffloadConfig( + dit=False, + dit_layerwise=False, + text_encoder=True, + vae=True, + pin_cpu_memory=False, + ), + ), + )) try: - sys.argv = [ - "verify_lora.py", - "--merged", - str(merged_dir), - "--ft", - "FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers", - ] - verify_lora_main() + summaries = generator.executor.collective_rpc( + _collect_lora_application, + args=(dense_param_name, expected_dense_values), + ) finally: - sys.argv = old_argv + generator.shutdown() - print("\nLoRA extraction pipeline test PASSED") + assert summaries == [{ + "adapted": 300, + "dense_matches_finetuned": True, + "pipeline": "WanDMDPipeline", + "unmatched": [], + }] diff --git a/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py b/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py new file mode 100644 index 0000000000..1a10cd09bc --- /dev/null +++ b/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py @@ -0,0 +1,596 @@ +"""Unit coverage for streaming and GPU-capable LoRA extraction.""" + +from __future__ import annotations + +import json +from pathlib import Path +import sys + +import pytest +from safetensors import safe_open +from safetensors.torch import load_file, save_file +import torch + +_REPO_ROOT = Path(__file__).parents[3] +sys.path.insert(0, str(_REPO_ROOT / "scripts" / "lora_extraction")) + +import extract_lora # noqa: E402 + + +def _write_transformer(root: Path, state: dict[str, torch.Tensor]) -> None: + transformer = root / "transformer" + transformer.mkdir(parents=True) + shard = "diffusion_pytorch_model-00001-of-00001.safetensors" + save_file(state, transformer / shard) + index = { + "metadata": { + "total_size": sum(tensor.numel() * tensor.element_size() for tensor in state.values()) + }, + "weight_map": {key: shard for key in state}, + } + (transformer / extract_lora.INDEX_FILENAME).write_text(json.dumps(index), encoding="utf-8") + + +def _toy_states() -> tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]: + generator = torch.Generator().manual_seed(7) + base = { + "blocks.0.linear.weight": torch.randn(9, 7, generator=generator), + "blocks.0.norm.weight": torch.randn(7, generator=generator), + "context.weight": torch.randn(8, 6, generator=generator), + "unchanged.bias": torch.randn(8, generator=generator), + } + finetuned = {key: value.clone() for key, value in base.items()} + finetuned["blocks.0.linear.weight"] += torch.randn(9, 2, generator=generator) @ torch.randn( + 2, 7, generator=generator) + finetuned["blocks.0.norm.weight"] += 0.125 + finetuned["context.weight"] += torch.randn(8, 6, generator=generator) * 0.01 + finetuned["blocks.0.attn.to_gate_compress.weight"] = torch.randn( + 9, 7, generator=generator).to(torch.bfloat16) + return base, finetuned + + +def _legacy_cpu_adapter( + base: dict[str, torch.Tensor], + finetuned: dict[str, torch.Tensor], + rank: int, + min_delta: float, +) -> dict[str, torch.Tensor]: + """Reproduce the pre-streaming extractor for CPU parity coverage.""" + adapter: dict[str, torch.Tensor] = {} + low_rank_keys: set[str] = set() + for key in sorted(finetuned): + if not extract_lora.is_extractable_weight(key) or key not in base: + continue + base_weight = base[key].detach().cpu().float().contiguous() + finetuned_weight = finetuned[key].detach().cpu().float().contiguous() + if base_weight.shape != finetuned_weight.shape: + continue + delta = (finetuned_weight - base_weight).contiguous() + if float(delta.abs().mean()) < min_delta: + continue + try: + u, singular_values, vh = torch.linalg.svd(delta, full_matrices=False) + except RuntimeError: + continue + chosen_rank = min(rank, singular_values.numel()) + sqrt_s = singular_values[:chosen_rank].float().sqrt() + lora_b = (u[:, :chosen_rank].float() * sqrt_s.unsqueeze(0)).contiguous() + lora_a = (vh[:chosen_rank].mT.float() * sqrt_s.unsqueeze(0)).mT.contiguous() + module_name = key.removesuffix(".weight") + adapter[f"{module_name}.lora_A.weight"] = lora_a + adapter[f"{module_name}.lora_B.weight"] = lora_b + low_rank_keys.add(key) + + adapter.update(extract_lora.build_dense_payload(base, finetuned, low_rank_keys, min_delta)) + return adapter + + +def _reconstructed_delta(adapter: dict[str, torch.Tensor], module_name: str) -> torch.Tensor: + return adapter[f"{module_name}.lora_B.weight"].float() @ adapter[f"{module_name}.lora_A.weight"].float() + + +def test_exact_cpu_extraction_matches_legacy_algorithm(tmp_path: Path) -> None: + base, finetuned = _toy_states() + base_dir = tmp_path / "base" + finetuned_dir = tmp_path / "finetuned" + output = tmp_path / "adapter.safetensors" + _write_transformer(base_dir, base) + _write_transformer(finetuned_dir, finetuned) + + extract_lora.extract_lora_adapter( + base=str(base_dir), + ft=str(finetuned_dir), + out=str(output), + rank=2, + min_delta=1e-8, + load_mode="indexed", + device="cpu", + svd_method="exact", + factor_dtype="float32", + dense_dtype="source", + replacement_dtype="source", + ) + + actual = load_file(output) + expected = _legacy_cpu_adapter(base, finetuned, rank=2, min_delta=1e-8) + assert actual.keys() == expected.keys() + for key in actual: + assert torch.equal(actual[key], expected[key]), key + + +def test_streaming_extraction_emits_lora_diff_and_replacement(tmp_path: Path) -> None: + base, finetuned = _toy_states() + base_dir = tmp_path / "base" + finetuned_dir = tmp_path / "finetuned" + output = tmp_path / "adapter.safetensors" + _write_transformer(base_dir, base) + _write_transformer(finetuned_dir, finetuned) + + result = extract_lora.extract_lora_adapter( + base=str(base_dir), + ft=str(finetuned_dir), + out=str(output), + rank=2, + min_delta=0.0, + load_mode="indexed", + device="cpu", + svd_method="exact", + factor_dtype="float16", + dense_dtype="float32", + replacement_dtype="source", + exact_tensor_patterns=(r"^context\.weight$", ), + ) + + assert result == output + adapter = load_file(output) + torch.testing.assert_close( + adapter["blocks.0.linear.lora_B.weight"].float() + @ adapter["blocks.0.linear.lora_A.weight"].float(), + finetuned["blocks.0.linear.weight"] - base["blocks.0.linear.weight"], + atol=6e-3, + rtol=6e-3, + ) + assert adapter["blocks.0.norm.diff"].dtype == torch.float32 + assert adapter["context.diff"].dtype == torch.float32 + assert adapter["blocks.0.attn.to_gate_compress.set_weight"].dtype == torch.bfloat16 + assert not any("unchanged" in key for key in adapter) + assert not any(key.endswith((".lora_rank", ".lora_alpha")) for key in adapter) + + report = json.loads(output.with_suffix(".safetensors.report.json").read_text()) + assert report["counts"] == {"diff": 2, "lora": 1, "set_weight": 1, "unchanged": 1} + with safe_open(output, framework="pt") as handle: + assert handle.metadata()["svd_method"] == "exact" + assert handle.metadata()["factor_dtype"] == "float16" + + +@pytest.mark.parametrize( + ("dtype", "base_value", "finetuned_value"), + [ + (torch.float16, -0.72412109375, 0.07098388671875), + (torch.bfloat16, 0.5078125, -0.0556640625), + ], +) +def test_dense_deltas_default_to_exact_float32( + tmp_path: Path, + dtype: torch.dtype, + base_value: float, + finetuned_value: float, +) -> None: + base = {"blocks.0.norm.weight": torch.tensor([base_value], dtype=dtype)} + finetuned = {"blocks.0.norm.weight": torch.tensor([finetuned_value], dtype=dtype)} + base_dir = tmp_path / "base" + finetuned_dir = tmp_path / "finetuned" + output = tmp_path / "adapter.safetensors" + _write_transformer(base_dir, base) + _write_transformer(finetuned_dir, finetuned) + + extract_lora.extract_lora_adapter( + base=str(base_dir), + ft=str(finetuned_dir), + out=str(output), + load_mode="indexed", + min_delta=0.0, + ) + + delta = load_file(output)["blocks.0.norm.diff"] + assert delta.dtype == torch.float32 + reconstructed = (base["blocks.0.norm.weight"].float() + delta).to(dtype) + assert torch.equal(reconstructed, finetuned["blocks.0.norm.weight"]) + + +def test_standalone_parameters_use_generic_dense_suffixes(tmp_path: Path) -> None: + base = {"blocks.0.scale_shift_table": torch.ones(4)} + finetuned = { + "blocks.0.scale_shift_table": torch.full((4, ), 1.25), + "blocks.0.extra_table": torch.arange(4, dtype=torch.float32), + } + base_dir = tmp_path / "base" + finetuned_dir = tmp_path / "finetuned" + output = tmp_path / "adapter.safetensors" + _write_transformer(base_dir, base) + _write_transformer(finetuned_dir, finetuned) + + extract_lora.extract_lora_adapter( + base=str(base_dir), + ft=str(finetuned_dir), + out=str(output), + rank=2, + min_delta=0.0, + load_mode="indexed", + dense_dtype="float32", + replacement_dtype="source", + ) + + adapter = load_file(output) + torch.testing.assert_close(adapter["blocks.0.scale_shift_table.diff_param"], torch.full((4, ), 0.25)) + torch.testing.assert_close(adapter["blocks.0.extra_table.set_param"], finetuned["blocks.0.extra_table"]) + report = json.loads(output.with_suffix(".safetensors.report.json").read_text()) + assert report["counts"] == {"diff": 1, "set_param": 1} + with safe_open(output, framework="pt") as handle: + assert handle.metadata()["set_param_tensors"] == "1" + + +def test_randomized_extraction_is_seeded_and_reports_residual(tmp_path: Path) -> None: + base, finetuned = _toy_states() + base_dir = tmp_path / "base" + finetuned_dir = tmp_path / "finetuned" + _write_transformer(base_dir, base) + _write_transformer(finetuned_dir, finetuned) + + outputs = [] + for index in range(2): + output = tmp_path / f"adapter-{index}.safetensors" + extract_lora.extract_lora_adapter( + base=str(base_dir), + ft=str(finetuned_dir), + out=str(output), + rank=2, + min_delta=0.0, + load_mode="indexed", + device="cpu", + svd_method="randomized", + randomized_q=4, + niter=2, + seed=123, + factor_dtype="float32", + dense_dtype="float32", + exact_tensor_patterns=(r"^context\.weight$", ), + ) + outputs.append(load_file(output)) + + assert outputs[0].keys() == outputs[1].keys() + for key in outputs[0]: + assert torch.equal(outputs[0][key], outputs[1][key]), key + report = json.loads((tmp_path / "adapter-0.safetensors.report.json").read_text()) + assert 0.0 <= report["factorized_weighted_relative_residual"] <= 1.0 + layer = report["layers"]["blocks.0.linear.weight"] + assert layer["method"] == "randomized-q4-niter2" + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is unavailable") +def test_exact_cpu_and_gpu_extraction_agree(tmp_path: Path) -> None: + base, finetuned = _toy_states() + base_dir = tmp_path / "base" + finetuned_dir = tmp_path / "finetuned" + _write_transformer(base_dir, base) + _write_transformer(finetuned_dir, finetuned) + + adapters = {} + reports = {} + for label, device in (("cpu", "cpu"), ("gpu", "cuda:0")): + output = tmp_path / f"adapter-{label}.safetensors" + extract_lora.extract_lora_adapter( + base=str(base_dir), + ft=str(finetuned_dir), + out=str(output), + rank=2, + min_delta=1e-8, + load_mode="indexed", + device=device, + svd_method="exact", + factor_dtype="float32", + dense_dtype="float32", + replacement_dtype="source", + ) + adapters[label] = load_file(output) + reports[label] = json.loads(output.with_suffix(".safetensors.report.json").read_text()) + + assert adapters["cpu"].keys() == adapters["gpu"].keys() + for key in adapters["cpu"]: + if ".lora_" not in key: + torch.testing.assert_close(adapters["cpu"][key], adapters["gpu"][key], atol=0, rtol=0) + + for module_name in ("blocks.0.linear", "context"): + cpu_delta = _reconstructed_delta(adapters["cpu"], module_name) + gpu_delta = _reconstructed_delta(adapters["gpu"], module_name) + torch.testing.assert_close(cpu_delta, gpu_delta, atol=2e-5, rtol=2e-4) + + cpu_residual = reports["cpu"]["factorized_weighted_relative_residual"] + gpu_residual = reports["gpu"]["factorized_weighted_relative_residual"] + # The aggregate residual is close to zero and therefore sensitive to + # backend-level singular-value rounding even when reconstructed deltas agree. + assert gpu_residual == pytest.approx(cpu_residual, abs=2e-4) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is unavailable") +def test_randomized_factorization_runs_on_gpu() -> None: + generator = torch.Generator().manual_seed(11) + delta = torch.randn(32, 24, generator=generator, dtype=torch.float32).cuda() + exact_a, exact_b, _, _ = extract_lora._factorize_delta( + delta, + rank=4, + full_rank=False, + method="exact", + randomized_q=None, + oversample=4, + niter=2, + seed=7, + ) + random_a, random_b, _, method = extract_lora._factorize_delta( + delta, + rank=4, + full_rank=False, + method="randomized", + randomized_q=12, + oversample=4, + niter=4, + seed=7, + ) + exact_error = torch.linalg.vector_norm(delta - exact_b @ exact_a) + randomized_error = torch.linalg.vector_norm(delta - random_b @ random_a) + assert randomized_error <= exact_error * 1.02 + assert method == "randomized-q12-niter4" + assert random_a.is_cuda and random_b.is_cuda + + +def test_pipeline_loading_disables_layerwise_offload(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[dict[str, object]] = [] + + class FakePipeline: + + @classmethod + def from_pretrained(cls, _model_path: str, **kwargs: object): + calls.append(kwargs) + pipeline = cls() + pipeline.pipeline = cls() + pipeline.pipeline.transformer = torch.nn.Linear(3, 2) + return pipeline + + monkeypatch.setattr(extract_lora, "get_pipeline_class_for_model", lambda _model_path: FakePipeline) + + state_dict = extract_lora.load_transformer_state_dict_from_model("org/model") + + assert calls[0]["dit_layerwise_offload"] is False + assert state_dict["weight"].shape == (2, 3) + + +def test_hub_resolution_downloads_only_transformer(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + snapshot = tmp_path / "snapshot" + (snapshot / "transformer").mkdir(parents=True) + calls: list[dict[str, object]] = [] + + def fake_snapshot_download(**kwargs: object) -> str: + calls.append(kwargs) + return str(snapshot) + + monkeypatch.setattr(extract_lora, "snapshot_download", fake_snapshot_download) + assert extract_lora._resolve_transformer_dir("org/model", "revision") == snapshot / "transformer" + assert calls == [{ + "repo_id": "org/model", + "revision": "revision", + "allow_patterns": ["transformer/*"], + }] + + +def _toy_checkpoints(tmp_path: Path) -> tuple[Path, Path]: + base, finetuned = _toy_states() + base_dir = tmp_path / "base" + finetuned_dir = tmp_path / "finetuned" + _write_transformer(base_dir, base) + _write_transformer(finetuned_dir, finetuned) + return base_dir, finetuned_dir + + +def _extract(base_dir: Path, finetuned_dir: Path, output: Path, **overrides: object) -> Path: + kwargs: dict[str, object] = { + "rank": 2, + "min_delta": 1e-8, + "load_mode": "indexed", + "device": "cpu", + "svd_method": "exact", + } + kwargs.update(overrides) + return extract_lora.extract_lora_adapter(base=str(base_dir), ft=str(finetuned_dir), out=str(output), **kwargs) + + +def _scratch_dir(work_root: Path, output: Path) -> Path: + return work_root / extract_lora.WORK_SUBDIR / extract_lora._work_namespace(output) + + +def test_work_dir_leaves_unrelated_files_alone(tmp_path: Path) -> None: + """--work-dir names where to put scratch, so its other contents must survive.""" + base_dir, finetuned_dir = _toy_checkpoints(tmp_path) + work_dir = tmp_path / "shared_scratch" + (work_dir / "nested").mkdir(parents=True) + keeper = work_dir / "unrelated.txt" + keeper.write_text("do not delete me", encoding="utf-8") + + _extract(base_dir, finetuned_dir, tmp_path / "adapter.safetensors", work_dir=str(work_dir)) + + assert keeper.read_text(encoding="utf-8") == "do not delete me" + assert (work_dir / "nested").is_dir() + assert not (work_dir / extract_lora.WORK_SUBDIR).exists() + + +def test_work_dir_scratch_is_confined_to_a_subdirectory(tmp_path: Path) -> None: + base_dir, finetuned_dir = _toy_checkpoints(tmp_path) + work_dir = tmp_path / "shared_scratch" + + _extract(base_dir, + finetuned_dir, + tmp_path / "adapter.safetensors", + work_dir=str(work_dir), + keep_work_dir=True) + + assert (_scratch_dir(work_dir, tmp_path / "adapter.safetensors") / "manifest.json").is_file() + + +def test_work_dir_namespaces_different_outputs(tmp_path: Path) -> None: + base_dir, finetuned_dir = _toy_checkpoints(tmp_path) + work_dir = tmp_path / "shared_scratch" + outputs = [tmp_path / "rank64.safetensors", tmp_path / "rank128.safetensors"] + + for output in outputs: + _extract(base_dir, finetuned_dir, output, work_dir=str(work_dir), keep_work_dir=True) + + scratch_dirs = {_scratch_dir(work_dir, output) for output in outputs} + assert len(scratch_dirs) == 2 + assert all((path / "manifest.json").is_file() for path in scratch_dirs) + + +def test_unmatched_exact_tensor_pattern_is_rejected(tmp_path: Path) -> None: + """A pattern that matches nothing silently factorizes the tensors it was meant to keep.""" + base_dir, finetuned_dir = _toy_checkpoints(tmp_path) + + with pytest.raises(ValueError, match="matched no tensor"): + _extract(base_dir, + finetuned_dir, + tmp_path / "adapter.safetensors", + exact_tensor_patterns=[r"^context\\.weight$"]) + + _extract(base_dir, + finetuned_dir, + tmp_path / "adapter.safetensors", + exact_tensor_patterns=[r"^context\.weight$"]) + assert "context.diff" in load_file(tmp_path / "adapter.safetensors") + + +def test_rerun_with_a_different_config_reuses_the_work_dir(tmp_path: Path) -> None: + """A leftover work dir must not make a fresh (non-resume) rerun fail.""" + base_dir, finetuned_dir = _toy_checkpoints(tmp_path) + output = tmp_path / "adapter.safetensors" + + _extract(base_dir, finetuned_dir, output, rank=2, keep_work_dir=True) + work_dir = output.parent / f".{output.name}.work" + assert (work_dir / "manifest.json").is_file() + + _extract(base_dir, finetuned_dir, output, rank=4) + assert output.is_file() + + +def test_resume_still_rejects_a_mismatched_config(tmp_path: Path) -> None: + base_dir, finetuned_dir = _toy_checkpoints(tmp_path) + output = tmp_path / "adapter.safetensors" + + _extract(base_dir, finetuned_dir, output, rank=2, keep_work_dir=True) + with pytest.raises(ValueError, match="Resume configuration does not match"): + _extract(base_dir, finetuned_dir, output, rank=4, resume=True) + + +@pytest.mark.parametrize("changed_side", ["base", "finetuned"]) +def test_resume_rejects_changed_checkpoint_contents(tmp_path: Path, changed_side: str) -> None: + base_dir, finetuned_dir = _toy_checkpoints(tmp_path) + output = tmp_path / "adapter.safetensors" + _extract(base_dir, finetuned_dir, output, keep_work_dir=True) + + base, finetuned = _toy_states() + changed = base if changed_side == "base" else finetuned + changed["blocks.0.linear.weight"] += 1.0 + changed_dir = base_dir if changed_side == "base" else finetuned_dir + shard = changed_dir / "transformer" / "diffusion_pytorch_model-00001-of-00001.safetensors" + save_file(changed, shard) + + with pytest.raises(ValueError, match="Resume configuration does not match"): + _extract(base_dir, finetuned_dir, output, resume=True) + + +@pytest.mark.parametrize("revision_kw", ["base_revision", "ft_revision"]) +def test_local_paths_reject_hugging_face_revisions(tmp_path: Path, revision_kw: str) -> None: + base_dir, finetuned_dir = _toy_checkpoints(tmp_path) + with pytest.raises(ValueError, match="cannot be used with local model path"): + _extract( + base_dir, + finetuned_dir, + tmp_path / "adapter.safetensors", + **{revision_kw: "abc123"}, + ) + + +def test_pipeline_mode_rejects_resume_before_loading(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + def forbidden_pipeline_load(*args, **kwargs): + raise AssertionError("resume must fail before materializing pipeline checkpoints") + + monkeypatch.setattr(extract_lora, "load_transformer_state_dict_from_model", forbidden_pipeline_load) + with pytest.raises(ValueError, match="resume requires indexed safetensors"): + extract_lora.extract_lora_adapter( + base="org/base", + ft="org/finetuned", + out=str(tmp_path / "adapter.safetensors"), + load_mode="pipeline", + resume=True, + ) + + +@pytest.mark.parametrize("revision_kw", ["base_revision", "ft_revision"]) +def test_pipeline_mode_rejects_revisions_before_loading(tmp_path: Path, revision_kw: str) -> None: + with pytest.raises(ValueError, match="require indexed loading"): + extract_lora.extract_lora_adapter( + base="org/base", + ft="org/finetuned", + out=str(tmp_path / "adapter.safetensors"), + load_mode="pipeline", + **{revision_kw: "abc123"}, + ) + + +def test_auto_resume_does_not_fall_back_to_pipeline(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + def indexed_failure(*args, **kwargs): + raise RuntimeError("indexed unavailable") + + def forbidden_pipeline_load(*args, **kwargs): + raise AssertionError("pipeline fallback cannot validate checkpoint identity") + + monkeypatch.setattr(extract_lora, "_resolve_transformer_dir", indexed_failure) + monkeypatch.setattr(extract_lora, "load_transformer_state_dict_from_model", forbidden_pipeline_load) + with pytest.raises(RuntimeError, match="indexed unavailable"): + extract_lora.extract_lora_adapter( + base="org/base", + ft="org/finetuned", + out=str(tmp_path / "adapter.safetensors"), + load_mode="auto", + resume=True, + ) + + +@pytest.mark.parametrize("revision_kw", ["base_revision", "ft_revision"]) +def test_auto_mode_does_not_drop_revision_during_fallback(monkeypatch: pytest.MonkeyPatch, tmp_path: Path, + revision_kw: str) -> None: + def indexed_failure(*args, **kwargs): + raise RuntimeError("indexed unavailable") + + def forbidden_pipeline_load(*args, **kwargs): + raise AssertionError("pipeline fallback would ignore the requested revision") + + monkeypatch.setattr(extract_lora, "_resolve_transformer_dir", indexed_failure) + monkeypatch.setattr(extract_lora, "load_transformer_state_dict_from_model", forbidden_pipeline_load) + with pytest.raises(RuntimeError, match="indexed unavailable"): + extract_lora.extract_lora_adapter( + base="org/base", + ft="org/finetuned", + out=str(tmp_path / "adapter.safetensors"), + load_mode="auto", + **{revision_kw: "abc123"}, + ) + + +def test_cli_defaults_exact_dense_deltas_to_float32(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(sys, "argv", ["extract_lora.py", "--base", "base", "--ft", "finetuned"]) + assert extract_lora.parse_args().dense_dtype == "float32" + + +def test_non_safetensors_output_is_rejected(tmp_path: Path) -> None: + """The writer is always safetensors, so a .pt name would be a mislabeled file.""" + base_dir, finetuned_dir = _toy_checkpoints(tmp_path) + + with pytest.raises(ValueError, match="must end in .safetensors"): + _extract(base_dir, finetuned_dir, tmp_path / "adapter.pt") diff --git a/scripts/lora_extraction/README.md b/scripts/lora_extraction/README.md index 01d82f3f1f..f4ab15353e 100644 --- a/scripts/lora_extraction/README.md +++ b/scripts/lora_extraction/README.md @@ -1,40 +1,104 @@ # LoRA Extraction and Merging -Tools for extracting and merging LoRA adapters for FastVideo models. +Generic runtime adapter extraction plus the existing legacy merge utilities for FastVideo models. ## Extract LoRA Adapter +The default remains exact CPU SVD: + ```bash python extract_lora.py \ --base Wan-AI/Wan2.2-TI2V-5B-Diffusers \ --ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \ --out adapter_r32.safetensors \ - --rank 32 + --rank 32 \ + --exact-tensor-pattern '^condition_embedder\.' \ + --exact-tensor-pattern '^proj_out\.weight$' ``` -**Options:** -- `--base`: Base model (HuggingFace ID or local path) -- `--ft`: Fine-tuned model (HuggingFace ID or local path) -- `--out`: Output adapter file -- `--rank`: LoRA rank (16, 32, 64, 128) -- `--full-rank`: Extract full-rank adapter (optional) +The extractor is runtime-agnostic: it cannot infer from checkpoint tensors whether a runtime wraps a +particular matrix as a LoRA layer. When extracting a full fine-tune, select matrices unsupported by the target runtime +with `--exact-tensor-pattern`; their changes remain exact `.diff` tensors instead of becoming factors the runtime +cannot apply. The Wan patterns above cover its excluded condition embedders and its unwrapped output projection. + +For large transformers, stream their indexed safetensors and factorize on a GPU: + +```bash +python extract_lora.py \ + --base MiniMaxAI/MiniMax-H3 \ + --base-revision 9bfb6693f2cf6de171db46d1aa586f67d773a1da \ + --ft FastVideo/FastVideo-FastH3-4-step-v1.1 \ + --ft-revision c9e910404950b42f627f07b0c4a09d9a3e087d47 \ + --out adapter_r64.safetensors \ + --rank 64 \ + --load-mode indexed \ + --device cuda:0 \ + --svd-method randomized \ + --randomized-q 320 \ + --niter 4 \ + --factor-dtype float16 \ + --dense-dtype float32 \ + --replacement-dtype source \ + --exact-tensor-pattern '^audio_proj_(in|out)\.weight$' \ + --exact-tensor-pattern '^context_embedder\.weight$' \ + --exact-tensor-pattern '^proj_(in|out)\.weight$' \ + --exact-tensor-pattern '^time_embedder\.' +``` + +`q=320, niter=4` retained 99.9355% of the energy captured by exact rank-64 SVD in a 362-matrix MiniMax-H3 comparison. Exact CPU SVD is still the default; randomized SVD must be requested explicitly. + +This FastH3 checkpoint contains VSA compression-gate replacements. Load the extracted adapter with +[`basic_fasth3_lora_preview.py`](../../examples/inference/basic/basic_fasth3_lora_preview.py), which inspects the payload +and selects the required MiniMax-H3 VSA backend. + +Important options: + +- `--base`, `--ft`: Hugging Face model IDs or local paths. +- `--rank`, `--full-rank`: truncated or full factorization rank. +- `--min-delta`: omit tensors whose maximum absolute FP32 delta is at or below this threshold (default: `1e-8`). +- `--load-mode indexed`: download/read only `transformer/*` and stream one tensor pair at a time. +- `--device`: factorization device, such as `cpu` or `cuda:0`. +- `--svd-method`: `exact` or `randomized`. +- `--randomized-q`, `--niter`, `--seed`: randomized SVD accuracy and reproducibility. +- `--factor-dtype`: storage dtype for `lora_A` and `lora_B`. +- `--dense-dtype`: storage dtype for exact `.diff`, `.diff_b`, and `.diff_param` payloads (default: `float32`). +- `--replacement-dtype`: storage dtype for fine-tuned-only `.set_weight` and `.set_param` parameters. +- `--exact-tensor-pattern`: repeatable regex selecting matrices to retain as exact dense deltas. +- `--base-revision`, `--ft-revision`: pin Hugging Face inputs in indexed mode. Revisions are rejected for local paths and pipeline loading rather than silently ignored. +- `--work-dir`, `--resume`: resume a partially completed streaming extraction. Scratch is written to an + output-specific directory under `fastvideo-lora-extract/`, and only that namespace is cleaned up. Resume requires indexed + safetensors and validates both checkpoints' index/shard fingerprints before reusing partial results. + +Fine-tuned parameters that cannot or should not be factorized are retained automatically: + +- a changed base weight becomes `.diff`; +- a changed base bias becomes `.diff_b`; +- another changed parameter, such as `scale_shift_table`, becomes `.diff_param`; +- a fine-tuned-only weight, such as a VSA compression gate, becomes `.set_weight`; +- another fine-tuned-only parameter becomes `.set_param`; +- a bit-identical parameter is omitted. +Indexed loading is preferred and downloads only the transformer component. `--load-mode auto` falls back to legacy pipeline loading when indexed safetensors are unavailable. -> **Note:** The script automatically handles architectural differences (e.g., FastWan has extra `gate_compress` layers) by falling back to direct safetensors loading for both models if pipeline loading fails. +Mixed low-rank/dense adapters produced by this extractor must be supplied when constructing FastVideo through +`ComponentConfig(lora_path=...)`; their dense payload cannot be swapped later with `set_lora_adapter`. The legacy +offline merger below retains its existing scope and is not part of the generic extraction workflow. +## Legacy Merge Adapter -## Merge Adapter +The command below documents the pre-existing merger for adapters it already supports. Do not pass a mixed adapter from +the generic extractor to it: the legacy merger does not apply the adapter's exact dense or replacement payloads. ```bash python merge_lora.py \ --base Wan-AI/Wan2.2-TI2V-5B-Diffusers \ - --adapter adapter_r32.safetensors \ + --adapter legacy_factor_only_adapter.safetensors \ --ft FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers \ --output merged_model ``` **Options:** -- `--base`: Base model (HuggingFace ID or local path) +- `--base`: Base model (Hugging Face ID or local path) - `--adapter`: LoRA adapter file (.safetensors) - `--ft`: Fine-tuned model (for configuration) - `--output`: Output directory diff --git a/scripts/lora_extraction/extract_lora.py b/scripts/lora_extraction/extract_lora.py index d506bf43bc..b6da85e517 100644 --- a/scripts/lora_extraction/extract_lora.py +++ b/scripts/lora_extraction/extract_lora.py @@ -1,114 +1,293 @@ -"""Extract FastVideo-style LoRA adapters from a fine-tuned model by SVDing (FT - base). - -Usage: - python scripts/lora_extraction/extract_lora.py \\ - --base --ft --out adapter.safetensors --rank 16 - -Example for models with architectural differences (fallback is automatic): - python extract_lora.py \\ - --base Wan-AI/Wan2.1-T2V-1.3B-Diffusers \\ - --ft FastVideo/FastWan2.1-T2V-1.3B-Diffusers \\ - --out fastvideo_adapter.safetensors \\ - --rank 32 +"""Extract a FastVideo LoRA adapter from the difference between two checkpoints. + +The extractor supports ordinary low-rank matrix deltas as well as parameters +that cannot be represented by a LoRA product: + +* ``.diff`` / ``.diff_b`` / ``.diff_param`` store exact additive deltas. +* ``.set_weight`` / ``.set_param`` store parameters absent from the base checkpoint. + +Indexed safetensors are streamed one tensor at a time, so extracting from large +transformers does not require both state dictionaries in host memory. Exact CPU +SVD remains the default for compatibility; GPU and randomized SVD are opt-in. """ from __future__ import annotations +import argparse +from collections.abc import Iterable, Iterator, Sequence +from contextlib import ExitStack, contextmanager, suppress +from dataclasses import asdict, dataclass +import hashlib +import json +import logging import os -# Set distributed env BEFORE any fastvideo imports +from pathlib import Path +import re +import shutil +from typing import Any, Protocol + +# Pipeline loading imports distributed code. Keep the single-process defaults in +# place for the legacy --load-mode pipeline fallback. os.environ.setdefault("MASTER_ADDR", "127.0.0.1") os.environ.setdefault("MASTER_PORT", "29500") os.environ.setdefault("WORLD_SIZE", "1") os.environ.setdefault("RANK", "0") os.environ.setdefault("LOCAL_RANK", "0") -import argparse -import sys -import logging -from pathlib import Path -from typing import Dict, Optional - import torch +from huggingface_hub import snapshot_download +from safetensors import safe_open +from safetensors.torch import load_file, save_file from tqdm import tqdm -# Optional safetensors support -_HAVE_SAFETENSORS = True -try: - from safetensors.torch import save_file as safetensors_save # type: ignore - from safetensors import safe_open # type: ignore -except Exception: - _HAVE_SAFETENSORS = False - safe_open = None # type: ignore - - -def load_transformer_state_dict_from_safetensors(model_path: str) -> Dict[str, torch.Tensor]: - """Load transformer weights directly from safetensors files. - - This bypasses the pipeline loader and works even when the model has - architectural differences (e.g., extra layers in fine-tuned model). - - Args: - model_path: HuggingFace model ID or local path - - Returns: - State dict with all transformer weights - """ - from huggingface_hub import snapshot_download - import os +LOG = logging.getLogger("extract_lora") +INDEX_FILENAME = "diffusion_pytorch_model.safetensors.index.json" +FORMAT_VERSION = "fastvideo-lora-v2" +# Scratch lives in a subdirectory we create so cleanup never reaches a caller's files +# when --work-dir points at a directory that already holds something else. +WORK_SUBDIR = "fastvideo-lora-extract" - # Download or locate the model - if os.path.isdir(model_path): - local_path = model_path - else: - local_path = snapshot_download(model_path) +DIFF_SUFFIX = ".diff" +DIFF_BIAS_SUFFIX = ".diff_b" +DIFF_PARAM_SUFFIX = ".diff_param" +SET_WEIGHT_SUFFIX = ".set_weight" +SET_PARAM_SUFFIX = ".set_param" - # Find transformer directory - transformer_dir = os.path.join(local_path, "transformer") - if not os.path.isdir(transformer_dir): - raise FileNotFoundError(f"Transformer directory not found at {transformer_dir}") +_DTYPE_MAP = { + "float32": torch.float32, + "float": torch.float32, + "fp32": torch.float32, + "float16": torch.float16, + "half": torch.float16, + "fp16": torch.float16, + "bfloat16": torch.bfloat16, + "bf16": torch.bfloat16, +} - # Load all safetensors files - state_dict: Dict[str, torch.Tensor] = {} - for fname in sorted(os.listdir(transformer_dir)): - if fname.endswith('.safetensors'): - fpath = os.path.join(transformer_dir, fname) - with safe_open(fpath, framework='pt', device='cpu') as f: - for key in f.keys(): - state_dict[key] = f.get_tensor(key) - if not state_dict: - raise ValueError(f"No safetensors files found in {transformer_dir}") +class TensorReader(Protocol): + """Random access to one checkpoint's transformer tensors.""" - LOG.info("Loaded %d keys directly from safetensors", len(state_dict)) - return state_dict + source: str + fingerprint: str + @property + def keys(self) -> set[str]: ... -# Configure minimal logging -LOG = logging.getLogger("extract_lora") + def get_tensor(self, key: str) -> torch.Tensor: ... + + def get_shape(self, key: str) -> tuple[int, ...]: ... + + def assert_unchanged(self) -> None: ... + + def __enter__(self) -> "TensorReader": ... + + def __exit__(self, *args: object) -> None: ... + + +class DictTensorReader: + """Reader wrapper for the legacy pipeline-loading path.""" + + def __init__(self, state_dict: dict[str, torch.Tensor], source: str) -> None: + self.state_dict = state_dict + self.source = source + self.fingerprint = f"pipeline:{source}" + + @property + def keys(self) -> set[str]: + return set(self.state_dict) + + def get_tensor(self, key: str) -> torch.Tensor: + return self.state_dict[key] + + def get_shape(self, key: str) -> tuple[int, ...]: + return tuple(self.state_dict[key].shape) + + def assert_unchanged(self) -> None: + return None + + def __enter__(self) -> "DictTensorReader": + return self + + def __exit__(self, *args: object) -> None: + return None + + +class IndexedSafetensorsReader: + """Stream tensors from an indexed or unsharded transformer component.""" + + def __init__(self, transformer_dir: Path) -> None: + self.transformer_dir = transformer_dir + self.source = str(transformer_dir.resolve()) + index_path = transformer_dir / INDEX_FILENAME + if index_path.is_file(): + index = json.loads(index_path.read_text(encoding="utf-8")) + self.weight_map: dict[str, str] = index["weight_map"] + else: + self.weight_map = {} + for path in sorted(transformer_dir.glob("*.safetensors")): + with safe_open(path, framework="pt") as handle: + for key in handle.keys(): + if key in self.weight_map: + raise ValueError(f"Tensor {key} occurs in multiple shards under {transformer_dir}") + self.weight_map[key] = path.name + if not self.weight_map: + raise ValueError(f"No transformer safetensors found under {transformer_dir}") + + shard_names = sorted(set(self.weight_map.values())) + identity_files = ([index_path] if index_path.is_file() else []) + [ + transformer_dir / shard for shard in shard_names + ] + before = _file_stats(identity_files) + self.fingerprint = _fingerprint_files(identity_files, transformer_dir) + after = _file_stats(identity_files) + if before != after: + raise RuntimeError(f"Checkpoint files changed while fingerprinting {transformer_dir}") + self._identity_files = identity_files + self._file_stats = after + + stack = ExitStack() + try: + shards = { + shard: stack.enter_context(safe_open(transformer_dir / shard, framework="pt", device="cpu")) + for shard in shard_names + } + except Exception: + stack.close() + raise + self._stack = stack + self._shards = shards + + @property + def keys(self) -> set[str]: + return set(self.weight_map) + + def get_tensor(self, key: str) -> torch.Tensor: + return self._shards[self.weight_map[key]].get_tensor(key) + + def get_shape(self, key: str) -> tuple[int, ...]: + return tuple(self._shards[self.weight_map[key]].get_slice(key).get_shape()) + + def assert_unchanged(self) -> None: + if _file_stats(self._identity_files) != self._file_stats: + raise RuntimeError(f"Checkpoint files changed during extraction: {self.transformer_dir}") + + def __enter__(self) -> "IndexedSafetensorsReader": + return self + + def __exit__(self, *args: object) -> None: + self._stack.close() + + +@dataclass(frozen=True) +class ExtractionConfig: + base_source: str + finetuned_source: str + base_fingerprint: str + finetuned_fingerprint: str + rank: int + full_rank: bool + min_delta: float + device: str + svd_method: str + randomized_q: int | None + oversample: int + niter: int + seed: int + factor_dtype: str + dense_dtype: str + replacement_dtype: str + dense_payload: bool + exact_tensor_patterns: tuple[str, ...] def configure_logging(level: str = "INFO") -> None: + if LOG.handlers: + LOG.setLevel(level) + return handler = logging.StreamHandler() - fmt = "%(asctime)s %(levelname)s %(message)s" - handler.setFormatter(logging.Formatter(fmt, datefmt="%Y-%m-%d %H:%M:%S")) + handler.setFormatter(logging.Formatter("%(asctime)s %(levelname)s %(message)s", datefmt="%Y-%m-%d %H:%M:%S")) LOG.addHandler(handler) LOG.setLevel(level) +def _atomic_json_dump(data: Any, path: Path) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_suffix(path.suffix + ".tmp") + temporary.write_text(json.dumps(data, indent=2, sort_keys=True) + "\n", encoding="utf-8") + temporary.replace(path) + + +def _file_stats(paths: Sequence[Path]) -> dict[str, tuple[int, int, int]]: + return { + str(path.resolve()): (path.stat().st_size, path.stat().st_mtime_ns, path.stat().st_ino) + for path in paths + } + + +def _fingerprint_files(paths: Sequence[Path], root: Path) -> str: + """Strong checkpoint identity used to reject stale resume shards.""" + digest = hashlib.sha256() + for path in sorted(paths, key=lambda item: str(item)): + digest.update(str(path.relative_to(root)).encode("utf-8")) + with path.open("rb") as handle: + while chunk := handle.read(8 * 1024 * 1024): + digest.update(chunk) + return digest.hexdigest() + + +def _torch_dtype(name: str) -> torch.dtype: + try: + return _DTYPE_MAP[name.lower()] + except KeyError as error: + raise ValueError(f"Unsupported dtype {name!r}; choose from {sorted(_DTYPE_MAP)}") from error + + +def _resolve_output_dtype(name: str, source_dtype: torch.dtype) -> torch.dtype: + return source_dtype if name.lower() == "source" else _torch_dtype(name) + + +def _resolve_transformer_dir(model: str, revision: str | None = None) -> Path: + """Resolve a local/HF model to its transformer directory. + + For a Hub model, download only the transformer component. This is material + for compound models: MiniMax-H3 is roughly 464 GiB in total while one + transformer is about 65 GiB. + """ + path = Path(model).expanduser() + if path.exists(): + if revision is not None: + raise ValueError(f"revision={revision!r} cannot be used with local model path {path}") + if (path / "transformer").is_dir(): + return path / "transformer" + if (path / INDEX_FILENAME).is_file() or any(path.glob("*.safetensors")): + return path + raise FileNotFoundError(f"No transformer safetensors found under {path}") + + snapshot = Path( + snapshot_download( + repo_id=model, + revision=revision, + allow_patterns=["transformer/*"], + )) + transformer_dir = snapshot / "transformer" + if not transformer_dir.is_dir(): + raise FileNotFoundError(f"Downloaded repository {model} has no transformer directory") + return transformer_dir + + def get_pipeline_class_for_model(model_path: str): - """Return appropriate FastVideo Pipeline class for the model.""" - from fastvideo.utils import maybe_download_model_index # local import - from fastvideo.pipelines.pipeline_registry import get_pipeline_registry, PipelineType + """Return the FastVideo pipeline class for the legacy loading mode.""" from fastvideo.fastvideo_args import WorkloadType + from fastvideo.pipelines.pipeline_registry import PipelineType, get_pipeline_registry + from fastvideo.utils import maybe_download_model_index config = maybe_download_model_index(model_path) pipeline_name = config.get("_class_name") if pipeline_name is None: - raise ValueError(f"Model config for {model_path} missing _class_name (diffusers format expected).") - - pipeline_registry = get_pipeline_registry(PipelineType.BASIC) - pipeline_cls = pipeline_registry.resolve_pipeline_cls(pipeline_name, PipelineType.BASIC, WorkloadType.T2V) - return pipeline_cls + raise ValueError(f"Model config for {model_path} is missing _class_name") + registry = get_pipeline_registry(PipelineType.BASIC) + return registry.resolve_pipeline_cls(pipeline_name, PipelineType.BASIC, WorkloadType.T2V) def load_transformer_state_dict_from_model( @@ -118,304 +297,467 @@ def load_transformer_state_dict_from_model( vae_cpu_offload: bool = True, text_encoder_cpu_offload: bool = True, pin_cpu_memory: bool = True, -) -> Dict[str, torch.Tensor]: - """Load pipeline and extract transformer.state_dict as CPU tensors.""" +) -> dict[str, torch.Tensor]: + """Load a transformer through FastVideo. Prefer indexed loading for extraction.""" pipeline_cls = get_pipeline_class_for_model(model_path) pipeline = pipeline_cls.from_pretrained( model_path, num_gpus=num_gpus, inference_mode=True, dit_cpu_offload=dit_cpu_offload, + # This helper must materialize real parameters. Layerwise offload may expose + # placeholder tensors through state_dict() before a layer is activated. + dit_layerwise_offload=False, vae_cpu_offload=vae_cpu_offload, text_encoder_cpu_offload=text_encoder_cpu_offload, pin_cpu_memory=pin_cpu_memory, ) - - # Try to locate transformer in several typical attributes transformer = getattr(pipeline, "transformer", None) if transformer is None: modules = getattr(pipeline, "modules", None) - if isinstance(modules, dict): - transformer = modules.get("transformer") + transformer = modules.get("transformer") if isinstance(modules, dict) else None if transformer is None: - pipeline_attr = getattr(pipeline, "pipeline", None) - transformer = getattr(pipeline_attr, "transformer", None) if pipeline_attr else None + nested_pipeline = getattr(pipeline, "pipeline", None) + transformer = getattr(nested_pipeline, "transformer", None) if transformer is None: - raise RuntimeError( - "Transformer not found in pipeline. Expected pipeline.transformer or pipeline.modules['transformer'].") - - state_dict = transformer.state_dict() + raise RuntimeError("Transformer not found in pipeline") - # DTensor safe handling try: - from torch.distributed.tensor import DTensor # type: ignore - _HAS_DTENSOR = True - except Exception: - DTensor = None # type: ignore - _HAS_DTENSOR = False - - state_dict_cpu: Dict[str, torch.Tensor] = {} - for k, v in state_dict.items(): - if _HAS_DTENSOR and isinstance(v, DTensor): # type: ignore - state_dict_cpu[k] = v.to_local().detach().cpu().contiguous() + from torch.distributed.tensor import DTensor + except ImportError: + DTensor = None # type: ignore[assignment,misc] + + result: dict[str, torch.Tensor] = {} + for key, value in transformer.state_dict().items(): + if DTensor is not None and isinstance(value, DTensor): + value = value.to_local() + result[key] = value.detach().cpu().contiguous() + del pipeline, transformer + torch.cuda.empty_cache() + return result + + +def load_transformer_state_dict_from_safetensors(model_path: str) -> dict[str, torch.Tensor]: + """Compatibility helper that materializes a directly loaded state dictionary.""" + with IndexedSafetensorsReader(_resolve_transformer_dir(model_path)) as reader: + return {key: reader.get_tensor(key) for key in sorted(reader.keys)} + + +@contextmanager +def _open_readers( + base: str, + finetuned: str, + base_revision: str | None, + finetuned_revision: str | None, + load_mode: str, +) -> Iterator[tuple[TensorReader, TensorReader]]: + revisions_requested = base_revision is not None or finetuned_revision is not None + if load_mode == "pipeline" and revisions_requested: + raise ValueError("--base-revision/--ft-revision require indexed loading; pipeline loading cannot honor them") + if load_mode in {"auto", "indexed"}: + stack = ExitStack() + try: + base_reader = stack.enter_context(IndexedSafetensorsReader(_resolve_transformer_dir(base, base_revision))) + finetuned_reader = stack.enter_context( + IndexedSafetensorsReader(_resolve_transformer_dir(finetuned, finetuned_revision))) + except Exception: + stack.close() + if load_mode == "indexed" or revisions_requested: + raise + LOG.warning("Indexed loading failed; falling back to pipeline loading", exc_info=True) else: - state_dict_cpu[k] = v.detach().cpu().contiguous() + LOG.info("Streaming indexed transformers: base=%s finetuned=%s", base_reader.source, + finetuned_reader.source) + try: + yield base_reader, finetuned_reader + finally: + stack.close() + return - # cleanup - try: - del pipeline, transformer - except Exception: - pass - torch.cuda.empty_cache() - return state_dict_cpu + base_state = load_transformer_state_dict_from_model(base) + finetuned_state = load_transformer_state_dict_from_model(finetuned) + yield DictTensorReader(base_state, base), DictTensorReader(finetuned_state, finetuned) def is_extractable_weight(key: str) -> bool: - """Return True if key represents a weight suitable for LoRA extraction.""" + """Backward-compatible name filter for matrices suitable for LoRA.""" if not key.endswith("weight"): return False - low = key.lower() - for skip in ("norm", "bias", "embedding"): - if skip in low: - return False - return True - + lowered = key.lower() + return not any(fragment in lowered for fragment in ("norm", "bias", "embedding")) -# Suffixes read by fastvideo.models.loader.lora_patch, and by ComfyUI's loader, for -# payload that is not a low-rank product. Keep these spellings in sync with that module. -DIFF_SUFFIX = ".diff" -DIFF_BIAS_SUFFIX = ".diff_b" -SET_WEIGHT_SUFFIX = ".set_weight" - -def dense_payload_key(param_name: str) -> Optional[str]: - """The adapter key that carries ``param_name`` whole, or None if we cannot name one.""" +def dense_payload_key(param_name: str) -> str: if param_name.endswith(".weight"): - return param_name[:-len(".weight")] + DIFF_SUFFIX + return param_name.removesuffix(".weight") + DIFF_SUFFIX if param_name.endswith(".bias"): - return param_name[:-len(".bias")] + DIFF_BIAS_SUFFIX - return None + return param_name.removesuffix(".bias") + DIFF_BIAS_SUFFIX + return param_name + DIFF_PARAM_SUFFIX def build_dense_payload( - base_sd: Dict[str, torch.Tensor], - ft_sd: Dict[str, torch.Tensor], - low_rank_keys: set, + base_sd: dict[str, torch.Tensor], + ft_sd: dict[str, torch.Tensor], + low_rank_keys: set[str], min_delta: float, -) -> Dict[str, torch.Tensor]: - """Capture the parameters low-rank extraction cannot represent. - - ``is_extractable_weight`` rejects norms, biases, and embeddings, and the SVD loop - additionally skips anything the base model does not have. Those exclusions are - correct -- a rank-``r`` factorization of a length-``n`` vector costs ``r(1 + n) > n``, - and a parameter with no base weight has no delta to factor -- but dropping the - tensors outright loses whatever the fine-tune did to them. A distillation that - retunes its norms, or a VSA student whose compression gate exists only after - distillation, comes out measurably wrong. - - Two rules: - - * present in the base and changed -> ``.diff`` / ``.diff_b``, an exact delta - * absent from the base -> ``.set_weight``, the parameter itself - - Anything bit-identical to the base is skipped. That is not an optimization: shipping - a delta of exactly zero states "this changed" in a file whose whole purpose is to - record what changed, and on real extractions it is the majority of the candidates. - """ - payload: Dict[str, torch.Tensor] = {} - identical = 0 - skipped: list = [] - - for key in sorted(ft_sd.keys()): +) -> dict[str, torch.Tensor]: + """Compatibility helper for callers using in-memory state dictionaries.""" + payload: dict[str, torch.Tensor] = {} + for key in sorted(ft_sd): if key in low_rank_keys: continue - - ft_tensor = ft_sd[key].detach().cpu() - base_tensor = base_sd.get(key) - - if base_tensor is None: - # No base weight exists, so no delta is expressible: ship the parameter. - # `.set_weight` is only defined for a module's weight; a bias with no base - # counterpart has no agreed spelling, so say so rather than invent one. - if not key.endswith(".weight"): - skipped.append(key) - continue - payload[key[:-len(".weight")] + SET_WEIGHT_SUFFIX] = ft_tensor.contiguous() + finetuned = ft_sd[key].detach().cpu() + base = base_sd.get(key) + if base is None: + output_key = (key.removesuffix(".weight") + SET_WEIGHT_SUFFIX + if key.endswith(".weight") else key + SET_PARAM_SUFFIX) + payload[output_key] = finetuned.contiguous() continue - - base_tensor = base_tensor.detach().cpu() - if base_tensor.shape != ft_tensor.shape: - skipped.append(key) + if base.shape != finetuned.shape or torch.equal(base.cpu(), finetuned): continue - if torch.equal(base_tensor, ft_tensor): - identical += 1 + delta = finetuned.float() - base.cpu().float() + if float(delta.abs().max()) <= min_delta: continue + output_key = dense_payload_key(key) + if output_key is not None: + payload[output_key] = delta.to(finetuned.dtype).contiguous() + return payload - delta = (ft_tensor.to(torch.float32) - base_tensor.to(torch.float32)) - if float(delta.abs().max()) < min_delta: - identical += 1 - continue - out_key = dense_payload_key(key) - if out_key is None: - skipped.append(key) - continue - payload[out_key] = delta.to(ft_tensor.dtype).contiguous() +def _seed_for_key(seed: int, key: str) -> int: + digest = hashlib.sha256(f"{seed}:{key}".encode()).digest() + return int.from_bytes(digest[:8], "little") % (2**31) - LOG.info( - "Dense payload: %d tensors emitted, %d unchanged and dropped, %d skipped (unnameable or reshaped)", - len(payload), identical, len(skipped)) - for key in skipped[:10]: - LOG.warning("Dense payload skipped %s", key) - return payload + +def _compile_patterns(patterns: Sequence[str]) -> tuple[re.Pattern[str], ...]: + return tuple(re.compile(pattern) for pattern in patterns) -def save_adapter_state(adapter_state: Dict[str, torch.Tensor], - out_path: Path, - metadata: Optional[Dict[str, str]] = None) -> None: - """Save adapter state dict to safetensors (if available) or torch.save. +def _validate_exact_patterns(patterns: Sequence[str], keys: Iterable[str]) -> None: + """Reject a pattern that matches no tensor rather than silently factorizing it anyway. - Provenance goes in the safetensors header rather than a sidecar file so that it - survives being copied, renamed, or downloaded on its own -- which is how adapters - actually travel. + A pattern is only ever used to *exclude* tensors from factorization, so one that + matches nothing is indistinguishable from not passing it at all -- the usual cause + is shell escaping, where a doubled backslash makes ``\\.`` mean "backslash, any + character" instead of a literal dot. """ - cleaned = {k: v.detach().cpu().contiguous() for k, v in adapter_state.items()} - out_str = str(out_path) - if out_path.suffix == ".safetensors" and _HAVE_SAFETENSORS: - safetensors_save(cleaned, out_str, metadata=metadata or None) - else: - torch.save(cleaned, out_str) + keys = sorted(keys) + unmatched = [pattern for pattern in patterns if not any(re.search(pattern, key) for key in keys)] + if unmatched: + rendered = ", ".join(repr(pattern) for pattern in unmatched) + raise ValueError( + "--exact-tensor-pattern matched no tensor in the fine-tuned checkpoint: " + rendered + + ". Those tensors would be rank-truncated instead of kept exact; check the escaping " + r"(inside shell single quotes write '\.', not '\\.').") + + +def _should_factor( + key: str, + shape: tuple[int, ...], + exact_patterns: Sequence[re.Pattern[str]], +) -> bool: + if not is_extractable_weight(key) or len(shape) != 2: + return False + return not any(pattern.search(key) for pattern in exact_patterns) -def build_adapter_from_states( - base_sd: Dict[str, torch.Tensor], - ft_sd: Dict[str, torch.Tensor], +def _factorize_delta( + delta: torch.Tensor, rank: int, full_rank: bool, - min_delta: float, - checkpoint_interval: int, - checkpoint_path: Optional[Path], - resume_from: int = 0, -) -> Dict[str, torch.Tensor]: - """Compute low-rank LoRA adapters by SVD on (ft - base) for extractable weights.""" - # DTensor detection - try: - from torch.distributed.tensor import DTensor # type: ignore - _HAS_DTENSOR = True - except Exception: - DTensor = None # type: ignore - _HAS_DTENSOR = False - - keys = sorted(ft_sd.keys()) - adapter_state: Dict[str, torch.Tensor] = {} - processed = 0 - mean_deltas = [] - - for idx, key in enumerate(tqdm(keys, desc="scanning keys", unit="keys")): - if idx < resume_from: - continue - if not is_extractable_weight(key): - continue - if key not in base_sd: - continue - - Wb_raw = base_sd[key] - Wf_raw = ft_sd[key] - - # Convert DTensor if present - if _HAS_DTENSOR and isinstance(Wb_raw, DTensor): # type: ignore - Wb = Wb_raw.to_local().detach().cpu().to(torch.float32).contiguous() + method: str, + randomized_q: int | None, + oversample: int, + niter: int, + seed: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, str]: + available_rank = min(delta.shape) + chosen_rank = available_rank if full_rank or rank <= 0 else min(rank, available_rank) + if chosen_rank == 0: + raise ValueError(f"Cannot factorize empty matrix with shape {tuple(delta.shape)}") + + if method == "exact": + u, singular_values, vh = torch.linalg.svd(delta, full_matrices=False) + v = vh.mT + method_description = "exact" + else: + q = randomized_q if randomized_q is not None else chosen_rank + oversample + q = min(available_rank, max(chosen_rank, q)) + if q == available_rank: + u, singular_values, vh = torch.linalg.svd(delta, full_matrices=False) + v = vh.mT + method_description = "exact-full-basis" else: - Wb = Wb_raw.detach().cpu().to(torch.float32).contiguous() + devices = [delta.device] if delta.device.type == "cuda" else [] + with torch.random.fork_rng(devices=devices): + torch.manual_seed(seed) + if delta.device.type == "cuda": + torch.cuda.manual_seed(seed) + u, singular_values, v = torch.svd_lowrank(delta, q=q, niter=niter) + method_description = f"randomized-q{q}-niter{niter}" - if _HAS_DTENSOR and isinstance(Wf_raw, DTensor): # type: ignore - Wf = Wf_raw.to_local().detach().cpu().to(torch.float32).contiguous() - else: - Wf = Wf_raw.detach().cpu().to(torch.float32).contiguous() + singular_values = singular_values[:chosen_rank].float() + sqrt_s = singular_values.sqrt() + lora_b = (u[:, :chosen_rank].float() * sqrt_s.unsqueeze(0)).contiguous() + lora_a = (v[:, :chosen_rank].float() * sqrt_s.unsqueeze(0)).mT.contiguous() + return lora_a, lora_b, singular_values, method_description - if Wb.shape != Wf.shape: - continue - delta = (Wf - Wb).contiguous() - mean_abs = float(delta.abs().mean().item()) - mean_deltas.append(mean_abs) - if mean_abs < min_delta: - continue +def _work_namespace(out_path: Path) -> str: + resolved = str(out_path.expanduser().resolve(strict=False)) + digest = hashlib.sha256(resolved.encode("utf-8")).hexdigest()[:12] + stem = re.sub(r"[^A-Za-z0-9_.-]+", "-", out_path.stem).strip("-.") or "adapter" + return f"{stem}-{digest}" - # SVD (CPU) - try: - U, S, Vh = torch.linalg.svd(delta, full_matrices=False) - except RuntimeError: - # skip layers that fail SVD - continue - max_rank = S.numel() - chosen_rank = max_rank if full_rank or rank <= 0 else min(rank, max_rank) - if chosen_rank == 0: - continue - - S_sqrt = torch.sqrt(S[:chosen_rank].to(torch.float32)) - U_r = U[:, :chosen_rank].to(torch.float32) # (out, r) - Vh_r = Vh[:chosen_rank, :].to(torch.float32) # (r, in) - - lora_B = (U_r * S_sqrt.unsqueeze(0)).contiguous() # (out, r) - tmp = (Vh_r.T * S_sqrt.unsqueeze(0)).contiguous() # (in, r) - lora_A = tmp.T.contiguous() # (r, in) - - base_name = key[:-len(".weight")] - a_key = f"{base_name}.lora_A.weight" - b_key = f"{base_name}.lora_B.weight" - rank_key = f"{base_name}.lora_rank" - alpha_key = f"{base_name}.lora_alpha" - - adapter_state[a_key] = lora_A.cpu() - adapter_state[b_key] = lora_B.cpu() - adapter_state[rank_key] = torch.tensor([chosen_rank], dtype=torch.int32) - adapter_state[alpha_key] = torch.tensor([float(chosen_rank)], dtype=torch.float32) +def _clear_work_dir(work_dir: Path) -> None: + """Remove the scratch artifacts this script writes, never the directory itself. - processed += 1 - - # checkpoint periodically - if checkpoint_path and checkpoint_interval > 0 and (idx + 1) % checkpoint_interval == 0: - try: - torch.save({"index": idx + 1, "adapter": adapter_state}, str(checkpoint_path)) - except Exception: - # non-fatal; continue - pass + Assembly is driven by the manifest rather than by a directory listing, so dropping + the manifest is what makes a rerun start clean; the tensor shards only cost disk. + """ + shutil.rmtree(work_dir / "tensors", ignore_errors=True) + (work_dir / "manifest.json").unlink(missing_ok=True) + + +def _prepare_work_dir(work_dir: Path, config: ExtractionConfig, resume: bool) -> tuple[Path, dict[str, Any]]: + manifest_path = work_dir / "manifest.json" + expected = asdict(config) + expected["exact_tensor_patterns"] = list(config.exact_tensor_patterns) + if resume and manifest_path.is_file(): + manifest = json.loads(manifest_path.read_text(encoding="utf-8")) + if manifest.get("config") != expected: + raise ValueError(f"Resume configuration does not match {manifest_path}") + return manifest_path, manifest + if not resume: + _clear_work_dir(work_dir) + + (work_dir / "tensors").mkdir(parents=True, exist_ok=True) + manifest = {"format": FORMAT_VERSION, "config": expected, "layers": {}} + _atomic_json_dump(manifest, manifest_path) + return manifest_path, manifest + + +def _validate_key_sets(base: TensorReader, finetuned: TensorReader) -> None: + missing = sorted(base.keys - finetuned.keys) + if missing: + raise ValueError(f"Fine-tuned checkpoint is missing {len(missing)} base tensors; first keys: {missing[:5]}") + for key in sorted(base.keys & finetuned.keys): + if base.get_shape(key) != finetuned.get_shape(key): + raise ValueError( + f"Shape mismatch for {key}: base={base.get_shape(key)}, finetuned={finetuned.get_shape(key)}") + + +def _save_layer_payload(path: Path, payload: dict[str, torch.Tensor], source_key: str) -> None: + save_file( + {key: tensor.detach().cpu().contiguous() for key, tensor in payload.items()}, + str(path), + metadata={"source_key": source_key, "format": FORMAT_VERSION}, + ) - # free local large tensors - del delta, U, S, Vh, U_r, Vh_r, tmp, lora_A, lora_B - # final checkpoint - if checkpoint_path: - try: - torch.save({"index": len(keys), "adapter": adapter_state}, str(checkpoint_path)) - except Exception: - pass +def _extract_layers( + base: TensorReader, + finetuned: TensorReader, + work_dir: Path, + manifest_path: Path, + manifest: dict[str, Any], + config: ExtractionConfig, +) -> None: + device = torch.device(config.device) + if device.type == "cuda" and not torch.cuda.is_available(): + raise RuntimeError(f"CUDA device requested but CUDA is unavailable: {device}") + factor_dtype = _torch_dtype(config.factor_dtype) + exact_patterns = _compile_patterns(config.exact_tensor_patterns) + + for index, key in enumerate(tqdm(sorted(finetuned.keys), desc="extracting LoRA", unit="tensor")): + tensor_file = work_dir / "tensors" / f"{index:05d}.safetensors" + existing = manifest["layers"].get(key) + if existing is not None and (existing.get("tensor_file") is None or tensor_file.is_file()): + continue - avg_delta = (sum(mean_deltas) / len(mean_deltas)) if mean_deltas else 0.0 - LOG.info("Extraction complete: processed_keys=%d, extracted_layers=%d, avg_abs_delta=%.6e", len(keys), processed, - avg_delta) - return adapter_state + finetuned_tensor = finetuned.get_tensor(key) + if key not in base.keys: + if not config.dense_payload: + manifest["layers"][key] = { + "kind": "skipped", + "shape": list(finetuned_tensor.shape), + "tensor_file": None, + "reason": "dense payload disabled", + } + _atomic_json_dump(manifest, manifest_path) + continue + is_weight = key.endswith(".weight") + output_key = (key.removesuffix(".weight") + SET_WEIGHT_SUFFIX if is_weight else key + SET_PARAM_SUFFIX) + output_dtype = _resolve_output_dtype(config.replacement_dtype, finetuned_tensor.dtype) + _save_layer_payload(tensor_file, {output_key: finetuned_tensor.to(output_dtype)}, key) + manifest["layers"][key] = { + "kind": "set_weight" if is_weight else "set_param", + "shape": list(finetuned_tensor.shape), + "tensor_file": tensor_file.name, + "output_keys": [output_key], + } + _atomic_json_dump(manifest, manifest_path) + continue + base_tensor = base.get_tensor(key) + delta = finetuned_tensor.to(device=device, dtype=torch.float32, copy=True) + delta.sub_(base_tensor.to(device=device, dtype=torch.float32)) + max_abs_delta = float(delta.abs().max().item()) if delta.numel() else 0.0 + if max_abs_delta <= config.min_delta: + manifest["layers"][key] = { + "kind": "unchanged", + "shape": list(delta.shape), + "tensor_file": None, + "max_abs_delta": max_abs_delta, + } + _atomic_json_dump(manifest, manifest_path) + del delta, base_tensor, finetuned_tensor + continue -def parse_args() -> argparse.Namespace: - p = argparse.ArgumentParser(description="Extract FastVideo-style LoRA adapter (CPU SVD).") - p.add_argument("--base", required=True, help="Base model id or local path") - p.add_argument("--ft", required=True, help="Fine-tuned model id or local path") - p.add_argument("--out", default="fastvideo_adapter.safetensors", help="Output adapter file (.safetensors or .pt)") - p.add_argument("--rank", type=int, default=16, help="Truncated SVD rank; <=0 for full rank") - p.add_argument("--full-rank", action="store_true", help="Use full SVD rank for every layer") - p.add_argument("--min-delta", type=float, default=1e-8, help="Minimum mean abs delta to consider a layer changed") - p.add_argument("--checkpoint", default="extract_lora_checkpoint.pt", help="Checkpoint path to resume/save progress") - p.add_argument("--resume", action="store_true", help="Resume from checkpoint if available") - p.add_argument("--no-dense-payload", - dest="dense_payload", - action="store_false", - help="Emit only low-rank factors, dropping changed norms/biases and any parameter " - "the base model lacks (pre-2026 behaviour)") - p.add_argument("--log-level", default="INFO", choices=["DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL"]) - return p.parse_args() + shape = tuple(delta.shape) + if _should_factor(key, shape, exact_patterns): + delta_fro_sq = float(delta.double().square().sum().item()) + lora_a, lora_b, singular_values, method = _factorize_delta( + delta, + rank=config.rank, + full_rank=config.full_rank, + method=config.svd_method, + randomized_q=config.randomized_q, + oversample=config.oversample, + niter=config.niter, + seed=_seed_for_key(config.seed, key), + ) + module_name = key.removesuffix(".weight") + actual_rank = lora_a.shape[0] + output_keys = [ + f"{module_name}.lora_A.weight", + f"{module_name}.lora_B.weight", + ] + # The extracted factors use alpha == rank, which is the loader's + # default when no scalar is present. Emitting per-layer rank/alpha + # bookkeeping adds hundreds of keys, and several adapter naming + # schemes map those scalars differently from their A/B factors. + payload = { + output_keys[0]: lora_a.to(factor_dtype), + output_keys[1]: lora_b.to(factor_dtype), + } + _save_layer_payload(tensor_file, payload, key) + captured = float(singular_values.double().square().sum().item()) + residual = (max(0.0, 1.0 - captured / delta_fro_sq)**0.5) if delta_fro_sq else 0.0 + manifest["layers"][key] = { + "kind": "lora", + "shape": list(shape), + "tensor_file": tensor_file.name, + "output_keys": output_keys, + "rank": actual_rank, + "method": method, + "delta_frobenius_norm": delta_fro_sq**0.5, + "relative_residual": residual, + "max_abs_delta": max_abs_delta, + } + del lora_a, lora_b, singular_values, payload + else: + output_key = dense_payload_key(key) + if config.dense_payload: + output_dtype = _resolve_output_dtype(config.dense_dtype, finetuned_tensor.dtype) + _save_layer_payload(tensor_file, {output_key: delta.to(output_dtype)}, key) + manifest["layers"][key] = { + "kind": "diff", + "shape": list(shape), + "tensor_file": tensor_file.name, + "output_keys": [output_key], + "max_abs_delta": max_abs_delta, + } + else: + manifest["layers"][key] = { + "kind": "skipped", + "shape": list(shape), + "tensor_file": None, + "reason": "dense payload disabled", + "max_abs_delta": max_abs_delta, + } + + _atomic_json_dump(manifest, manifest_path) + del delta, base_tensor, finetuned_tensor + if device.type == "cuda": + torch.cuda.empty_cache() + + +def _assemble_adapter( + out_path: Path, + work_dir: Path, + manifest: dict[str, Any], + metadata: dict[str, str], +) -> None: + adapter: dict[str, torch.Tensor] = {} + for report in tqdm(manifest["layers"].values(), desc="assembling adapter", unit="tensor"): + tensor_file = report.get("tensor_file") + if tensor_file is None: + continue + payload = load_file(str(work_dir / "tensors" / tensor_file), device="cpu") + overlap = set(adapter) & set(payload) + if overlap: + raise ValueError(f"Duplicate adapter keys while assembling {out_path}: {sorted(overlap)[:5]}") + adapter.update(payload) + out_path.parent.mkdir(parents=True, exist_ok=True) + temporary = out_path.with_suffix(out_path.suffix + ".tmp") + save_file(adapter, str(temporary), metadata=metadata) + temporary.replace(out_path) + out_path.chmod(0o644) + + +def _verify_adapter(out_path: Path, manifest: dict[str, Any]) -> None: + expected_keys = { + key + for layer in manifest["layers"].values() + for key in layer.get("output_keys", []) + } + with safe_open(out_path, framework="pt") as adapter: + actual_keys = set(adapter.keys()) + if actual_keys != expected_keys: + missing = sorted(expected_keys - actual_keys) + extra = sorted(actual_keys - expected_keys) + raise ValueError(f"Output key mismatch: missing={missing[:5]}, extra={extra[:5]}") + for source_key, layer in manifest["layers"].items(): + kind = layer["kind"] + shape = tuple(layer["shape"]) + if kind == "lora": + module_name = source_key.removesuffix(".weight") + rank = int(layer["rank"]) + a_shape = tuple(adapter.get_slice(f"{module_name}.lora_A.weight").get_shape()) + b_shape = tuple(adapter.get_slice(f"{module_name}.lora_B.weight").get_shape()) + if a_shape != (rank, shape[1]) or b_shape != (shape[0], rank): + raise ValueError(f"Invalid factor shapes for {source_key}: A={a_shape}, B={b_shape}") + elif kind in {"diff", "set_weight", "set_param"}: + output_key = layer["output_keys"][0] + output_shape = tuple(adapter.get_slice(output_key).get_shape()) + if output_shape != shape: + raise ValueError(f"Invalid dense payload shape for {source_key}: {output_shape} != {shape}") + + +def _build_report(manifest: dict[str, Any], out_path: Path) -> dict[str, Any]: + counts: dict[str, int] = {} + delta_energy = 0.0 + residual_energy = 0.0 + for layer in manifest["layers"].values(): + kind = layer["kind"] + counts[kind] = counts.get(kind, 0) + 1 + if kind == "lora": + norm = float(layer["delta_frobenius_norm"]) + residual = float(layer["relative_residual"]) + delta_energy += norm * norm + residual_energy += (norm * residual)**2 + weighted_residual = (residual_energy / delta_energy)**0.5 if delta_energy else 0.0 + return { + "format": FORMAT_VERSION, + "adapter": str(out_path.resolve()), + "adapter_size_bytes": out_path.stat().st_size, + "counts": counts, + "factorized_weighted_relative_residual": weighted_residual, + "config": manifest["config"], + "layers": manifest["layers"], + } def extract_lora_adapter( @@ -425,111 +767,163 @@ def extract_lora_adapter( rank: int = 32, full_rank: bool = False, min_delta: float = 1e-6, - checkpoint: Optional[str] = None, + checkpoint: str | None = None, resume: bool = False, log_level: str = "INFO", dense_payload: bool = True, -) -> None: - """Extract LoRA adapter from fine-tuned model. - - Args: - base: Base model path or HuggingFace ID - ft: Fine-tuned model path or HuggingFace ID - out: Output adapter file path - rank: LoRA rank (default: 32) - full_rank: Extract full-rank adapter - min_delta: Minimum delta for extraction - checkpoint: Checkpoint file path - resume: Resume from checkpoint - log_level: Logging level - dense_payload: Also capture parameters the SVD path cannot represent -- norms, - biases, and anything absent from the base model -- as `.diff` / `.diff_b` / - `.set_weight` keys. On by default: without it a distillation that retunes - its norms, or a VSA student whose compression gate exists only after - distillation, is silently reproduced without those changes. - """ + *, + base_revision: str | None = None, + ft_revision: str | None = None, + load_mode: str = "auto", + device: str = "cpu", + svd_method: str = "exact", + randomized_q: int | None = None, + oversample: int = 64, + niter: int = 4, + seed: int = 42, + factor_dtype: str = "float32", + dense_dtype: str = "float32", + replacement_dtype: str = "source", + exact_tensor_patterns: Sequence[str] = (), + work_dir: str | None = None, + keep_work_dir: bool = False, +) -> Path: + """Extract one adapter while streaming transformer tensors.""" configure_logging(log_level) + if load_mode not in {"auto", "indexed", "pipeline"}: + raise ValueError(f"Unsupported load mode: {load_mode}") + if svd_method not in {"exact", "randomized"}: + raise ValueError(f"Unsupported SVD method: {svd_method}") + if randomized_q is not None and randomized_q < 1: + raise ValueError("randomized_q must be positive") + if resume and load_mode == "pipeline": + raise ValueError("--resume requires indexed safetensors; pipeline loading cannot validate checkpoint identity") + reader_load_mode = "indexed" if resume else load_mode + + out_path = Path(out).expanduser() + if out_path.suffix != ".safetensors": + raise ValueError(f"--out must end in .safetensors; adapters are always written as safetensors: {out_path}") + work_root: Path | None = None + if work_dir is not None: + # --work-dir is a shared root; each output gets an independent namespace. + work_root = Path(work_dir).expanduser() / WORK_SUBDIR + effective_work_dir = work_root / _work_namespace(out_path) + elif checkpoint is not None: + effective_work_dir = Path(checkpoint).expanduser().with_suffix(".work") + else: + effective_work_dir = out_path.parent / f".{out_path.name}.work" + + with _open_readers(base, ft, base_revision, ft_revision, reader_load_mode) as (base_reader, finetuned_reader): + if resume and (not isinstance(base_reader, IndexedSafetensorsReader) + or not isinstance(finetuned_reader, IndexedSafetensorsReader)): + raise ValueError("--resume requires indexed safetensors so checkpoint identity can be validated") + _validate_key_sets(base_reader, finetuned_reader) + _validate_exact_patterns(exact_tensor_patterns, finetuned_reader.keys) + config = ExtractionConfig( + base_source=base_reader.source, + finetuned_source=finetuned_reader.source, + base_fingerprint=base_reader.fingerprint, + finetuned_fingerprint=finetuned_reader.fingerprint, + rank=rank, + full_rank=full_rank, + min_delta=min_delta, + device=str(torch.device(device)), + svd_method=svd_method, + randomized_q=randomized_q, + oversample=oversample, + niter=niter, + seed=seed, + factor_dtype=str(_torch_dtype(factor_dtype)).removeprefix("torch."), + dense_dtype=dense_dtype, + replacement_dtype=replacement_dtype, + dense_payload=dense_payload, + exact_tensor_patterns=tuple(exact_tensor_patterns), + ) + manifest_path, manifest = _prepare_work_dir(effective_work_dir, config, resume) + _extract_layers(base_reader, finetuned_reader, effective_work_dir, manifest_path, manifest, config) + base_reader.assert_unchanged() + finetuned_reader.assert_unchanged() + + counts: dict[str, int] = {} + for layer in manifest["layers"].values(): + counts[layer["kind"]] = counts.get(layer["kind"], 0) + 1 + metadata = { + "format": FORMAT_VERSION, + "base_model": base, + "base_revision": base_revision or "unspecified", + "finetuned_model": ft, + "finetuned_revision": ft_revision or "unspecified", + "requested_rank": str(rank), + "full_rank": str(bool(full_rank)), + "svd_method": svd_method, + "randomized_q": str(randomized_q) if randomized_q is not None else "automatic", + "niter": str(niter), + "seed": str(seed), + "factor_dtype": config.factor_dtype, + "dense_diff_dtype": dense_dtype, + "replacement_dtype": replacement_dtype, + "lora_layers": str(counts.get("lora", 0)), + "diff_tensors": str(counts.get("diff", 0)), + "set_weight_tensors": str(counts.get("set_weight", 0)), + "set_param_tensors": str(counts.get("set_param", 0)), + "dropped_unchanged": str(counts.get("unchanged", 0)), + "application": "W = W_base + lora_B @ lora_A; then dense diffs added and replacements assigned", + } + _assemble_adapter(out_path, effective_work_dir, manifest, metadata) + _verify_adapter(out_path, manifest) + report = _build_report(manifest, out_path) + report_path = out_path.with_suffix(out_path.suffix + ".report.json") + _atomic_json_dump(report, report_path) + LOG.info("Saved adapter to %s (%.2f GiB); report=%s", out_path, out_path.stat().st_size / 2**30, report_path) + + if not keep_work_dir: + _clear_work_dir(effective_work_dir) + with suppress(OSError): + effective_work_dir.rmdir() + if work_root is not None: + with suppress(OSError): + work_root.rmdir() + return out_path - out_path = Path(out) - checkpoint_path = Path(checkpoint) if checkpoint else None - - # Load both models - ensure consistent loading method for matching keys - try: - import fastvideo # noqa: F401 - LOG.info("Loading base model via pipeline: %s", base) - base_sd = load_transformer_state_dict_from_model(base) - LOG.info("Loading fine-tuned model via pipeline: %s", ft) - ft_sd = load_transformer_state_dict_from_model(ft) - except Exception as exc: - LOG.warning("Pipeline loading failed: %s", exc) - LOG.info("Falling back to direct safetensors loading for BOTH models...") - # Direct loading - both models use same method for consistent keys - LOG.info("Loading base model directly from safetensors: %s", base) - base_sd = load_transformer_state_dict_from_safetensors(base) - LOG.info("Loading fine-tuned model directly from safetensors: %s", ft) - ft_sd = load_transformer_state_dict_from_safetensors(ft) - - resume_idx = 0 - adapter_existing: Dict[str, torch.Tensor] = {} - if resume and checkpoint_path and checkpoint_path.exists(): - try: - ck = torch.load(str(checkpoint_path), map_location="cpu") - adapter_existing = ck.get("adapter", {}) or {} - resume_idx = int(ck.get("index", 0) or 0) - LOG.info("Resuming from checkpoint index=%d with %d existing entries", resume_idx, len(adapter_existing)) - except Exception: - adapter_existing = {} - - adapter_state = dict(adapter_existing) if adapter_existing else {} - new_adapter = build_adapter_from_states( - base_sd=base_sd, - ft_sd=ft_sd, - rank=rank, - full_rank=full_rank, - min_delta=min_delta, - checkpoint_interval=50, - checkpoint_path=checkpoint_path, - resume_from=resume_idx, - ) - adapter_state.update(new_adapter) - - if dense_payload: - low_rank_keys = {k.split(".lora_")[0] + ".weight" for k in adapter_state if ".lora_" in k} - adapter_state.update( - build_dense_payload( - base_sd=base_sd, - ft_sd=ft_sd, - low_rank_keys=low_rank_keys, - min_delta=min_delta, - )) - - # final save - save_adapter_state( - adapter_state, - out_path, - metadata={ - "base_model": base, - "finetuned_model": ft, - "requested_rank": str(rank), - "full_rank": str(bool(full_rank)), - "dense_payload": str(bool(dense_payload)), - "application": "W_effective = W_base + lora_B @ lora_A, then .diff added and .set_weight assigned", - "format": "fastvideo-lora-v2", - }, - ) - - # cleanup checkpoint if present - if checkpoint_path and checkpoint_path.exists(): - try: - checkpoint_path.unlink() - except Exception: - pass - LOG.info("Saved adapter to %s (entries=%d)", str(out_path), len(adapter_state) // 4) +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base", required=True, help="Base model ID or local path") + parser.add_argument("--ft", required=True, help="Fine-tuned model ID or local path") + parser.add_argument("--out", default="fastvideo_adapter.safetensors") + parser.add_argument("--base-revision", default=None) + parser.add_argument("--ft-revision", default=None) + parser.add_argument("--rank", type=int, default=16) + parser.add_argument("--full-rank", action="store_true") + parser.add_argument("--min-delta", type=float, default=1e-8, + help="Drop tensors whose maximum absolute FP32 delta does not exceed this value") + parser.add_argument("--load-mode", choices=("auto", "indexed", "pipeline"), default="auto") + parser.add_argument("--device", default="cpu", help="Factorization device, for example cpu or cuda:0") + parser.add_argument("--svd-method", choices=("exact", "randomized"), default="exact") + parser.add_argument("--randomized-q", type=int, default=None, + help="Randomized basis width; q=320 is validated for rank-64 MiniMax-H3") + parser.add_argument("--oversample", type=int, default=64, + help="Used when --randomized-q is omitted: q = rank + oversample") + parser.add_argument("--niter", type=int, default=4) + parser.add_argument("--seed", type=int, default=42) + parser.add_argument("--factor-dtype", choices=("float32", "float16", "bfloat16"), default="float32") + parser.add_argument("--dense-dtype", choices=("source", "float32", "float16", "bfloat16"), default="float32") + parser.add_argument("--replacement-dtype", choices=("source", "float32", "float16", "bfloat16"), + default="source") + parser.add_argument("--exact-tensor-pattern", action="append", default=[], + help="Regex for a tensor to retain as an exact dense delta instead of factorizing; repeatable") + parser.add_argument("--work-dir", default=None) + parser.add_argument("--checkpoint", default=None, + help="Deprecated work-directory alias retained for CLI compatibility") + parser.add_argument("--resume", action="store_true") + parser.add_argument("--keep-work-dir", action="store_true") + parser.add_argument("--no-dense-payload", dest="dense_payload", action="store_false", + help="Emit low-rank factors only; changed norms/biases and fine-tuned-only weights are omitted") + parser.add_argument("--log-level", default="INFO", choices=("DEBUG", "INFO", "WARNING", "ERROR", "CRITICAL")) + return parser.parse_args() def main() -> None: - """CLI wrapper for extract_lora_adapter.""" args = parse_args() extract_lora_adapter( base=args.base, @@ -542,6 +936,21 @@ def main() -> None: resume=args.resume, log_level=args.log_level, dense_payload=args.dense_payload, + base_revision=args.base_revision, + ft_revision=args.ft_revision, + load_mode=args.load_mode, + device=args.device, + svd_method=args.svd_method, + randomized_q=args.randomized_q, + oversample=args.oversample, + niter=args.niter, + seed=args.seed, + factor_dtype=args.factor_dtype, + dense_dtype=args.dense_dtype, + replacement_dtype=args.replacement_dtype, + exact_tensor_patterns=args.exact_tensor_pattern, + work_dir=args.work_dir, + keep_work_dir=args.keep_work_dir, )