From 22bcc38afbf49b067a4408ff57416897048f2a57 Mon Sep 17 00:00:00 2001 From: shaoxiongduan Date: Sat, 29 Aug 2026 12:12:10 +0000 Subject: [PATCH 1/6] [feat]: add streaming GPU LoRA extraction --- docs/training/finetune.md | 7 + docs/utilities/lora.md | 39 +- .../test_streaming_lora_extraction.py | 180 +++ scripts/lora_extraction/README.md | 53 +- scripts/lora_extraction/extract_lora.py | 1117 +++++++++++------ 5 files changed, 970 insertions(+), 426 deletions(-) create mode 100644 fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py diff --git a/docs/training/finetune.md b/docs/training/finetune.md index 7389242bf6..6777736692 100644 --- a/docs/training/finetune.md +++ b/docs/training/finetune.md @@ -131,6 +131,13 @@ python scripts/lora_extraction/extract_lora.py \ | `--out` | Output adapter file (.safetensors) | | `--rank` | LoRA rank (16, 32, 64, 128) | | `--full-rank` | Extract full-rank adapter (optional) | +| `--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` payloads | + +For large checkpoints, indexed loading streams one transformer tensor pair at a time and downloads only `transformer/*`. The extractor also preserves changed norms and biases as exact deltas and fine-tuned-only weights as `.set_weight`. See [LoRA Extraction and Merging](../utilities/lora.md) for the GPU/randomized-SVD command, resume options, and accuracy controls. ### Merge LoRA Adapter diff --git a/docs/utilities/lora.md b/docs/utilities/lora.md index e9109db7d9..a25083da8e 100644 --- a/docs/utilities/lora.md +++ b/docs/utilities/lora.md @@ -12,13 +12,40 @@ python scripts/lora_extraction/extract_lora.py \ --rank 32 ``` -**Options:** +Exact CPU SVD remains the default. For a large transformer, stream its indexed safetensors and factorize on a GPU: -- `--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) +```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. +- `--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-tensor-pattern`: repeatable regex for a matrix that should remain an exact dense delta. +- `--work-dir`, `--resume`: resume an interrupted streaming extraction. + +The adapter retains changes that do not fit a low-rank product: `.diff` and `.diff_b` hold exact additive weight/bias deltas, while `.set_weight` holds a parameter absent from the base checkpoint, such as a VSA compression gate. 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). ## Merge Adapter 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..62e8d72a83 --- /dev/null +++ b/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py @@ -0,0 +1,180 @@ +"""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 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" + + +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_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_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/*"], + }] diff --git a/scripts/lora_extraction/README.md b/scripts/lora_extraction/README.md index 01d82f3f1f..f5863a3723 100644 --- a/scripts/lora_extraction/README.md +++ b/scripts/lora_extraction/README.md @@ -4,6 +4,8 @@ Tools for extracting and merging LoRA adapters 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 \ @@ -12,15 +14,52 @@ python extract_lora.py \ --rank 32 ``` -**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) +For large transformers, stream their indexed safetensors and factorize on a GPU: + +```bash +python extract_lora.py \ + --base MiniMaxAI/MiniMax-H3 \ + --ft FastVideo/FastVideo-FastH3-8-step-Preview-v1-VSA-DataFree \ + --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. + +Important options: + +- `--base`, `--ft`: Hugging Face model IDs or local paths. +- `--rank`: requested LoRA rank. +- `--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` and `.diff_b` payloads. +- `--replacement-dtype`: storage dtype for fine-tuned-only `.set_weight` parameters. +- `--exact-tensor-pattern`: repeatable regex selecting matrices to retain as exact dense deltas. +- `--work-dir`, `--resume`: resume a partially completed streaming extraction. + +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`; +- a fine-tuned-only weight, such as a VSA compression gate, becomes `.set_weight`; +- a bit-identical parameter is omitted. -> **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. +Indexed loading is preferred and downloads only the transformer component. `--load-mode auto` falls back to legacy pipeline loading when indexed safetensors are unavailable. ## Merge Adapter diff --git a/scripts/lora_extraction/extract_lora.py b/scripts/lora_extraction/extract_lora.py index d506bf43bc..4200915191 100644 --- a/scripts/lora_extraction/extract_lora.py +++ b/scripts/lora_extraction/extract_lora.py @@ -1,114 +1,236 @@ -"""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`` store exact additive deltas. +* ``.set_weight`` stores a parameter 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 Iterator, Sequence +from contextlib import ExitStack, contextmanager +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" +DIFF_SUFFIX = ".diff" +DIFF_BIAS_SUFFIX = ".diff_b" +SET_WEIGHT_SUFFIX = ".set_weight" - # Download or locate the model - if os.path.isdir(model_path): - local_path = model_path - else: - local_path = snapshot_download(model_path) +_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, +} - # 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}") - # 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) +class TensorReader(Protocol): + """Random access to one checkpoint's transformer tensors.""" - if not state_dict: - raise ValueError(f"No safetensors files found in {transformer_dir}") + source: str - LOG.info("Loaded %d keys directly from safetensors", len(state_dict)) - return state_dict + @property + def keys(self) -> set[str]: ... + def get_tensor(self, key: str) -> torch.Tensor: ... -# Configure minimal logging -LOG = logging.getLogger("extract_lora") + def get_shape(self, key: str) -> tuple[int, ...]: ... + + 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 + + @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 __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}") + + self._stack = ExitStack() + self._shards = { + shard: self._stack.enter_context(safe_open(transformer_dir / shard, framework="pt", device="cpu")) + for shard in sorted(set(self.weight_map.values())) + } + + @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 __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 + 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 _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 (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,8 +240,8 @@ 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, @@ -130,292 +252,416 @@ def load_transformer_state_dict_from_model( 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") - if transformer is None: - pipeline_attr = getattr(pipeline, "pipeline", None) - transformer = getattr(pipeline_attr, "transformer", None) if pipeline_attr else None + transformer = modules.get("transformer") if isinstance(modules, dict) else None if transformer is None: - raise RuntimeError( - "Transformer not found in pipeline. Expected pipeline.transformer or pipeline.modules['transformer'].") + raise RuntimeError("Transformer not found in pipeline") - state_dict = transformer.state_dict() - - # 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]]: + 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": + 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 - - -# 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" + lowered = key.lower() + return not any(fragment in lowered for fragment in ("norm", "bias", "embedding")) -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 | None: 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 param_name.removesuffix(".bias") + DIFF_BIAS_SUFFIX return None 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: + if key.endswith(".weight"): + payload[key.removesuffix(".weight") + SET_WEIGHT_SUFFIX] = 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. - 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. - """ - 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) +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() - - if _HAS_DTENSOR and isinstance(Wf_raw, DTensor): # type: ignore - Wf = Wf_raw.to_local().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}" + + 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 + + +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 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}") + if not resume: + shutil.rmtree(work_dir) else: - Wf = Wf_raw.detach().cpu().to(torch.float32).contiguous() + return manifest_path, manifest + elif work_dir.exists() and not resume: + shutil.rmtree(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}, + ) - 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: +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 - # SVD (CPU) - try: - U, S, Vh = torch.linalg.svd(delta, full_matrices=False) - except RuntimeError: - # skip layers that fail SVD + 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 + if not key.endswith(".weight"): + raise ValueError(f"Fine-tuned-only parameter has no supported replacement suffix: {key}") + output_key = key.removesuffix(".weight") + SET_WEIGHT_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", + "shape": list(finetuned_tensor.shape), + "tensor_file": tensor_file.name, + "output_keys": [output_key], + } + _atomic_json_dump(manifest, manifest_path) continue - max_rank = S.numel() - chosen_rank = max_rank if full_rank or rank <= 0 else min(rank, max_rank) - if chosen_rank == 0: + 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 - 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) - - 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 - - # 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 - - 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 - - -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 output_key is None: + raise ValueError(f"Changed parameter has no supported dense suffix: {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"}: + 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 +671,141 @@ 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 = "source", + 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") + + out_path = Path(out).expanduser() + if work_dir is not None: + effective_work_dir = Path(work_dir).expanduser() + 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, load_mode) as (base_reader, finetuned_reader): + _validate_key_sets(base_reader, finetuned_reader) + config = ExtractionConfig( + base_source=base_reader.source, + finetuned_source=finetuned_reader.source, + 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) + + 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)), + "dropped_unchanged": str(counts.get("unchanged", 0)), + "application": "W = W_base + lora_B @ lora_A; then .diff/.diff_b added and .set_weight 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: + shutil.rmtree(effective_work_dir) + 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="source") + 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 +818,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, ) From 78948fb757a6e7218365ca3be57bc0cba1192221 Mon Sep 17 00:00:00 2001 From: shaoxiongduan Date: Sat, 29 Aug 2026 12:59:31 +0000 Subject: [PATCH 2/6] [bugfix]: preserve standalone LoRA parameters --- docs/training/finetune.md | 4 +- docs/utilities/lora.md | 2 +- fastvideo/models/loader/lora_patch.py | 21 +-- fastvideo/tests/loader/test_lora_patch.py | 32 +++- .../lora_extraction/test_lora_extraction.py | 28 +++- .../test_streaming_lora_extraction.py | 142 ++++++++++++++++++ scripts/lora_extraction/README.md | 6 +- scripts/lora_extraction/extract_lora.py | 24 +-- 8 files changed, 224 insertions(+), 35 deletions(-) diff --git a/docs/training/finetune.md b/docs/training/finetune.md index 6777736692..fee384088f 100644 --- a/docs/training/finetune.md +++ b/docs/training/finetune.md @@ -135,9 +135,9 @@ python scripts/lora_extraction/extract_lora.py \ | `--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` payloads | +| `--dense-dtype` | Storage dtype for exact `.diff`/`.diff_b`/`.diff_param` payloads | -For large checkpoints, indexed loading streams one transformer tensor pair at a time and downloads only `transformer/*`. The extractor also preserves changed norms and biases as exact deltas and fine-tuned-only weights as `.set_weight`. See [LoRA Extraction and Merging](../utilities/lora.md) for the GPU/randomized-SVD command, resume options, and accuracy controls. +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`. See [LoRA Extraction and Merging](../utilities/lora.md) for the GPU/randomized-SVD command, resume options, and accuracy controls. ### Merge LoRA Adapter diff --git a/docs/utilities/lora.md b/docs/utilities/lora.md index a25083da8e..9125256ee7 100644 --- a/docs/utilities/lora.md +++ b/docs/utilities/lora.md @@ -43,7 +43,7 @@ Important options: - `--exact-tensor-pattern`: repeatable regex for a matrix that should remain an exact dense delta. - `--work-dir`, `--resume`: resume an interrupted streaming extraction. -The adapter retains changes that do not fit a low-rank product: `.diff` and `.diff_b` hold exact additive weight/bias deltas, while `.set_weight` holds a parameter absent from the base checkpoint, such as a VSA compression gate. Bit-identical parameters are omitted. The extractor writes an adjacent `*.report.json` with tensor counts, settings, and reconstruction residuals. +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). diff --git a/fastvideo/models/loader/lora_patch.py b/fastvideo/models/loader/lora_patch.py index 40925020ef..8d65cf083d 100644 --- a/fastvideo/models/loader/lora_patch.py +++ b/fastvideo/models/loader/lora_patch.py @@ -9,13 +9,14 @@ Two payload kinds cover the gap, named after the convention ComfyUI's loader already reads so one file works in both places: -``.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,11 +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"} +# 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": ""} # Recognized elsewhere in an adapter and deliberately not our business: the low-rank # half, which ``LoRAPipeline`` merges through the wrapped-module path. @@ -163,7 +164,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) diff --git a/fastvideo/tests/loader/test_lora_patch.py b/fastvideo/tests/loader/test_lora_patch.py index 55e33a38da..b705ff6bb7 100644 --- a/fastvideo/tests/loader/test_lora_patch.py +++ b/fastvideo/tests/loader/test_lora_patch.py @@ -48,6 +48,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 +84,23 @@ 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_param_names_mapping_is_applied_to_dense_keys(tmp_path): @@ -134,6 +146,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 +187,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..3f25a96289 100644 --- a/fastvideo/tests/lora_extraction/test_lora_extraction.py +++ b/fastvideo/tests/lora_extraction/test_lora_extraction.py @@ -2,6 +2,9 @@ import sys from pathlib import Path +import pytest +import torch + # Add scripts/lora_extraction to path for imports repo_root = Path(__file__).parents[3] lora_scripts = repo_root / "scripts" / "lora_extraction" @@ -13,23 +16,38 @@ from verify_lora import main as verify_lora_main -def test_lora_extraction_pipeline(): - """Test end-to-end LoRA extraction workflow.""" +@pytest.mark.parametrize( + "extraction_device", + [ + pytest.param("cpu", id="cpu"), + pytest.param( + "cuda:0", + id="gpu", + marks=pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is unavailable"), + ), + ], +) +def test_lora_extraction_pipeline(extraction_device: str): + """Test the existing Wan2.2 extraction workflow on CPU and GPU.""" import tempfile # 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" + device_name = extraction_device.replace(":", "-") + adapter_path = tmpdir_path / f"adapter_r16_{device_name}.safetensors" + merged_dir = tmpdir_path / f"merged_r16_{device_name}" # 1. Extract rank-16 adapter - print("\nExtracting rank-16 adapter") + print(f"\nExtracting rank-16 adapter on {extraction_device}") extract_lora_adapter( base="Wan-AI/Wan2.2-TI2V-5B-Diffusers", ft="FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers", out=str(adapter_path), rank=16, + load_mode="indexed", + device=extraction_device, + svd_method="exact", ) assert adapter_path.exists(), "Adapter file was not created" diff --git a/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py b/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py index 62e8d72a83..3a484a8a0d 100644 --- a/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py +++ b/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py @@ -49,6 +49,75 @@ def _toy_states() -> tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]: 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" @@ -94,6 +163,34 @@ def test_streaming_extraction_emits_lora_diff_and_replacement(tmp_path: Path) -> assert handle.metadata()["factor_dtype"] == "float16" +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"]) + + def test_randomized_extraction_is_seeded_and_reports_residual(tmp_path: Path) -> None: base, finetuned = _toy_states() base_dir = tmp_path / "base" @@ -131,6 +228,51 @@ def test_randomized_extraction_is_seeded_and_reports_residual(tmp_path: Path) -> 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) diff --git a/scripts/lora_extraction/README.md b/scripts/lora_extraction/README.md index f5863a3723..78b06dc455 100644 --- a/scripts/lora_extraction/README.md +++ b/scripts/lora_extraction/README.md @@ -47,8 +47,8 @@ Important options: - `--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` and `.diff_b` payloads. -- `--replacement-dtype`: storage dtype for fine-tuned-only `.set_weight` parameters. +- `--dense-dtype`: storage dtype for exact `.diff`, `.diff_b`, and `.diff_param` payloads. +- `--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. - `--work-dir`, `--resume`: resume a partially completed streaming extraction. @@ -56,7 +56,9 @@ Fine-tuned parameters that cannot or should not be factorized are retained autom - 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. diff --git a/scripts/lora_extraction/extract_lora.py b/scripts/lora_extraction/extract_lora.py index 4200915191..f815884be5 100644 --- a/scripts/lora_extraction/extract_lora.py +++ b/scripts/lora_extraction/extract_lora.py @@ -3,8 +3,8 @@ The extractor supports ordinary low-rank matrix deltas as well as parameters that cannot be represented by a LoRA product: -* ``.diff`` / ``.diff_b`` store exact additive deltas. -* ``.set_weight`` stores a parameter absent from the base checkpoint. +* ``.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 @@ -45,7 +45,9 @@ FORMAT_VERSION = "fastvideo-lora-v2" DIFF_SUFFIX = ".diff" DIFF_BIAS_SUFFIX = ".diff_b" +DIFF_PARAM_SUFFIX = ".diff_param" SET_WEIGHT_SUFFIX = ".set_weight" +SET_PARAM_SUFFIX = ".set_param" _DTYPE_MAP = { "float32": torch.float32, @@ -321,12 +323,12 @@ def is_extractable_weight(key: str) -> bool: return not any(fragment in lowered for fragment in ("norm", "bias", "embedding")) -def dense_payload_key(param_name: str) -> str | None: +def dense_payload_key(param_name: str) -> str: if param_name.endswith(".weight"): return param_name.removesuffix(".weight") + DIFF_SUFFIX if param_name.endswith(".bias"): return param_name.removesuffix(".bias") + DIFF_BIAS_SUFFIX - return None + return param_name + DIFF_PARAM_SUFFIX def build_dense_payload( @@ -343,8 +345,9 @@ def build_dense_payload( finetuned = ft_sd[key].detach().cpu() base = base_sd.get(key) if base is None: - if key.endswith(".weight"): - payload[key.removesuffix(".weight") + SET_WEIGHT_SUFFIX] = finetuned.contiguous() + output_key = (key.removesuffix(".weight") + SET_WEIGHT_SUFFIX + if key.endswith(".weight") else key + SET_PARAM_SUFFIX) + payload[output_key] = finetuned.contiguous() continue if base.shape != finetuned.shape or torch.equal(base.cpu(), finetuned): continue @@ -488,9 +491,8 @@ def _extract_layers( } _atomic_json_dump(manifest, manifest_path) continue - if not key.endswith(".weight"): - raise ValueError(f"Fine-tuned-only parameter has no supported replacement suffix: {key}") - output_key = key.removesuffix(".weight") + SET_WEIGHT_SUFFIX + output_key = (key.removesuffix(".weight") + SET_WEIGHT_SUFFIX + if key.endswith(".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] = { @@ -561,8 +563,6 @@ def _extract_layers( del lora_a, lora_b, singular_values, payload else: output_key = dense_payload_key(key) - if output_key is None: - raise ValueError(f"Changed parameter has no supported dense suffix: {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) @@ -754,7 +754,7 @@ def extract_lora_adapter( "diff_tensors": str(counts.get("diff", 0)), "set_weight_tensors": str(counts.get("set_weight", 0)), "dropped_unchanged": str(counts.get("unchanged", 0)), - "application": "W = W_base + lora_B @ lora_A; then .diff/.diff_b added and .set_weight assigned", + "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) From aa57512bd893c089fa181b8acb784d5818ac98c9 Mon Sep 17 00:00:00 2001 From: shaoxiongduan Date: Sat, 29 Aug 2026 14:54:22 +0000 Subject: [PATCH 3/6] [bugfix]: preserve runtime-unsupported LoRA deltas --- docs/training/finetune.md | 7 +- docs/utilities/lora.md | 9 +- .../lora_extraction/test_lora_extraction.py | 122 ++++++++++-------- scripts/lora_extraction/README.md | 9 +- 4 files changed, 88 insertions(+), 59 deletions(-) diff --git a/docs/training/finetune.md b/docs/training/finetune.md index fee384088f..b231d3c397 100644 --- a/docs/training/finetune.md +++ b/docs/training/finetune.md @@ -121,7 +121,9 @@ 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 | @@ -136,8 +138,9 @@ python scripts/lora_extraction/extract_lora.py \ | `--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 | +| `--exact-tensor-pattern` | Repeatable regex for matrices the target runtime cannot load as LoRA factors | -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`. See [LoRA Extraction and Merging](../utilities/lora.md) for the GPU/randomized-SVD command, resume options, and accuracy controls. +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 [LoRA Extraction and Merging](../utilities/lora.md) for the GPU/randomized-SVD command, resume options, and accuracy controls. ### Merge LoRA Adapter diff --git a/docs/utilities/lora.md b/docs/utilities/lora.md index 9125256ee7..ac12aa4465 100644 --- a/docs/utilities/lora.md +++ b/docs/utilities/lora.md @@ -9,9 +9,16 @@ 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$' ``` +The extractor is runtime-agnostic by default 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. 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 diff --git a/fastvideo/tests/lora_extraction/test_lora_extraction.py b/fastvideo/tests/lora_extraction/test_lora_extraction.py index 3f25a96289..17b1b35c3c 100644 --- a/fastvideo/tests/lora_extraction/test_lora_extraction.py +++ b/fastvideo/tests/lora_extraction/test_lora_extraction.py @@ -1,80 +1,92 @@ -"""Test LoRA extraction, merging, and verification pipeline.""" -import sys +"""Test extraction through the real FastVideo LoRA loading path.""" from pathlib import Path +import sys +import tempfile import pytest 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 extract_lora_adapter # noqa: E402 -@pytest.mark.parametrize( - "extraction_device", - [ - pytest.param("cpu", id="cpu"), - pytest.param( - "cuda:0", - id="gpu", - marks=pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is unavailable"), - ), - ], -) -def test_lora_extraction_pipeline(extraction_device: str): - """Test the existing Wan2.2 extraction workflow on CPU and GPU.""" - import tempfile +def _collect_lora_application(worker) -> 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) + return { + "adapted": adapted, + "pipeline": type(pipeline).__name__, + "unmatched": unmatched, + } - # Use temp directory for outputs to avoid polluting repo - with tempfile.TemporaryDirectory() as tmpdir: - tmpdir_path = Path(tmpdir) - device_name = extraction_device.replace(":", "-") - adapter_path = tmpdir_path / f"adapter_r16_{device_name}.safetensors" - merged_dir = tmpdir_path / f"merged_r16_{device_name}" - # 1. Extract rank-16 adapter - print(f"\nExtracting rank-16 adapter on {extraction_device}") +@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" + with tempfile.TemporaryDirectory() as tmpdir: + adapter_path = Path(tmpdir) / "adapter_r16.safetensors" extract_lora_adapter( - base="Wan-AI/Wan2.2-TI2V-5B-Diffusers", + base=base, ft="FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers", out=str(adapter_path), rank=16, load_mode="indexed", - device=extraction_device, + 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 + generator = VideoGenerator.from_config( + GeneratorConfig( + model_path=base, + 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) finally: - sys.argv = old_argv + generator.shutdown() - print("\nLoRA extraction pipeline test PASSED") + assert summaries == [{ + "adapted": 300, + "pipeline": "WanDMDPipeline", + "unmatched": [], + }] diff --git a/scripts/lora_extraction/README.md b/scripts/lora_extraction/README.md index 78b06dc455..4e6d17ec88 100644 --- a/scripts/lora_extraction/README.md +++ b/scripts/lora_extraction/README.md @@ -11,9 +11,16 @@ 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$' ``` +The extractor is runtime-agnostic by default: 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 rather than being discarded. 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 From d0dcee752c492fc8b868edc1aa877dd3c59f4bda Mon Sep 17 00:00:00 2001 From: shaoxiongduan Date: Sun, 30 Aug 2026 03:15:56 +0000 Subject: [PATCH 4/6] [bugfix]: guard scratch cleanup, unmatched patterns, and dense-tensor merges Review follow-ups on the streaming extractor. - extract_lora removed the whole --work-dir tree, so pointing it at a directory that held anything else destroyed those files. Scratch now lives in a dedicated subdirectory and cleanup only removes the manifest and tensor shards this script writes. Assembly is manifest-driven, so dropping the manifest is what makes a rerun start clean. - An --exact-tensor-pattern that matched nothing silently rank-truncated the tensors it was meant to keep exact, and the doubled backslashes in the MiniMax-H3 README command did exactly that. Patterns are now validated against the checkpoint keys before any SVD runs, and the README uses single-escaped dots. - merge_lora only understood the lora_A/lora_B half of an adapter, so every .diff / .diff_b / .diff_param / .set_weight / .set_param tensor was dropped with no warning, including the two the docs now tell users to extract. It applies them and reports anything it still cannot place. - A stale work dir made a fresh (non-resume) rerun fail with a resume error; the config check is now gated on --resume. - --out with a non-.safetensors suffix wrote a safetensors file under that name; it is rejected instead. Tests: fastvideo/tests/lora_extraction/ -> 17 passed, 1 failed. The failure is test_lora_extraction_pipeline, which needs an HF download and fails identically on the parent commit. --- docs/utilities/lora.md | 3 +- .../tests/lora_extraction/test_merge_lora.py | 84 ++++++++++++++++ .../test_streaming_lora_extraction.py | 96 +++++++++++++++++++ scripts/lora_extraction/README.md | 11 ++- scripts/lora_extraction/extract_lora.py | 57 ++++++++--- scripts/lora_extraction/merge_lora.py | 71 +++++++++++++- 6 files changed, 304 insertions(+), 18 deletions(-) create mode 100644 fastvideo/tests/lora_extraction/test_merge_lora.py diff --git a/docs/utilities/lora.md b/docs/utilities/lora.md index ac12aa4465..d615cdbf69 100644 --- a/docs/utilities/lora.md +++ b/docs/utilities/lora.md @@ -48,7 +48,8 @@ Important options: - `--randomized-q`, `--niter`, `--seed`: randomized SVD accuracy and reproducibility. - `--factor-dtype`, `--dense-dtype`, `--replacement-dtype`: adapter storage precision. - `--exact-tensor-pattern`: repeatable regex for a matrix that should remain an exact dense delta. -- `--work-dir`, `--resume`: resume an interrupted streaming extraction. +- `--work-dir`, `--resume`: resume an interrupted streaming extraction. Scratch is written to a + `fastvideo-lora-extract/` subdirectory of `--work-dir`, and only that subdirectory is cleaned up. 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. diff --git a/fastvideo/tests/lora_extraction/test_merge_lora.py b/fastvideo/tests/lora_extraction/test_merge_lora.py new file mode 100644 index 0000000000..da8427c3dd --- /dev/null +++ b/fastvideo/tests/lora_extraction/test_merge_lora.py @@ -0,0 +1,84 @@ +"""Coverage for merging the non-factorized half of an adapter into base weights.""" + +from __future__ import annotations + +import logging +from pathlib import Path +import sys + +import torch + +_REPO_ROOT = Path(__file__).parents[3] +sys.path.insert(0, str(_REPO_ROOT / "scripts" / "lora_extraction")) + +import merge_lora # noqa: E402 + + +def _base_state() -> dict[str, torch.Tensor]: + generator = torch.Generator().manual_seed(11) + return { + "blocks.0.linear.weight": torch.randn(9, 7, generator=generator), + "audio_proj_in.weight": torch.randn(8, 6, generator=generator), + "time_embedder.linear.bias": torch.randn(8, generator=generator), + "scale_param": torch.randn(4, generator=generator), + } + + +def _adapter(base: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: + generator = torch.Generator().manual_seed(12) + return { + "blocks.0.linear.lora_A.weight": torch.randn(2, 7, generator=generator), + "blocks.0.linear.lora_B.weight": torch.randn(9, 2, generator=generator), + "audio_proj_in.diff": torch.randn(8, 6, generator=generator), + "time_embedder.linear.diff_b": torch.randn(8, generator=generator), + "scale_param.diff_param": torch.randn(4, generator=generator), + "new_proj.set_weight": torch.randn(3, 3, generator=generator), + } + + +def test_group_dense_keys_splits_additive_and_replacement() -> None: + additive, replacement, unrecognized = merge_lora.group_dense_keys(_adapter(_base_state())) + + assert set(additive) == {"audio_proj_in.weight", "time_embedder.linear.bias", "scale_param"} + assert set(replacement) == {"new_proj.weight"} + assert unrecognized == [] + + +def test_group_dense_keys_reports_unrecognized_suffixes() -> None: + _, _, unrecognized = merge_lora.group_dense_keys({"mystery.tensor": torch.zeros(2)}) + assert unrecognized == ["mystery.tensor"] + + +def test_merge_applies_dense_tensors_alongside_lora() -> None: + """--exact-tensor-pattern keeps tensors as dense deltas; the merge must still apply them.""" + base = _base_state() + adapter = _adapter(base) + + merged = merge_lora.merge_lora_into_base(base, adapter) + + expected_lora = base["blocks.0.linear.weight"] + adapter["blocks.0.linear.lora_B.weight"] @ adapter[ + "blocks.0.linear.lora_A.weight"] + assert torch.allclose(merged["blocks.0.linear.weight"], expected_lora, atol=1e-6) + assert torch.allclose(merged["audio_proj_in.weight"], + base["audio_proj_in.weight"] + adapter["audio_proj_in.diff"], + atol=1e-6) + assert torch.allclose(merged["time_embedder.linear.bias"], + base["time_embedder.linear.bias"] + adapter["time_embedder.linear.diff_b"], + atol=1e-6) + assert torch.allclose(merged["scale_param"], base["scale_param"] + adapter["scale_param.diff_param"], atol=1e-6) + assert torch.equal(merged["new_proj.weight"], adapter["new_proj.set_weight"]) + # the caller's state dict must not be mutated + assert torch.equal(base["audio_proj_in.weight"], _base_state()["audio_proj_in.weight"]) + + +def test_merge_warns_instead_of_silently_dropping_unknown_keys(caplog) -> None: + base = _base_state() + adapter = {"mystery.tensor": torch.zeros(2), "shape_mismatch.diff": torch.zeros(1, 1)} + + with caplog.at_level(logging.WARNING, logger=merge_lora.LOG.name): + merge_lora.merge_lora_into_base(base, adapter) + + messages = " ".join(record.getMessage() for record in caplog.records) + assert "unrecognized suffix" in messages + assert "mystery.tensor" in messages + assert "shape_mismatch.weight" in messages diff --git a/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py b/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py index 3a484a8a0d..133fa5ffe5 100644 --- a/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py +++ b/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py @@ -320,3 +320,99 @@ def fake_snapshot_download(**kwargs: object) -> str: "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 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 (work_dir / extract_lora.WORK_SUBDIR / "manifest.json").is_file() + + +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) + + +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 4e6d17ec88..7c5befc284 100644 --- a/scripts/lora_extraction/README.md +++ b/scripts/lora_extraction/README.md @@ -37,10 +37,10 @@ python extract_lora.py \ --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\\.' + --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. @@ -57,7 +57,8 @@ Important options: - `--dense-dtype`: storage dtype for exact `.diff`, `.diff_b`, and `.diff_param` payloads. - `--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. -- `--work-dir`, `--resume`: resume a partially completed streaming extraction. +- `--work-dir`, `--resume`: resume a partially completed streaming extraction. Scratch is written to a + `fastvideo-lora-extract/` subdirectory of `--work-dir`, and only that subdirectory is cleaned up. Fine-tuned parameters that cannot or should not be factorized are retained automatically: diff --git a/scripts/lora_extraction/extract_lora.py b/scripts/lora_extraction/extract_lora.py index f815884be5..6a80ab592f 100644 --- a/scripts/lora_extraction/extract_lora.py +++ b/scripts/lora_extraction/extract_lora.py @@ -14,8 +14,8 @@ from __future__ import annotations import argparse -from collections.abc import Iterator, Sequence -from contextlib import ExitStack, contextmanager +from collections.abc import Iterable, Iterator, Sequence +from contextlib import ExitStack, contextmanager, suppress from dataclasses import asdict, dataclass import hashlib import json @@ -43,6 +43,10 @@ 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" + DIFF_SUFFIX = ".diff" DIFF_BIAS_SUFFIX = ".diff_b" DIFF_PARAM_SUFFIX = ".diff_param" @@ -369,6 +373,24 @@ def _compile_patterns(patterns: Sequence[str]) -> tuple[re.Pattern[str], ...]: return tuple(re.compile(pattern) for pattern in patterns) +def _validate_exact_patterns(patterns: Sequence[str], keys: Iterable[str]) -> None: + """Reject a pattern that matches no tensor rather than silently factorizing it anyway. + + 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. + """ + 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, ...], @@ -421,20 +443,27 @@ def _factorize_delta( return lora_a, lora_b, singular_values, method_description +def _clear_work_dir(work_dir: Path) -> None: + """Remove the scratch artifacts this script writes, never the directory itself. + + 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 manifest_path.is_file(): + 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}") - if not resume: - shutil.rmtree(work_dir) - else: - return manifest_path, manifest - elif work_dir.exists() and not resume: - shutil.rmtree(work_dir) + 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": {}} @@ -702,8 +731,11 @@ def extract_lora_adapter( raise ValueError("randomized_q must be positive") 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}") if work_dir is not None: - effective_work_dir = Path(work_dir).expanduser() + # --work-dir names where to put scratch, not a directory to take over. + effective_work_dir = Path(work_dir).expanduser() / WORK_SUBDIR elif checkpoint is not None: effective_work_dir = Path(checkpoint).expanduser().with_suffix(".work") else: @@ -711,6 +743,7 @@ def extract_lora_adapter( with _open_readers(base, ft, base_revision, ft_revision, load_mode) as (base_reader, finetuned_reader): _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, @@ -764,7 +797,9 @@ def extract_lora_adapter( 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: - shutil.rmtree(effective_work_dir) + _clear_work_dir(effective_work_dir) + with suppress(OSError): + effective_work_dir.rmdir() return out_path diff --git a/scripts/lora_extraction/merge_lora.py b/scripts/lora_extraction/merge_lora.py index f06a45618d..eda8aef96d 100644 --- a/scripts/lora_extraction/merge_lora.py +++ b/scripts/lora_extraction/merge_lora.py @@ -95,6 +95,65 @@ def load_adapter(adapter_path: str) -> dict: return fix_adapter_naming(adapter) +# Suffix -> the parameter suffix it targets, mirroring fastvideo.models.loader.lora_patch. +# An empty target keeps the full name, for standalone nn.Parameters. +ADDITIVE_SUFFIXES: dict[str, str] = {".diff_param": "", ".diff_b": ".bias", ".diff": ".weight"} +REPLACEMENT_SUFFIXES: dict[str, str] = {".set_weight": ".weight", ".set_param": ""} +LORA_SUFFIXES = (".lora_A.weight", ".lora_B.weight", ".lora_rank", ".lora_alpha") + + +def group_dense_keys(adapter: dict) -> tuple[dict, dict, list]: + """Split the non-factorized half of an adapter into additive and replacement params. + + ``--exact-tensor-pattern`` keeps selected matrices as exact dense deltas instead of + LoRA factors, so an adapter merged without these is missing those tensors entirely. + """ + additive: dict = {} + replacement: dict = {} + unrecognized: list = [] + + for key, tensor in adapter.items(): + if key.endswith(LORA_SUFFIXES): + continue + for suffix, param_suffix in ADDITIVE_SUFFIXES.items(): + if key.endswith(suffix): + additive[key.removesuffix(suffix) + param_suffix] = tensor + break + else: + for suffix, param_suffix in REPLACEMENT_SUFFIXES.items(): + if key.endswith(suffix): + replacement[key.removesuffix(suffix) + param_suffix] = tensor + break + else: + unrecognized.append(key) + + return additive, replacement, unrecognized + + +def merge_dense_into_base(merged_sd: dict, additive: dict, replacement: dict) -> tuple[int, list]: + """Apply exact dense deltas and replacement parameters in place.""" + merged_count = 0 + skipped: list = [] + + for param_name, tensor in additive.items(): + target = merged_sd.get(param_name) + if target is None or tuple(target.shape) != tuple(tensor.shape): + skipped.append(param_name) + continue + merged_sd[param_name] = (target.to(torch.float32) + tensor.to(torch.float32)).to(target.dtype) + merged_count += 1 + + for param_name, tensor in replacement.items(): + target = merged_sd.get(param_name) + if target is not None and tuple(target.shape) != tuple(tensor.shape): + skipped.append(param_name) + continue + merged_sd[param_name] = tensor.to(target.dtype) if target is not None else tensor + merged_count += 1 + + return merged_count, skipped + + def group_adapter_keys(adapter: dict) -> dict: grouped = defaultdict(dict) @@ -214,7 +273,17 @@ def merge_lora_into_base(base_sd: dict, adapter: dict) -> dict: merged_sd[weight_key] = merged_weight.to(base_sd[weight_key].dtype) merged_count += 1 - LOG.info(f"Merged {merged_count} layers, skipped {skipped_count}") + additive, replacement, unrecognized = group_dense_keys(adapter) + dense_merged, dense_skipped = merge_dense_into_base(merged_sd, additive, replacement) + + LOG.info(f"Merged {merged_count} LoRA layers, skipped {skipped_count}") + LOG.info(f"Merged {dense_merged} dense tensors ({len(additive)} additive, {len(replacement)} replacement)") + if dense_skipped: + LOG.warning(f"Skipped {len(dense_skipped)} dense tensors absent or mismatched in the base model: " + f"{dense_skipped[:5]}") + if unrecognized: + LOG.warning(f"Ignored {len(unrecognized)} adapter keys with an unrecognized suffix: {unrecognized[:5]}") + return merged_sd From 3906bd8451ed9e8f7772efa4b2b2920ea6825fb5 Mon Sep 17 00:00:00 2001 From: shaoxiongduan Date: Mon, 31 Aug 2026 02:59:58 +0000 Subject: [PATCH 5/6] [bugfix]: harden resumable extraction and adapter merging --- .buildkite/scripts/lanes/lora_extraction.sh | 2 +- docs/training/finetune.md | 2 +- docs/utilities/lora.md | 9 +- fastvideo/models/loader/fsdp_load.py | 2 + fastvideo/models/loader/lora_patch.py | 34 +++- fastvideo/tests/loader/test_lora_patch.py | 14 +- .../tests/lora_extraction/test_merge_lora.py | 89 +++++++++- .../test_streaming_lora_extraction.py | 108 +++++++++++- scripts/lora_extraction/README.md | 9 +- scripts/lora_extraction/extract_lora.py | 93 ++++++++-- scripts/lora_extraction/merge_lora.py | 165 +++++++++++++----- 11 files changed, 457 insertions(+), 70 deletions(-) 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 b231d3c397..4544257272 100644 --- a/docs/training/finetune.md +++ b/docs/training/finetune.md @@ -137,7 +137,7 @@ python scripts/lora_extraction/extract_lora.py \ | `--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 | +| `--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 | 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 [LoRA Extraction and Merging](../utilities/lora.md) for the GPU/randomized-SVD command, resume options, and accuracy controls. diff --git a/docs/utilities/lora.md b/docs/utilities/lora.md index d615cdbf69..248fe342d8 100644 --- a/docs/utilities/lora.md +++ b/docs/utilities/lora.md @@ -46,10 +46,12 @@ Important options: - `--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. +- `--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. -- `--work-dir`, `--resume`: resume an interrupted streaming extraction. Scratch is written to a - `fastvideo-lora-extract/` subdirectory of `--work-dir`, and only that subdirectory is cleaned up. +- `--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. @@ -71,6 +73,7 @@ python scripts/lora_extraction/merge_lora.py \ - `--adapter`: LoRA adapter file (.safetensors) - `--ft`: Fine-tuned model (for configuration) - `--output`: Output directory +- `--allow-unmatched`: Allow an output even when adapter keys cannot be applied (strict matching is the default) ## Validate Quality (Optional) diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index 05d8222cad..d2b11a613c 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -241,9 +241,11 @@ 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) + lora_param_names_mapping_fn = get_param_names_mapping(model.lora_param_names_mapping) 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: diff --git a/fastvideo/models/loader/lora_patch.py b/fastvideo/models/loader/lora_patch.py index 8d65cf083d..7d0a3b95dc 100644 --- a/fastvideo/models/loader/lora_patch.py +++ b/fastvideo/models/loader/lora_patch.py @@ -127,14 +127,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 @@ -151,7 +151,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 @@ -235,6 +235,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")``. @@ -246,20 +247,37 @@ def _resolve( return None 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/tests/loader/test_lora_patch.py b/fastvideo/tests/loader/test_lora_patch.py index b705ff6bb7..6cd23b6260 100644 --- a/fastvideo/tests/loader/test_lora_patch.py +++ b/fastvideo/tests/loader/test_lora_patch.py @@ -9,8 +9,9 @@ import torch from safetensors.torch import save_file +from fastvideo.configs.models.dits.wanvideo import WanVideoConfig 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"): @@ -114,6 +115,17 @@ def mapping(name): assert set(patch._additive) == {"blocks.0.ff.fc_in.weight"} +def test_official_wan_dense_key_uses_lora_then_checkpoint_mapping(tmp_path): + path = write_adapter(tmp_path, {"blocks.0.self_attn.q.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_q.weight"} + + 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)}) diff --git a/fastvideo/tests/lora_extraction/test_merge_lora.py b/fastvideo/tests/lora_extraction/test_merge_lora.py index da8427c3dd..f7219eb0df 100644 --- a/fastvideo/tests/lora_extraction/test_merge_lora.py +++ b/fastvideo/tests/lora_extraction/test_merge_lora.py @@ -2,15 +2,22 @@ from __future__ import annotations +import json import logging from pathlib import Path import sys +import pytest +from safetensors.torch import save_file import torch +from fastvideo.configs.models.dits.wanvideo import WanVideoConfig +from fastvideo.models.loader.utils import get_param_names_mapping + _REPO_ROOT = Path(__file__).parents[3] sys.path.insert(0, str(_REPO_ROOT / "scripts" / "lora_extraction")) +import extract_lora # noqa: E402 import merge_lora # noqa: E402 @@ -71,14 +78,92 @@ def test_merge_applies_dense_tensors_alongside_lora() -> None: assert torch.equal(base["audio_proj_in.weight"], _base_state()["audio_proj_in.weight"]) +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 test_indexed_wan_extraction_merges_into_fastvideo_namespace(tmp_path: Path) -> None: + generator = torch.Generator().manual_seed(21) + base_hf = { + "blocks.0.attn1.to_q.weight": torch.randn(5, 4, generator=generator), + "condition_embedder.time_proj.weight": torch.randn(3, 4, generator=generator), + } + finetuned_hf = {name: tensor.clone() for name, tensor in base_hf.items()} + finetuned_hf["blocks.0.attn1.to_q.weight"] += torch.randn( + 5, 2, generator=generator) @ torch.randn(2, 4, generator=generator) + finetuned_hf["condition_embedder.time_proj.weight"] += 0.25 + base_dir = tmp_path / "base" + finetuned_dir = tmp_path / "finetuned" + output = tmp_path / "adapter.safetensors" + _write_transformer(base_dir, base_hf) + _write_transformer(finetuned_dir, finetuned_hf) + 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", + exact_tensor_patterns=(r"^condition_embedder\.time_proj\.weight$", ), + ) + + base_custom = { + "blocks.0.to_q.weight": base_hf["blocks.0.attn1.to_q.weight"], + "condition_embedder.time_modulation.linear.weight": base_hf["condition_embedder.time_proj.weight"], + } + config = WanVideoConfig() + merged = merge_lora.merge_lora_into_base( + base_custom, + merge_lora.load_adapter(str(output)), + lora_param_names_mapping=get_param_names_mapping(config.lora_param_names_mapping), + param_names_mapping=get_param_names_mapping(config.param_names_mapping), + ) + + torch.testing.assert_close(merged["blocks.0.to_q.weight"], finetuned_hf["blocks.0.attn1.to_q.weight"]) + torch.testing.assert_close(merged["condition_embedder.time_modulation.linear.weight"], + finetuned_hf["condition_embedder.time_proj.weight"]) + + +def test_merge_accumulates_multiple_sources_into_a_fused_parameter() -> None: + base = {"fused.weight": torch.zeros(6, 4)} + adapter = { + "q.lora_A.weight": torch.ones(1, 4), + "q.lora_B.weight": torch.ones(3, 1), + "k.lora_A.weight": torch.full((1, 4), 2.0), + "k.lora_B.weight": torch.ones(3, 1), + } + + def mapping(name: str): + return "fused.weight", 0 if name.startswith("q.") else 1, 2 + + merged = merge_lora.merge_lora_into_base(base, adapter, param_names_mapping=mapping) + torch.testing.assert_close(merged["fused.weight"][:3], torch.ones(3, 4)) + torch.testing.assert_close(merged["fused.weight"][3:], torch.full((3, 4), 2.0)) + + +def test_merge_is_strict_about_unapplied_keys() -> None: + with pytest.raises(ValueError, match="unapplied"): + merge_lora.merge_lora_into_base(_base_state(), {"missing.diff": torch.zeros(2)}) + + def test_merge_warns_instead_of_silently_dropping_unknown_keys(caplog) -> None: base = _base_state() adapter = {"mystery.tensor": torch.zeros(2), "shape_mismatch.diff": torch.zeros(1, 1)} with caplog.at_level(logging.WARNING, logger=merge_lora.LOG.name): - merge_lora.merge_lora_into_base(base, adapter) + merge_lora.merge_lora_into_base(base, adapter, strict=False) messages = " ".join(record.getMessage() for record in caplog.records) - assert "unrecognized suffix" in messages + assert "unrecognized adapter key" in messages assert "mystery.tensor" in messages assert "shape_mismatch.weight" in messages diff --git a/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py b/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py index 133fa5ffe5..f0aaa6c895 100644 --- a/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py +++ b/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py @@ -163,6 +163,41 @@ def test_streaming_extraction_emits_lora_diff_and_replacement(tmp_path: Path) -> 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 = { @@ -343,6 +378,10 @@ def _extract(base_dir: Path, finetuned_dir: Path, output: Path, **overrides: obj 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) @@ -368,7 +407,20 @@ def test_work_dir_scratch_is_confined_to_a_subdirectory(tmp_path: Path) -> None: work_dir=str(work_dir), keep_work_dir=True) - assert (work_dir / extract_lora.WORK_SUBDIR / "manifest.json").is_file() + 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: @@ -410,6 +462,60 @@ def test_resume_still_rejects_a_mismatched_config(tmp_path: Path) -> None: _extract(base_dir, finetuned_dir, output, rank=4, resume=True) +def test_resume_rejects_changed_checkpoint_contents(tmp_path: Path) -> None: + base_dir, finetuned_dir = _toy_checkpoints(tmp_path) + output = tmp_path / "adapter.safetensors" + _extract(base_dir, finetuned_dir, output, keep_work_dir=True) + + _, changed = _toy_states() + changed["blocks.0.linear.weight"] += 1.0 + shard = finetuned_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) + + +def test_local_paths_reject_hugging_face_revisions(tmp_path: Path) -> 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", + base_revision="abc123") + + +def test_pipeline_mode_rejects_revisions_before_loading(tmp_path: Path) -> 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", + base_revision="abc123", + ) + + +def test_auto_mode_does_not_drop_revision_during_fallback(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 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", + base_revision="abc123", + ) + + 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) diff --git a/scripts/lora_extraction/README.md b/scripts/lora_extraction/README.md index 7c5befc284..91553868b9 100644 --- a/scripts/lora_extraction/README.md +++ b/scripts/lora_extraction/README.md @@ -54,11 +54,13 @@ Important options: - `--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. +- `--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. -- `--work-dir`, `--resume`: resume a partially completed streaming extraction. Scratch is written to a - `fastvideo-lora-extract/` subdirectory of `--work-dir`, and only that subdirectory is cleaned up. +- `--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: @@ -87,6 +89,7 @@ python merge_lora.py \ - `--adapter`: LoRA adapter file (.safetensors) - `--ft`: Fine-tuned model (for configuration) - `--output`: Output directory +- `--allow-unmatched`: Opt in to writing an output when adapter keys cannot be applied; strict matching is the default. ## Validate Quality (Optional) diff --git a/scripts/lora_extraction/extract_lora.py b/scripts/lora_extraction/extract_lora.py index 6a80ab592f..198e14e65c 100644 --- a/scripts/lora_extraction/extract_lora.py +++ b/scripts/lora_extraction/extract_lora.py @@ -69,6 +69,7 @@ class TensorReader(Protocol): """Random access to one checkpoint's transformer tensors.""" source: str + fingerprint: str @property def keys(self) -> set[str]: ... @@ -77,6 +78,8 @@ 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: ... @@ -88,6 +91,7 @@ class DictTensorReader: 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]: @@ -99,6 +103,9 @@ def get_tensor(self, key: str) -> torch.Tensor: 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 @@ -127,11 +134,29 @@ def __init__(self, transformer_dir: Path) -> None: if not self.weight_map: raise ValueError(f"No transformer safetensors found under {transformer_dir}") - self._stack = ExitStack() - self._shards = { - shard: self._stack.enter_context(safe_open(transformer_dir / shard, framework="pt", device="cpu")) - for shard in sorted(set(self.weight_map.values())) - } + 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]: @@ -143,6 +168,10 @@ def get_tensor(self, key: str) -> torch.Tensor: 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 @@ -154,6 +183,8 @@ def __exit__(self, *args: object) -> None: class ExtractionConfig: base_source: str finetuned_source: str + base_fingerprint: str + finetuned_fingerprint: str rank: int full_rank: bool min_delta: float @@ -187,6 +218,24 @@ def _atomic_json_dump(data: Any, path: Path) -> None: 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()] @@ -207,6 +256,8 @@ def _resolve_transformer_dir(model: str, revision: str | None = None) -> Path: """ 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")): @@ -294,6 +345,9 @@ def _open_readers( 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: @@ -302,7 +356,7 @@ def _open_readers( IndexedSafetensorsReader(_resolve_transformer_dir(finetuned, finetuned_revision))) except Exception: stack.close() - if load_mode == "indexed": + if load_mode == "indexed" or revisions_requested: raise LOG.warning("Indexed loading failed; falling back to pipeline loading", exc_info=True) else: @@ -443,6 +497,13 @@ def _factorize_delta( return lora_a, lora_b, singular_values, method_description +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}" + + def _clear_work_dir(work_dir: Path) -> None: """Remove the scratch artifacts this script writes, never the directory itself. @@ -715,7 +776,7 @@ def extract_lora_adapter( niter: int = 4, seed: int = 42, factor_dtype: str = "float32", - dense_dtype: str = "source", + dense_dtype: str = "float32", replacement_dtype: str = "source", exact_tensor_patterns: Sequence[str] = (), work_dir: str | None = None, @@ -733,20 +794,27 @@ def extract_lora_adapter( 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 names where to put scratch, not a directory to take over. - effective_work_dir = Path(work_dir).expanduser() / WORK_SUBDIR + # --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, 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, @@ -764,6 +832,8 @@ def extract_lora_adapter( ) 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(): @@ -800,6 +870,9 @@ def extract_lora_adapter( _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 @@ -824,7 +897,7 @@ def parse_args() -> argparse.Namespace: 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="source") + 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=[], diff --git a/scripts/lora_extraction/merge_lora.py b/scripts/lora_extraction/merge_lora.py index eda8aef96d..33d6e8bcfe 100644 --- a/scripts/lora_extraction/merge_lora.py +++ b/scripts/lora_extraction/merge_lora.py @@ -16,6 +16,8 @@ import logging from pathlib import Path from collections import defaultdict +from collections.abc import Callable +from typing import Any os.environ.setdefault("MASTER_ADDR", "127.0.0.1") os.environ.setdefault("MASTER_PORT", "29500") @@ -130,25 +132,78 @@ def group_dense_keys(adapter: dict) -> tuple[dict, dict, list]: return additive, replacement, unrecognized -def merge_dense_into_base(merged_sd: dict, additive: dict, replacement: dict) -> tuple[int, list]: +ParamMapping = Callable[[str], tuple[str, Any, Any]] + + +def map_adapter_parameter( + name: str, + lora_param_names_mapping: ParamMapping | None, + param_names_mapping: ParamMapping | None, +) -> tuple[str, int | None, int | None]: + """Apply the same official-LoRA -> HF -> FastVideo mapping order as runtime loading.""" + if lora_param_names_mapping is not None: + name, merge_index, total = lora_param_names_mapping(name) + if merge_index is not None: + raise NotImplementedError(f"Adapter-specific mapping unexpectedly fused {name} ({merge_index}/{total})") + if param_names_mapping is None: + return name, None, None + return param_names_mapping(name) + + +def _mapped_slice(target: torch.Tensor, merge_index: int | None, + total: int | None) -> tuple[Any, tuple[int, ...]] | None: + if merge_index is None: + return Ellipsis, tuple(target.shape) + if total is None or total < 1 or target.ndim < 1 or target.shape[0] % total: + return None + chunk = target.shape[0] // total + return slice(merge_index * chunk, (merge_index + 1) * chunk), (chunk, *target.shape[1:]) + + +def merge_dense_into_base( + merged_sd: dict, + additive: dict, + replacement: dict, + lora_param_names_mapping: ParamMapping | None = None, + param_names_mapping: ParamMapping | None = None, +) -> tuple[int, list[str]]: """Apply exact dense deltas and replacement parameters in place.""" merged_count = 0 - skipped: list = [] + skipped: list[str] = [] - for param_name, tensor in additive.items(): + for source_name, tensor in additive.items(): + param_name, merge_index, total = map_adapter_parameter(source_name, lora_param_names_mapping, + param_names_mapping) target = merged_sd.get(param_name) - if target is None or tuple(target.shape) != tuple(tensor.shape): - skipped.append(param_name) + target_slice = _mapped_slice(target, merge_index, total) if target is not None else None + if target is None or target_slice is None or target_slice[1] != tuple(tensor.shape): + skipped.append(f"{source_name} -> {param_name}") continue - merged_sd[param_name] = (target.to(torch.float32) + tensor.to(torch.float32)).to(target.dtype) + index, _ = target_slice + updated = target.to(torch.float32).clone() + updated[index] += tensor.to(torch.float32) + merged_sd[param_name] = updated.to(target.dtype) merged_count += 1 - for param_name, tensor in replacement.items(): + for source_name, tensor in replacement.items(): + param_name, merge_index, total = map_adapter_parameter(source_name, lora_param_names_mapping, + param_names_mapping) target = merged_sd.get(param_name) - if target is not None and tuple(target.shape) != tuple(tensor.shape): - skipped.append(param_name) + if target is None: + if merge_index is not None: + skipped.append(f"{source_name} -> {param_name}") + continue + merged_sd[param_name] = tensor + merged_count += 1 continue - merged_sd[param_name] = tensor.to(target.dtype) if target is not None else tensor + target_slice = _mapped_slice(target, merge_index, total) + if target_slice is None or target_slice[1] != tuple(tensor.shape): + skipped.append(f"{source_name} -> {param_name}") + continue + index, _ = target_slice + updated = target.clone() + updated[index] = tensor.to(target.dtype) + merged_sd[param_name] = updated merged_count += 1 return merged_count, skipped @@ -171,7 +226,7 @@ def group_adapter_keys(adapter: dict) -> dict: return grouped -def get_reverse_param_mapping(base_model_path: str): +def get_reverse_param_mapping(base_model_path: str) -> tuple[dict, ParamMapping, ParamMapping]: LOG.info("Loading base model for parameter mapping") pipeline_cls = get_pipeline_class_for_model(base_model_path) @@ -195,16 +250,16 @@ def get_reverse_param_mapping(base_model_path: str): if transformer is None: raise RuntimeError("Could not find transformer in pipeline") - if hasattr(transformer, "reverse_param_names_mapping"): + param_names_mapping_fn = get_param_names_mapping(transformer.param_names_mapping) + lora_param_names_mapping_fn = get_param_names_mapping(transformer.lora_param_names_mapping) + + if getattr(transformer, "reverse_param_names_mapping", None): reverse_mapping = transformer.reverse_param_names_mapping elif hasattr(transformer, "config") and hasattr(transformer.config, "arch_config"): arch_config = transformer.config.arch_config - if hasattr(arch_config, "reverse_param_names_mapping"): + if getattr(arch_config, "reverse_param_names_mapping", None): reverse_mapping = arch_config.reverse_param_names_mapping else: - param_mapping = arch_config.param_names_mapping - param_names_mapping_fn = get_param_names_mapping(param_mapping) - from diffusers import DiffusionPipeline from huggingface_hub import snapshot_download @@ -229,36 +284,48 @@ def get_reverse_param_mapping(base_model_path: str): del transformer torch.cuda.empty_cache() - return reverse_mapping + return reverse_mapping, param_names_mapping_fn, lora_param_names_mapping_fn -def merge_lora_into_base(base_sd: dict, adapter: dict) -> dict: +def merge_lora_into_base( + base_sd: dict, + adapter: dict, + *, + lora_param_names_mapping: ParamMapping | None = None, + param_names_mapping: ParamMapping | None = None, + strict: bool = True, +) -> dict: LOG.info("Merging LoRA into base weights") adapter_layers = group_adapter_keys(adapter) merged_sd = dict(base_sd) merged_count = 0 - skipped_count = 0 + skipped: list[str] = [] for base_name, parts in adapter_layers.items(): - weight_key = base_name if base_name.endswith(".weight") else base_name + ".weight" - - if weight_key not in base_sd: - skipped_count += 1 + source_weight_key = base_name if base_name.endswith(".weight") else base_name + ".weight" + weight_key, merge_index, total = map_adapter_parameter(source_weight_key, lora_param_names_mapping, + param_names_mapping) + base_tensor = merged_sd.get(weight_key) + if base_tensor is None: + skipped.append(f"{source_weight_key} -> {weight_key} (missing base parameter)") continue if "A" not in parts or "B" not in parts: - skipped_count += 1 + skipped.append(f"{source_weight_key} (incomplete factor pair)") continue lora_A = parts["A"].to(torch.float32) lora_B = parts["B"].to(torch.float32) - base_weight = base_sd[weight_key].to(torch.float32) - - out_dim, in_dim = base_weight.shape + target_slice = _mapped_slice(base_tensor, merge_index, total) + if target_slice is None: + skipped.append(f"{source_weight_key} -> {weight_key} (invalid fused mapping)") + continue + index, expected_shape = target_slice + out_dim, in_dim = expected_shape if lora_B.shape[0] != out_dim or lora_A.shape[1] != in_dim or lora_B.shape[1] != lora_A.shape[0]: - skipped_count += 1 + skipped.append(f"{source_weight_key} -> {weight_key} (factor shape mismatch)") continue delta = lora_B @ lora_A @@ -269,20 +336,27 @@ def merge_lora_into_base(base_sd: dict, adapter: dict) -> dict: if rank != 0 and alpha != rank: delta = delta * (alpha / float(rank)) - merged_weight = base_weight + delta - merged_sd[weight_key] = merged_weight.to(base_sd[weight_key].dtype) + updated = base_tensor.to(torch.float32).clone() + updated[index] += delta + merged_sd[weight_key] = updated.to(base_tensor.dtype) merged_count += 1 additive, replacement, unrecognized = group_dense_keys(adapter) - dense_merged, dense_skipped = merge_dense_into_base(merged_sd, additive, replacement) + dense_merged, dense_skipped = merge_dense_into_base( + merged_sd, + additive, + replacement, + lora_param_names_mapping, + param_names_mapping, + ) - LOG.info(f"Merged {merged_count} LoRA layers, skipped {skipped_count}") - LOG.info(f"Merged {dense_merged} dense tensors ({len(additive)} additive, {len(replacement)} replacement)") - if dense_skipped: - LOG.warning(f"Skipped {len(dense_skipped)} dense tensors absent or mismatched in the base model: " - f"{dense_skipped[:5]}") - if unrecognized: - LOG.warning(f"Ignored {len(unrecognized)} adapter keys with an unrecognized suffix: {unrecognized[:5]}") + LOG.info("Merged %d LoRA layers, skipped %d", merged_count, len(skipped)) + LOG.info("Merged %d dense tensors (%d additive, %d replacement)", dense_merged, len(additive), len(replacement)) + problems = skipped + dense_skipped + [f"unrecognized adapter key {key}" for key in unrecognized] + if problems and strict: + raise ValueError(f"Adapter merge left {len(problems)} keys/layers unapplied: {problems[:5]}") + if problems: + LOG.warning("Adapter merge left %d keys/layers unapplied: %s", len(problems), problems[:5]) return merged_sd @@ -350,6 +424,7 @@ def merge_lora( ft: str, output: str, log_level: str = "INFO", + allow_unmatched: bool = False, ) -> None: """Merge LoRA adapter into base model. @@ -359,6 +434,7 @@ def merge_lora( ft: Finetuned model ID (for config) output: Output directory log_level: Logging level + allow_unmatched: Write output even if recognized adapter payloads cannot be applied """ configure_logging(log_level) @@ -366,14 +442,20 @@ def merge_lora( LOG.info(f"Adapter: {adapter}") LOG.info(f"Output: {output}") - reverse_mapping = get_reverse_param_mapping(base) + reverse_mapping, param_names_mapping_fn, lora_param_names_mapping_fn = get_reverse_param_mapping(base) LOG.info(f"Loading base model: {base}") base_sd = load_transformer_state_dict_from_model(base) LOG.info(f"Loaded { len(base_sd)} parameters") adapter_sd = load_adapter(adapter) - merged_sd = merge_lora_into_base(base_sd, adapter_sd) + merged_sd = merge_lora_into_base( + base_sd, + adapter_sd, + lora_param_names_mapping=lora_param_names_mapping_fn, + param_names_mapping=param_names_mapping_fn, + strict=not allow_unmatched, + ) save_merged_model(merged_sd, base, ft, output, reverse_mapping) LOG.info("Merge complete") @@ -387,6 +469,8 @@ def main(): parser.add_argument("--ft", required=True, help="Finetuned model ID (for config)") parser.add_argument("--output", required=True, help="Output directory") parser.add_argument("--log-level", default="INFO", help="Logging level") + parser.add_argument("--allow-unmatched", action="store_true", + help="Write the merged model even when adapter keys cannot be applied") args = parser.parse_args() merge_lora( @@ -395,6 +479,7 @@ def main(): ft=args.ft, output=args.output, log_level=args.log_level, + allow_unmatched=args.allow_unmatched, ) From c7e527d363950fa6eb93b050ed5308f2a45042d8 Mon Sep 17 00:00:00 2001 From: shaoxiongduan Date: Tue, 1 Sep 2026 09:19:37 +0000 Subject: [PATCH 6/6] [bugfix]: finalize generic LoRA extraction runtime path --- docs/training/finetune.md | 24 ++- docs/utilities/lora.md | 23 +- fastvideo/models/loader/fsdp_load.py | 20 +- fastvideo/models/loader/lora_patch.py | 14 +- fastvideo/pipelines/lora_pipeline.py | 4 +- fastvideo/tests/loader/test_lora_patch.py | 122 ++++++++++- .../lora_extraction/test_lora_extraction.py | 55 ++++- .../tests/lora_extraction/test_merge_lora.py | 169 --------------- .../test_streaming_lora_extraction.py | 98 +++++++-- scripts/lora_extraction/README.md | 32 ++- scripts/lora_extraction/extract_lora.py | 20 +- scripts/lora_extraction/merge_lora.py | 200 ++---------------- 12 files changed, 363 insertions(+), 418 deletions(-) delete mode 100644 fastvideo/tests/lora_extraction/test_merge_lora.py diff --git a/docs/training/finetune.md b/docs/training/finetune.md index 4544257272..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 @@ -128,11 +128,12 @@ python scripts/lora_extraction/extract_lora.py \ | 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 | @@ -140,16 +141,21 @@ python scripts/lora_extraction/extract_lora.py \ | `--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 | -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 [LoRA Extraction and Merging](../utilities/lora.md) for the GPU/randomized-SVD command, resume options, and accuracy controls. +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 LoRA Adapter +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. -Merge an adapter back into a base model: +### 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 248fe342d8..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 @@ -14,10 +14,10 @@ python scripts/lora_extraction/extract_lora.py \ --exact-tensor-pattern '^proj_out\.weight$' ``` -The extractor is runtime-agnostic by default and cannot determine from checkpoint tensors whether the target runtime +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. The Wan patterns above cover its excluded condition embedders -and its unwrapped output projection. +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: @@ -43,6 +43,7 @@ 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. @@ -57,23 +58,29 @@ The adapter retains changes that do not fit a low-rank product: `.diff` and `.di 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). -## Merge Adapter +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. + +## Legacy 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 -- `--allow-unmatched`: Allow an output even when adapter keys cannot be applied (strict matching is the default) ## Validate Quality (Optional) diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index d2b11a613c..f5a6c156e6 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -241,13 +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) - lora_param_names_mapping_fn = get_param_names_mapping(model.lora_param_names_mapping) - dense_lora_patch = DenseLoRAPatch.from_adapter( - lora_path, - param_names_mapping_fn, - lora_param_names_mapping=lora_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 @@ -697,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 7d0a3b95dc..6154cca941 100644 --- a/fastvideo/models/loader/lora_patch.py +++ b/fastvideo/models/loader/lora_patch.py @@ -6,8 +6,8 @@ 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_param`` An exact additive delta for a parameter the base model has. Used where a rank-``r`` @@ -56,12 +56,6 @@ ADDITIVE_SUFFIXES: dict[str, str] = {".diff_param": "", ".diff_b": ".bias", ".diff": ".weight"} REPLACEMENT_SUFFIXES: dict[str, str] = {".set_weight": ".weight", ".set_param": ""} -# 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") - - # 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 # ComfyUI write ``.lora_down.weight`` with a bare ``.alpha``. Normalizing on the way in @@ -243,8 +237,8 @@ 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( 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 6cd23b6260..8fee777474 100644 --- a/fastvideo/tests/loader/test_lora_patch.py +++ b/fastvideo/tests/loader/test_lora_patch.py @@ -10,6 +10,7 @@ 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 @@ -104,6 +105,12 @@ def test_from_adapter_splits_additive_from_replacement(tmp_path): } +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): """An adapter written against the published layout needs no separate rename table.""" path = write_adapter(tmp_path, {"blocks.0.ff.net.0.proj.diff": torch.zeros(4)}) @@ -115,15 +122,124 @@ def mapping(name): assert set(patch._additive) == {"blocks.0.ff.fc_in.weight"} -def test_official_wan_dense_key_uses_lora_then_checkpoint_mapping(tmp_path): - path = write_adapter(tmp_path, {"blocks.0.self_attn.q.diff": torch.zeros(4)}) +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_q.weight"} + 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): diff --git a/fastvideo/tests/lora_extraction/test_lora_extraction.py b/fastvideo/tests/lora_extraction/test_lora_extraction.py index 17b1b35c3c..4ff39980a7 100644 --- a/fastvideo/tests/lora_extraction/test_lora_extraction.py +++ b/fastvideo/tests/lora_extraction/test_lora_extraction.py @@ -1,9 +1,11 @@ """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 @@ -14,10 +16,31 @@ lora_scripts = repo_root / "scripts" / "lora_extraction" sys.path.insert(0, str(lora_scripts)) -from extract_lora import extract_lora_adapter # noqa: E402 +from extract_lora import INDEX_FILENAME, _resolve_transformer_dir, extract_lora_adapter # noqa: E402 -def _collect_lora_application(worker) -> dict[str, object]: +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}") + + +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] @@ -30,8 +53,14 @@ def _collect_lora_application(worker) -> dict[str, object]: 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, } @@ -41,22 +70,36 @@ def _collect_lora_application(worker) -> dict[str, object]: 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=base, - ft="FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers", + 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$"), ) + 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), @@ -81,12 +124,16 @@ def test_lora_extraction_pipeline() -> None: ), )) try: - summaries = generator.executor.collective_rpc(_collect_lora_application) + summaries = generator.executor.collective_rpc( + _collect_lora_application, + args=(dense_param_name, expected_dense_values), + ) finally: generator.shutdown() assert summaries == [{ "adapted": 300, + "dense_matches_finetuned": True, "pipeline": "WanDMDPipeline", "unmatched": [], }] diff --git a/fastvideo/tests/lora_extraction/test_merge_lora.py b/fastvideo/tests/lora_extraction/test_merge_lora.py deleted file mode 100644 index f7219eb0df..0000000000 --- a/fastvideo/tests/lora_extraction/test_merge_lora.py +++ /dev/null @@ -1,169 +0,0 @@ -"""Coverage for merging the non-factorized half of an adapter into base weights.""" - -from __future__ import annotations - -import json -import logging -from pathlib import Path -import sys - -import pytest -from safetensors.torch import save_file -import torch - -from fastvideo.configs.models.dits.wanvideo import WanVideoConfig -from fastvideo.models.loader.utils import get_param_names_mapping - -_REPO_ROOT = Path(__file__).parents[3] -sys.path.insert(0, str(_REPO_ROOT / "scripts" / "lora_extraction")) - -import extract_lora # noqa: E402 -import merge_lora # noqa: E402 - - -def _base_state() -> dict[str, torch.Tensor]: - generator = torch.Generator().manual_seed(11) - return { - "blocks.0.linear.weight": torch.randn(9, 7, generator=generator), - "audio_proj_in.weight": torch.randn(8, 6, generator=generator), - "time_embedder.linear.bias": torch.randn(8, generator=generator), - "scale_param": torch.randn(4, generator=generator), - } - - -def _adapter(base: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: - generator = torch.Generator().manual_seed(12) - return { - "blocks.0.linear.lora_A.weight": torch.randn(2, 7, generator=generator), - "blocks.0.linear.lora_B.weight": torch.randn(9, 2, generator=generator), - "audio_proj_in.diff": torch.randn(8, 6, generator=generator), - "time_embedder.linear.diff_b": torch.randn(8, generator=generator), - "scale_param.diff_param": torch.randn(4, generator=generator), - "new_proj.set_weight": torch.randn(3, 3, generator=generator), - } - - -def test_group_dense_keys_splits_additive_and_replacement() -> None: - additive, replacement, unrecognized = merge_lora.group_dense_keys(_adapter(_base_state())) - - assert set(additive) == {"audio_proj_in.weight", "time_embedder.linear.bias", "scale_param"} - assert set(replacement) == {"new_proj.weight"} - assert unrecognized == [] - - -def test_group_dense_keys_reports_unrecognized_suffixes() -> None: - _, _, unrecognized = merge_lora.group_dense_keys({"mystery.tensor": torch.zeros(2)}) - assert unrecognized == ["mystery.tensor"] - - -def test_merge_applies_dense_tensors_alongside_lora() -> None: - """--exact-tensor-pattern keeps tensors as dense deltas; the merge must still apply them.""" - base = _base_state() - adapter = _adapter(base) - - merged = merge_lora.merge_lora_into_base(base, adapter) - - expected_lora = base["blocks.0.linear.weight"] + adapter["blocks.0.linear.lora_B.weight"] @ adapter[ - "blocks.0.linear.lora_A.weight"] - assert torch.allclose(merged["blocks.0.linear.weight"], expected_lora, atol=1e-6) - assert torch.allclose(merged["audio_proj_in.weight"], - base["audio_proj_in.weight"] + adapter["audio_proj_in.diff"], - atol=1e-6) - assert torch.allclose(merged["time_embedder.linear.bias"], - base["time_embedder.linear.bias"] + adapter["time_embedder.linear.diff_b"], - atol=1e-6) - assert torch.allclose(merged["scale_param"], base["scale_param"] + adapter["scale_param.diff_param"], atol=1e-6) - assert torch.equal(merged["new_proj.weight"], adapter["new_proj.set_weight"]) - # the caller's state dict must not be mutated - assert torch.equal(base["audio_proj_in.weight"], _base_state()["audio_proj_in.weight"]) - - -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 test_indexed_wan_extraction_merges_into_fastvideo_namespace(tmp_path: Path) -> None: - generator = torch.Generator().manual_seed(21) - base_hf = { - "blocks.0.attn1.to_q.weight": torch.randn(5, 4, generator=generator), - "condition_embedder.time_proj.weight": torch.randn(3, 4, generator=generator), - } - finetuned_hf = {name: tensor.clone() for name, tensor in base_hf.items()} - finetuned_hf["blocks.0.attn1.to_q.weight"] += torch.randn( - 5, 2, generator=generator) @ torch.randn(2, 4, generator=generator) - finetuned_hf["condition_embedder.time_proj.weight"] += 0.25 - base_dir = tmp_path / "base" - finetuned_dir = tmp_path / "finetuned" - output = tmp_path / "adapter.safetensors" - _write_transformer(base_dir, base_hf) - _write_transformer(finetuned_dir, finetuned_hf) - 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", - exact_tensor_patterns=(r"^condition_embedder\.time_proj\.weight$", ), - ) - - base_custom = { - "blocks.0.to_q.weight": base_hf["blocks.0.attn1.to_q.weight"], - "condition_embedder.time_modulation.linear.weight": base_hf["condition_embedder.time_proj.weight"], - } - config = WanVideoConfig() - merged = merge_lora.merge_lora_into_base( - base_custom, - merge_lora.load_adapter(str(output)), - lora_param_names_mapping=get_param_names_mapping(config.lora_param_names_mapping), - param_names_mapping=get_param_names_mapping(config.param_names_mapping), - ) - - torch.testing.assert_close(merged["blocks.0.to_q.weight"], finetuned_hf["blocks.0.attn1.to_q.weight"]) - torch.testing.assert_close(merged["condition_embedder.time_modulation.linear.weight"], - finetuned_hf["condition_embedder.time_proj.weight"]) - - -def test_merge_accumulates_multiple_sources_into_a_fused_parameter() -> None: - base = {"fused.weight": torch.zeros(6, 4)} - adapter = { - "q.lora_A.weight": torch.ones(1, 4), - "q.lora_B.weight": torch.ones(3, 1), - "k.lora_A.weight": torch.full((1, 4), 2.0), - "k.lora_B.weight": torch.ones(3, 1), - } - - def mapping(name: str): - return "fused.weight", 0 if name.startswith("q.") else 1, 2 - - merged = merge_lora.merge_lora_into_base(base, adapter, param_names_mapping=mapping) - torch.testing.assert_close(merged["fused.weight"][:3], torch.ones(3, 4)) - torch.testing.assert_close(merged["fused.weight"][3:], torch.full((3, 4), 2.0)) - - -def test_merge_is_strict_about_unapplied_keys() -> None: - with pytest.raises(ValueError, match="unapplied"): - merge_lora.merge_lora_into_base(_base_state(), {"missing.diff": torch.zeros(2)}) - - -def test_merge_warns_instead_of_silently_dropping_unknown_keys(caplog) -> None: - base = _base_state() - adapter = {"mystery.tensor": torch.zeros(2), "shape_mismatch.diff": torch.zeros(1, 1)} - - with caplog.at_level(logging.WARNING, logger=merge_lora.LOG.name): - merge_lora.merge_lora_into_base(base, adapter, strict=False) - - messages = " ".join(record.getMessage() for record in caplog.records) - assert "unrecognized adapter key" in messages - assert "mystery.tensor" in messages - assert "shape_mismatch.weight" in messages diff --git a/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py b/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py index f0aaa6c895..1a10cd09bc 100644 --- a/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py +++ b/fastvideo/tests/lora_extraction/test_streaming_lora_extraction.py @@ -224,6 +224,10 @@ def test_standalone_parameters_use_generic_dense_suffixes(tmp_path: Path) -> Non 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: @@ -339,6 +343,27 @@ def test_randomized_factorization_runs_on_gpu() -> None: 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) @@ -462,42 +487,84 @@ def test_resume_still_rejects_a_mismatched_config(tmp_path: Path) -> None: _extract(base_dir, finetuned_dir, output, rank=4, resume=True) -def test_resume_rejects_changed_checkpoint_contents(tmp_path: Path) -> None: +@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) - _, changed = _toy_states() + base, finetuned = _toy_states() + changed = base if changed_side == "base" else finetuned changed["blocks.0.linear.weight"] += 1.0 - shard = finetuned_dir / "transformer" / "diffusion_pytorch_model-00001-of-00001.safetensors" + 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) -def test_local_paths_reject_hugging_face_revisions(tmp_path: Path) -> None: +@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", - base_revision="abc123") + _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, + ) -def test_pipeline_mode_rejects_revisions_before_loading(tmp_path: Path) -> None: +@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", - base_revision="abc123", + **{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, ) -def test_auto_mode_does_not_drop_revision_during_fallback(monkeypatch: pytest.MonkeyPatch, - tmp_path: Path) -> None: +@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") @@ -512,10 +579,15 @@ def forbidden_pipeline_load(*args, **kwargs): ft="org/finetuned", out=str(tmp_path / "adapter.safetensors"), load_mode="auto", - base_revision="abc123", + **{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) diff --git a/scripts/lora_extraction/README.md b/scripts/lora_extraction/README.md index 91553868b9..f4ab15353e 100644 --- a/scripts/lora_extraction/README.md +++ b/scripts/lora_extraction/README.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 @@ -16,17 +16,19 @@ python extract_lora.py \ --exact-tensor-pattern '^proj_out\.weight$' ``` -The extractor is runtime-agnostic by default: it cannot infer from checkpoint tensors whether a runtime wraps a +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 rather than being discarded. The Wan patterns -above cover its excluded condition embedders and its unwrapped output projection. +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 \ - --ft FastVideo/FastVideo-FastH3-8-step-Preview-v1-VSA-DataFree \ + --base-revision 9bfb6693f2cf6de171db46d1aa586f67d773a1da \ + --ft FastVideo/FastVideo-FastH3-4-step-v1.1 \ + --ft-revision c9e910404950b42f627f07b0c4a09d9a3e087d47 \ --out adapter_r64.safetensors \ --rank 64 \ --load-mode indexed \ @@ -45,10 +47,15 @@ python extract_lora.py \ `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`: requested LoRA rank. +- `--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`. @@ -73,23 +80,28 @@ Fine-tuned parameters that cannot or should not be factorized are retained autom Indexed loading is preferred and downloads only the transformer component. `--load-mode auto` falls back to legacy pipeline loading when indexed safetensors are unavailable. +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 -- `--allow-unmatched`: Opt in to writing an output when adapter keys cannot be applied; strict matching is the default. ## Validate Quality (Optional) diff --git a/scripts/lora_extraction/extract_lora.py b/scripts/lora_extraction/extract_lora.py index 198e14e65c..b6da85e517 100644 --- a/scripts/lora_extraction/extract_lora.py +++ b/scripts/lora_extraction/extract_lora.py @@ -305,6 +305,9 @@ def load_transformer_state_dict_from_model( 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, @@ -313,6 +316,9 @@ def load_transformer_state_dict_from_model( if transformer is None: modules = getattr(pipeline, "modules", None) transformer = modules.get("transformer") if isinstance(modules, dict) else None + if transformer is None: + nested_pipeline = getattr(pipeline, "pipeline", None) + transformer = getattr(nested_pipeline, "transformer", None) if transformer is None: raise RuntimeError("Transformer not found in pipeline") @@ -581,12 +587,12 @@ def _extract_layers( } _atomic_json_dump(manifest, manifest_path) continue - output_key = (key.removesuffix(".weight") + SET_WEIGHT_SUFFIX - if key.endswith(".weight") else key + SET_PARAM_SUFFIX) + 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", + "kind": "set_weight" if is_weight else "set_param", "shape": list(finetuned_tensor.shape), "tensor_file": tensor_file.name, "output_keys": [output_key], @@ -723,7 +729,7 @@ def _verify_adapter(out_path: Path, manifest: dict[str, Any]) -> None: 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"}: + 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: @@ -790,6 +796,9 @@ def extract_lora_adapter( 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": @@ -804,7 +813,7 @@ def extract_lora_adapter( else: effective_work_dir = out_path.parent / f".{out_path.name}.work" - with _open_readers(base, ft, base_revision, ft_revision, load_mode) as (base_reader, finetuned_reader): + 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") @@ -856,6 +865,7 @@ def extract_lora_adapter( "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", } diff --git a/scripts/lora_extraction/merge_lora.py b/scripts/lora_extraction/merge_lora.py index 33d6e8bcfe..f06a45618d 100644 --- a/scripts/lora_extraction/merge_lora.py +++ b/scripts/lora_extraction/merge_lora.py @@ -16,8 +16,6 @@ import logging from pathlib import Path from collections import defaultdict -from collections.abc import Callable -from typing import Any os.environ.setdefault("MASTER_ADDR", "127.0.0.1") os.environ.setdefault("MASTER_PORT", "29500") @@ -97,118 +95,6 @@ def load_adapter(adapter_path: str) -> dict: return fix_adapter_naming(adapter) -# Suffix -> the parameter suffix it targets, mirroring fastvideo.models.loader.lora_patch. -# An empty target keeps the full name, for standalone nn.Parameters. -ADDITIVE_SUFFIXES: dict[str, str] = {".diff_param": "", ".diff_b": ".bias", ".diff": ".weight"} -REPLACEMENT_SUFFIXES: dict[str, str] = {".set_weight": ".weight", ".set_param": ""} -LORA_SUFFIXES = (".lora_A.weight", ".lora_B.weight", ".lora_rank", ".lora_alpha") - - -def group_dense_keys(adapter: dict) -> tuple[dict, dict, list]: - """Split the non-factorized half of an adapter into additive and replacement params. - - ``--exact-tensor-pattern`` keeps selected matrices as exact dense deltas instead of - LoRA factors, so an adapter merged without these is missing those tensors entirely. - """ - additive: dict = {} - replacement: dict = {} - unrecognized: list = [] - - for key, tensor in adapter.items(): - if key.endswith(LORA_SUFFIXES): - continue - for suffix, param_suffix in ADDITIVE_SUFFIXES.items(): - if key.endswith(suffix): - additive[key.removesuffix(suffix) + param_suffix] = tensor - break - else: - for suffix, param_suffix in REPLACEMENT_SUFFIXES.items(): - if key.endswith(suffix): - replacement[key.removesuffix(suffix) + param_suffix] = tensor - break - else: - unrecognized.append(key) - - return additive, replacement, unrecognized - - -ParamMapping = Callable[[str], tuple[str, Any, Any]] - - -def map_adapter_parameter( - name: str, - lora_param_names_mapping: ParamMapping | None, - param_names_mapping: ParamMapping | None, -) -> tuple[str, int | None, int | None]: - """Apply the same official-LoRA -> HF -> FastVideo mapping order as runtime loading.""" - if lora_param_names_mapping is not None: - name, merge_index, total = lora_param_names_mapping(name) - if merge_index is not None: - raise NotImplementedError(f"Adapter-specific mapping unexpectedly fused {name} ({merge_index}/{total})") - if param_names_mapping is None: - return name, None, None - return param_names_mapping(name) - - -def _mapped_slice(target: torch.Tensor, merge_index: int | None, - total: int | None) -> tuple[Any, tuple[int, ...]] | None: - if merge_index is None: - return Ellipsis, tuple(target.shape) - if total is None or total < 1 or target.ndim < 1 or target.shape[0] % total: - return None - chunk = target.shape[0] // total - return slice(merge_index * chunk, (merge_index + 1) * chunk), (chunk, *target.shape[1:]) - - -def merge_dense_into_base( - merged_sd: dict, - additive: dict, - replacement: dict, - lora_param_names_mapping: ParamMapping | None = None, - param_names_mapping: ParamMapping | None = None, -) -> tuple[int, list[str]]: - """Apply exact dense deltas and replacement parameters in place.""" - merged_count = 0 - skipped: list[str] = [] - - for source_name, tensor in additive.items(): - param_name, merge_index, total = map_adapter_parameter(source_name, lora_param_names_mapping, - param_names_mapping) - target = merged_sd.get(param_name) - target_slice = _mapped_slice(target, merge_index, total) if target is not None else None - if target is None or target_slice is None or target_slice[1] != tuple(tensor.shape): - skipped.append(f"{source_name} -> {param_name}") - continue - index, _ = target_slice - updated = target.to(torch.float32).clone() - updated[index] += tensor.to(torch.float32) - merged_sd[param_name] = updated.to(target.dtype) - merged_count += 1 - - for source_name, tensor in replacement.items(): - param_name, merge_index, total = map_adapter_parameter(source_name, lora_param_names_mapping, - param_names_mapping) - target = merged_sd.get(param_name) - if target is None: - if merge_index is not None: - skipped.append(f"{source_name} -> {param_name}") - continue - merged_sd[param_name] = tensor - merged_count += 1 - continue - target_slice = _mapped_slice(target, merge_index, total) - if target_slice is None or target_slice[1] != tuple(tensor.shape): - skipped.append(f"{source_name} -> {param_name}") - continue - index, _ = target_slice - updated = target.clone() - updated[index] = tensor.to(target.dtype) - merged_sd[param_name] = updated - merged_count += 1 - - return merged_count, skipped - - def group_adapter_keys(adapter: dict) -> dict: grouped = defaultdict(dict) @@ -226,7 +112,7 @@ def group_adapter_keys(adapter: dict) -> dict: return grouped -def get_reverse_param_mapping(base_model_path: str) -> tuple[dict, ParamMapping, ParamMapping]: +def get_reverse_param_mapping(base_model_path: str): LOG.info("Loading base model for parameter mapping") pipeline_cls = get_pipeline_class_for_model(base_model_path) @@ -250,16 +136,16 @@ def get_reverse_param_mapping(base_model_path: str) -> tuple[dict, ParamMapping, if transformer is None: raise RuntimeError("Could not find transformer in pipeline") - param_names_mapping_fn = get_param_names_mapping(transformer.param_names_mapping) - lora_param_names_mapping_fn = get_param_names_mapping(transformer.lora_param_names_mapping) - - if getattr(transformer, "reverse_param_names_mapping", None): + if hasattr(transformer, "reverse_param_names_mapping"): reverse_mapping = transformer.reverse_param_names_mapping elif hasattr(transformer, "config") and hasattr(transformer.config, "arch_config"): arch_config = transformer.config.arch_config - if getattr(arch_config, "reverse_param_names_mapping", None): + if hasattr(arch_config, "reverse_param_names_mapping"): reverse_mapping = arch_config.reverse_param_names_mapping else: + param_mapping = arch_config.param_names_mapping + param_names_mapping_fn = get_param_names_mapping(param_mapping) + from diffusers import DiffusionPipeline from huggingface_hub import snapshot_download @@ -284,48 +170,36 @@ def get_reverse_param_mapping(base_model_path: str) -> tuple[dict, ParamMapping, del transformer torch.cuda.empty_cache() - return reverse_mapping, param_names_mapping_fn, lora_param_names_mapping_fn + return reverse_mapping -def merge_lora_into_base( - base_sd: dict, - adapter: dict, - *, - lora_param_names_mapping: ParamMapping | None = None, - param_names_mapping: ParamMapping | None = None, - strict: bool = True, -) -> dict: +def merge_lora_into_base(base_sd: dict, adapter: dict) -> dict: LOG.info("Merging LoRA into base weights") adapter_layers = group_adapter_keys(adapter) merged_sd = dict(base_sd) merged_count = 0 - skipped: list[str] = [] + skipped_count = 0 for base_name, parts in adapter_layers.items(): - source_weight_key = base_name if base_name.endswith(".weight") else base_name + ".weight" - weight_key, merge_index, total = map_adapter_parameter(source_weight_key, lora_param_names_mapping, - param_names_mapping) - base_tensor = merged_sd.get(weight_key) - if base_tensor is None: - skipped.append(f"{source_weight_key} -> {weight_key} (missing base parameter)") + weight_key = base_name if base_name.endswith(".weight") else base_name + ".weight" + + if weight_key not in base_sd: + skipped_count += 1 continue if "A" not in parts or "B" not in parts: - skipped.append(f"{source_weight_key} (incomplete factor pair)") + skipped_count += 1 continue lora_A = parts["A"].to(torch.float32) lora_B = parts["B"].to(torch.float32) - target_slice = _mapped_slice(base_tensor, merge_index, total) - if target_slice is None: - skipped.append(f"{source_weight_key} -> {weight_key} (invalid fused mapping)") - continue - index, expected_shape = target_slice - out_dim, in_dim = expected_shape + base_weight = base_sd[weight_key].to(torch.float32) + + out_dim, in_dim = base_weight.shape if lora_B.shape[0] != out_dim or lora_A.shape[1] != in_dim or lora_B.shape[1] != lora_A.shape[0]: - skipped.append(f"{source_weight_key} -> {weight_key} (factor shape mismatch)") + skipped_count += 1 continue delta = lora_B @ lora_A @@ -336,28 +210,11 @@ def merge_lora_into_base( if rank != 0 and alpha != rank: delta = delta * (alpha / float(rank)) - updated = base_tensor.to(torch.float32).clone() - updated[index] += delta - merged_sd[weight_key] = updated.to(base_tensor.dtype) + merged_weight = base_weight + delta + merged_sd[weight_key] = merged_weight.to(base_sd[weight_key].dtype) merged_count += 1 - additive, replacement, unrecognized = group_dense_keys(adapter) - dense_merged, dense_skipped = merge_dense_into_base( - merged_sd, - additive, - replacement, - lora_param_names_mapping, - param_names_mapping, - ) - - LOG.info("Merged %d LoRA layers, skipped %d", merged_count, len(skipped)) - LOG.info("Merged %d dense tensors (%d additive, %d replacement)", dense_merged, len(additive), len(replacement)) - problems = skipped + dense_skipped + [f"unrecognized adapter key {key}" for key in unrecognized] - if problems and strict: - raise ValueError(f"Adapter merge left {len(problems)} keys/layers unapplied: {problems[:5]}") - if problems: - LOG.warning("Adapter merge left %d keys/layers unapplied: %s", len(problems), problems[:5]) - + LOG.info(f"Merged {merged_count} layers, skipped {skipped_count}") return merged_sd @@ -424,7 +281,6 @@ def merge_lora( ft: str, output: str, log_level: str = "INFO", - allow_unmatched: bool = False, ) -> None: """Merge LoRA adapter into base model. @@ -434,7 +290,6 @@ def merge_lora( ft: Finetuned model ID (for config) output: Output directory log_level: Logging level - allow_unmatched: Write output even if recognized adapter payloads cannot be applied """ configure_logging(log_level) @@ -442,20 +297,14 @@ def merge_lora( LOG.info(f"Adapter: {adapter}") LOG.info(f"Output: {output}") - reverse_mapping, param_names_mapping_fn, lora_param_names_mapping_fn = get_reverse_param_mapping(base) + reverse_mapping = get_reverse_param_mapping(base) LOG.info(f"Loading base model: {base}") base_sd = load_transformer_state_dict_from_model(base) LOG.info(f"Loaded { len(base_sd)} parameters") adapter_sd = load_adapter(adapter) - merged_sd = merge_lora_into_base( - base_sd, - adapter_sd, - lora_param_names_mapping=lora_param_names_mapping_fn, - param_names_mapping=param_names_mapping_fn, - strict=not allow_unmatched, - ) + merged_sd = merge_lora_into_base(base_sd, adapter_sd) save_merged_model(merged_sd, base, ft, output, reverse_mapping) LOG.info("Merge complete") @@ -469,8 +318,6 @@ def main(): parser.add_argument("--ft", required=True, help="Finetuned model ID (for config)") parser.add_argument("--output", required=True, help="Output directory") parser.add_argument("--log-level", default="INFO", help="Logging level") - parser.add_argument("--allow-unmatched", action="store_true", - help="Write the merged model even when adapter keys cannot be applied") args = parser.parse_args() merge_lora( @@ -479,7 +326,6 @@ def main(): ft=args.ft, output=args.output, log_level=args.log_level, - allow_unmatched=args.allow_unmatched, )