Skip to content

[Model] CosyVoice3: fold generator and F0 weight norms after loading - #7871

Closed
zhongc1211 wants to merge 4 commits into
vllm-project:mainfrom
zhongc1211:fix/cosyvoice3-fold-hift-weight-norm
Closed

zhongc1211 wants to merge 4 commits into
vllm-project:mainfrom
zhongc1211:fix/cosyvoice3-fold-hift-weight-norm

Conversation

@zhongc1211

@zhongc1211 zhongc1211 commented Sep 20, 2026 •

Copy link
Copy Markdown
Contributor

Purpose

CosyVoice3's HiFT generator and F0 predictor no longer recompute g * v / ||v|| for every streamed chunk. This PR folds the 77 generator weight norms and five F0 weight norms after strict checkpoint loading, while keeping F0 on CPU in FP32.

This now covers C4 and C5. Generator and F0 folding use separate entry points: the loader places and folds the generator, moves F0 to CPU/FP32 before folding its weights, then sets HiFT to eval. C2 backend routing is not included; integrating it later needs another placement/dtype check before folding.

Both PyTorch APIs use a shared normalizer, following @0z5a's tested fix. It keeps ordinary Parameters and their requires_grad, and clones plain/inference tensors outside inference mode so autograd can use the folded weights. Repeated removal is a no-op. Thanks @linyueqian for identifying the legacy coverage gap.

Folding changes checkpoint keys. Load original checkpoints into fresh instances; strict reloading into a folded module is unsupported. Sleep/wake inspection does not establish inherited/RPC reload safety. Later device or dtype changes are outside the verified contract.

Related: #6870 (C4+C5), #7827 (closed overlapping PR), #6927 (generic helper), #7518 (C2 backend routing), #7521 (incremental streaming and historical E2E harness).

Test Plan

vLLM Version: 0.29.0; Python 3.12.3; torch 2.13.0+cu130; CUDA 13.0.

vLLM-Omni Commit: 5c7a151f7b05949fd73c0a96893272e71031fe4d.

The 2026-09-25 remote run used the C4 commit 2e021d07 plus the integration patch. All four source/test file hashes match the commit above. Hardware: RTX 4000 Ada, with four CPU threads. Use the repository's test dependencies and pre-commit with the versions above; CUDA is required for the CUDA cases. The synthetic tests need no model download or HF cache.

From the repository root:

export OMP_NUM_THREADS=4 MKL_NUM_THREADS=4 OPENBLAS_NUM_THREADS=4
files=(
  tests/model_executor/models/cosyvoice3/test_cosyvoice3_hift_weight_norm.py
  tests/model_executor/models/cosyvoice3/test_cosyvoice3_f0_weight_norm.py
)
python3 -m pytest "${files[@]}" -q -o addopts=''
python3 -m pytest "${files[@]}" -q -o addopts='' -m 'core_model and cpu' --run-level=core_model
python3 -m pytest "${files[@]}" -q -o addopts='' -m 'core_model and cuda' --run-level=core_model
python3 -m pytest tests/model_executor/models/cosyvoice3 -q -o addopts=''
pre-commit run --files "${files[@]}" \
  vllm_omni/model_executor/models/cosyvoice3/code2wav_core/hifigan.py \
  vllm_omni/model_executor/models/cosyvoice3/cosyvoice3_code2wav.py

A separate probe exercised the actual loader and shipped HiFT with FunAudioLLM/Fun-CosyVoice3-0.5B-2512, revision 29e01c4e8d000f4bcd70751be16fa94bf3d85a18. The hift.pt SHA256 was b279d7641eb97ae55b3b540cfba4f953c26492a2df758328a89a4d007ab87a65. It used synthetic mel, a Linear Flow stub, FP32, seeds 5101/5102, 192 mel frames split into eight 24-frame chunks, a 64-frame window, deterministic CuDNN and TF32 disabled. For each generator backend (CPU or CUDA), three separate processes ran unfolded, generator-only folded and jointly folded arms; F0 stayed on CPU in all arms. Loading and streaming ran under inference_mode.

The checkpoint probe harness and raw logs are retained locally, not included in this PR or uploaded. The commands above reproduce the committed tests, not that separate probe.

Test Result

Current source check Result
Joint focused tests 69 passed
CPU CI-like selection 56 passed, 13 deselected
CUDA CI-like selection 13 passed, 56 deselected
CosyVoice3 directory suite 168 passed
Final tests against unchanged C4-only production 47 expected failures, 22 passes; no errors/skips
Applicable pre-commit checks before and after runtime validation Passed; source/test bytes unchanged

The CPU and CUDA selections partition the focused 69, which are included in the directory suite's 168. These counts are not additive; the directory suite is not the whole repository suite.

  • Coverage includes both weight-norm APIs, modern/legacy checkpoint keys, cold-loaded changed gains, strict-load failure before either fold, idempotence, non-inference folded Parameters and backward across grad contexts, and CPU/CUDA-generator placement with CPU/FP32 F0. Reduced-vocoder streaming tests compare audio and caches against both unfolded and C4-only references, including window movement and finalization. The shipped-config count test constructs and folds all 77+5 layers without a forward pass.
  • The real-checkpoint probe matched all four same-backend reference-versus-joint comparisons exactly: input/source RNG tensors, audio chunks, concatenation and retained caches. Each arm produced eight chunks and 92,160 samples, shifted the window and cleared the final cache. Generator/F0 norm counts were 77/5, 0/5 and 0/0. References disabled selected folding calls on the candidate source, not pristine main.
  • The first probe failed an overly broad assertion that every HiFT parameter must be non-inference. A control reproduced inference-marked biases from inference-mode CUDA placement without folding. Only the harness changed: it checks all 82 folded weights and compares other parameter flags across arms. All six processes were then rerun successfully; production/tests were unchanged. The failed attempt is retained.
  • Pytest emitted 16 warnings, including TorchScript deprecations and a vLLM/vLLM-Omni version-metadata mismatch (0.29.0 versus 0.1.dev1+g2e021d073). Passing these tests does not establish general version compatibility.

These results do not establish full text-to-audio serving correctness, real Flow inference, cross-backend parity, audio quality or a speedup. No performance, NPU or TRT-estimator validation was run for this integration. Remote test results are separate from GitHub CI status.

Earlier C4 evidence and corrections (not measurements of this commit)

The C4-only 2e021d07 snapshot passed 26 focused CPU cases, including three CI-filtered loader cases, plus applicable pre-commit checks. That run used Linux x86_64, 2 vCPU/8 GB, two threads and no CUDA. An older RTX 2000 Ada snapshot passed 124 CosyVoice3 cases, including 25 focused cases, and 13 separate MiniCPM-o CUDA Graph cases. None of these counts is added to the current results.

An earlier negative control on d3c6eb1 with expanded tests produced 24 passes and one failure at the legacy inference-Parameter assertion, before forward/backward. The earlier separate-process real-checkpoint probe matched one seeded float32 eight-chunk synthetic mel stream and caches with generator/F0 counts 77/5 versus 0/5. It used candidate source with folding disabled, not pristine main, and did not exercise Flow synthesis or full serving. The older same-instance probe with empty logger capture remains a failed run.

Historical random-initialized vocoder decode measurements, not rerun for this integration:

GPU Frames Pairs Baseline median (ms) Folded median (ms) Median paired saving (ms)
RTX A4500 41 200 26.671 22.868 3.829
RTX A4500 191 200 30.165 25.668 4.511
RTX 2000 Ada 41 30 17.832 15.019 2.819
RTX 2000 Ada 191 30 25.830 23.460 2.352

Paired savings are medians of per-pair differences, not differences of medians. Raw pairs were not saved, so the bootstrap intervals cannot be recomputed independently. No serving speedup is established. Raw logs and the historical reproduction bundle remain local and have not been uploaded.

Correction to the earlier write-up: ok=30,corr=27 means 30 completed requests and 27 correctness problems, not 27 passes. First requests create references, so the remaining three are not established successful comparisons. Equal problem counts do not prove parity. The jitter attribution, 0.2–0.7% serving projection, ±28% run-to-run-noise claim and “within one SD” reassurance are withdrawn. Request-level SD is not run-to-run uncertainty. From rounded equal-size means, aggregate E2E increased about 2.76% at c1 and 26.27% at c8; the 11.90% increase was c8 mean RTF, not E2E. That run established neither serving correctness nor absence of regression; new tests cannot validate it.

Post-Deploy Monitoring & Validation

At first load with the shipped configuration, expect Folded 77 weight-norm layers in HiFT generator and Folded 5 weight-norm layers in HiFT F0 predictor, with no remaining HiFT weight norms. Compare finite audio, window updates and finalization on the same streaming workload before and after the change. Investigate an unexpected zero-fold warning, load failure or reproducible audio regression and roll back if needed. Rollout owner and validation window are unassigned.

BEFORE SUBMITTING: read CONTRIBUTING.md and run the precheck-pr skill with the code agent for a self-check against project conventions.
(anything written below this line will be removed by GitHub Actions)

…kpoint load

CosyVoice3's HiFT vocoder keeps 77 weight-norm parametrizations on frozen
inference weights, so every decode recomputes g*v/||v|| per layer: profiled
as 77 aten::_weight_norm + 77 CUDA weight_norm_fwd_first_dim_kernel events
per chunk (RFC 6870, workstream C4).

The shipped remove_weight_norm() was unusable: it calls the legacy
torch.nn.utils spelling on modules built with the parametrizations API
(ValueError on any current torch), invokes a non-existent
SourceModuleHnNSF.remove_weight_norm(), and walks source_downs layers that
were never weight-normed. Nothing called it.

Fix and wire it into the load path:

- per-layer _fold_weight_norm() supports both APIs (parametrizations API;
  legacy WeightNorm forward-pre-hook via isinstance, same idea as the
  merged indextts2 _strip_weight_norm), returns 0/1, is idempotent, and
  normalizes the materialized weight to a frozen nn.Parameter regardless
  of the caller's grad mode (no_grad/inference_mode callers would
  otherwise get a buffer or an inference tensor);
- ResBlock/HiFTGenerator.remove_weight_norm() return the folded count so a
  silent no-op shows up in the log (same convention as minimax_music3);
- CosyVoice3Code2Wav.load_weights() folds after the strict checkpoint load
  and device move;
- f0_predictor's 5 weight norms are deliberately left parametrized
  (pinned-CPU precision path; RFC C5, to be done with PR 7518).
  step_audio2 and glm_tts import HiFTGenerator but never fold, so they
  are unaffected.

One-time post-load transform: folding rewrites the state_dict schema
(parametrizations.weight.original0/1 -> weight). vllm-omni has no
reload_weights path and sleep/wake restores physical pages via
CuMemAllocator without re-invoking the loader, so nothing in-tree loads
twice; a second load of the original checkpoint would fail the strict
load_state_dict loudly.

Validation (Runpod, vllm 0.29.0, torch 2.13.0+cu130, base d4ffde1):
- tests/model_executor/models/cosyvoice3/: 114 passed (RTX 2000 Ada host);
  shared-consumer checks: minicpmo cuda-graph test 13 passed, step_audio2
  and glm_tts imports OK
- pre-commit run --files <3 changed files>: all gates passed (ruff,
  typos, mypy-3.10, CI marks, SPDX, forbidden imports, torch.cuda guard)
- new L1 regressions (16 cases): fold exactness atol=rtol=0 across
  parametrized and legacy APIs, 22050 and 24000 Hz, finalize and
  streaming; fold counts, idempotency, Parameter-ness (including folding
  under inference_mode); shipped-config generator folds exactly 77 with
  only the 5 f0 layers left; load_weights folds loaded (not initial)
  weights for modern and legacy checkpoint key schemas. All 16 fail on
  unpatched d4ffde1.
- real checkpoint (FunAudioLLM/Fun-CosyVoice3-0.5B-2512, GPU): production
  load_weights logs 'Folded 77 weight-norm layers'; inference() on the
  real hift weights is bit-identical before vs after fold; generator 0 /
  f0 5 parametrizations left; folded weight is a Parameter on CUDA
- CUDA profiler (real-sized generator, decode path): weight-norm events
  231 -> 0; inference output bit-identical (torch.equal=True)
- wall-clock A/B, paired interleaved rounds, alternating arm order,
  bootstrap 95% CI: RTX A4500 (200 rounds, initial fold revision):
  decode() 3.83 ms [3.74, 3.89] chunk 41 (200/200 faster) and 4.51 ms
  [4.33, 4.62] chunk 191 (197/200); RTX 2000 Ada (30 rounds, final head):
  decode() 2.82 and 2.35 ms (30/30 each). Full inference() path is
  dominated by unchanged f0_predictor jitter (C2/7518 scope); no e2e RTF
  claim made (7521's harness cannot resolve this effect on shared GPUs).

Refs: RFC 6870 (C4).
Related: PR 6927 (generic fold helper for MiniCPM-o),
PR 7518 (C2/C5 f0 predictor).

Signed-off-by: zack <huixindaddy@yahoo.com>
@vllm-omni-review-bot

Copy link
Copy Markdown

This PR appears to belong to: docs/design/module/model_integration.md, docs/design/module/ar_runtime.md.

Module owners: @tzhouam @fake0fan @Gaohan123

Routing: @tzhouam via module of the changed files, CODEOWNERS; @fake0fan via module of the changed files; @Gaohan123 via module of the changed files

@zhongc1211, please review your own changes and leave a short self-review comment describing what you checked. PRs without author self-review may not be assigned a reviewer.

Please take a look when you have a chance. If you would like an automated review, mention @vllm-omni-review-bot in a comment.

@zhongc1211

zhongc1211 commented Sep 20, 2026 •

Copy link
Copy Markdown
Contributor Author

Validation update for 5c7a151f7b05

The scope is now 77 generator folds plus five F0 folds. F0 remains on CPU in FP32; C2 routing is not included. The shared normalizer from @0z5a and the legacy coverage requested by @linyueqian are retained.

  • The exact four source/test files now pushed passed 69 focused cases: 56 CPU and 13 CUDA. All 69 are included in the 168 passing CosyVoice3 directory cases. Applicable pre-commit checks passed without changing those files.
  • The final tests against unchanged C4-only production produced 47 expected failures and 22 passes, with no errors or skips. This checks that the new assertions detect missing C5 behavior.
  • Separate-process probes of the real HiFT checkpoint matched audio and caches exactly against unfolded and C4-only references on each of two backends: CPU generator/CPU F0 and CUDA generator/CPU F0. Each arm used eight chunks and produced 92,160 samples. This used synthetic mel and a Linear Flow stub, not full serving or a performance benchmark.
  • The first checkpoint probe failed because its assertion covered all HiFT parameters rather than folded weights. A control reproduced inference-marked biases from inference-mode device placement without folding. The harness was corrected and all six processes rerun; source/tests did not change, and the failed attempt is retained.
  • Environment: RTX 4000 Ada, four CPU threads, Python 3.12.3, torch 2.13.0+cu130 and vLLM 0.29.0. The runs retained 16 warnings, including the vLLM/vLLM-Omni version-metadata mismatch. No full-serving, perceptual, NPU or current performance validation was run. Original-checkpoint reload into folded instances and subsequent device/dtype changes remain outside the verified contract.

The earlier C4-only results and E2E corrections below remain historical; they are not measurements of this head. No serving-speedup or end-to-end parity claim is added.

Earlier self-review for `2e021d07` (historical)

Self-review update for 2e021d07:

I have reviewed and understand the code changes.

  • Scope and loading: strict checkpoint load → device placement → fold → eval. The shipped generator folds 77 norms and leaves five F0 norms unchanged. Folding changes state-dict keys: reloading the original checkpoint into the folded instance is unsupported. Sleep/wake inspection does not establish inherited or RPC reload safety.
  • Legacy fix: both APIs now use the shared normalizer from @0z5a's tested patch. Ordinary non-inference Parameters retain their flags. Six API/grad-context cases check for non-inference weights and complete backward. The earlier negative control failed at the inference-Parameter assertion, before forward/backward.
  • Current validation: the exact source bytes now pushed passed 26 focused CPU cases, no skips, plus applicable pre-commit checks without source changes. The three CI-filtered loader cases are included in those 26. Tests cover actual legacy construction, loading before any forward, the transition from training to eval, and multichunk finalization. The small audio fixture folds 41 layers; the shipped-config 77-layer test is construction/folding-only.
  • Environment and limits: Linux x86_64, 2 vCPU/8 GB, Python 3.12.3, torch 2.13.0+cu130, vLLM 0.29.0, CUDA unavailable. Each pytest run emitted 16 warnings, including TorchScript deprecations and a version-metadata mismatch. Earlier GPU results—124 CosyVoice3 cases, including 25 focused cases, and 13 separate MiniCPM-o CUDA Graph cases—belong to a different snapshot, not this commit. The earlier separate-process real-checkpoint probe matched one seeded eight-chunk synthetic mel stream and its caches; it did not exercise Flow synthesis or full serving. The older same-instance run with empty logger capture remains a failed run.

I also need to correct the earlier PR description and self-review. In that E2E harness, ok=30,corr=27 means 30 completed requests and 27 correctness problems, not 27 passes. The first request for each text creates a reference, so the remaining three are not established successful comparisons. Equal problem counts do not prove parity, and attributing them to shared-host jitter was unsupported.

From the rounded equal-size means, aggregate E2E increased about 2.76% at c1 and 26.27% at c8; the 11.90% increase was c8 mean RTF, not E2E. Request-level SD is not run-to-run uncertainty. I withdraw the 0.2–0.7% serving projection, ±28% run-to-run-noise claim and “within one SD” reassurance. That experiment established neither serving correctness nor absence of regression; the new CPU tests do not validate it.

The A4500/Ada numbers are historical random-initialized vocoder microbenchmarks, not new measurements or a serving-speedup result. Raw pairs were not retained, so the bootstrap intervals cannot be recomputed independently. Raw logs and the historical reproduction bundle remain local and have not been uploaded.

@hsliuustc0106 hsliuustc0106 added the enhancement New feature or request label Sep 20, 2026

@linyueqian linyueqian left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The fold itself is sound: it runs after strict checkpoint loading and device placement, covers the 77 generator convolutions, leaves the five f0 predictor parametrizations in place, and the current sleep/wake path does not reload checkpoints, so there is no re-parametrization hazard. The 16 parametrized cases cover the folded layers; one optional coverage suggestion is inline. Static review found no regression.

One coordination item with no line anchor, and the reason this is a comment rather than an approval for now: #7827 by @0z5a (opened a day earlier for #6870 workstream C4) folds the same HiFT weight norms after load, but removes all 82 parametrizations including the f0 ones, where this PR keeps those five. Two PRs for one workstream item should not both land; @timzsu, as the RFC owner, and the two authors should settle which one goes forward (or whether the f0 layers should be folded too), and I will approve the surviving one on its head. Static read at d3c6eb10 against merge-base e01655f4; fork head, no PR code executed; pre-commit and DCO green; no ready label yet.

torch.testing.assert_close(value, flow_state[name], rtol=0, atol=0)


def test_fold_under_inference_mode_stays_usable_outside_it():

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

[suggestion] This regression exercises only the modern parametrization API, and the legacy branch returns before the inference-tensor normalization logic. Using the weight_norm_api fixture here and asserting explicitly that the folded weights are not inference tensors would cover cross-context behaviour for both supported APIs.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

You're right that the legacy branch never reaches the normalization. I ran both APIs through the same fold under inference_mode() (torch 2.13, CPU) to see what each leaves behind:

API normalization folded weight
parametrized none Tensor, is_inference=True
parametrized current Parameter, requires_grad=False, is_inference=False
legacy none Parameter, requires_grad=True, is_inference=True
legacy with the patch Parameter, requires_grad=False, is_inference=False

So the legacy path is worse than just un-normalized: it hands back a grad-tracking inference tensor.

Patch below. _normalize_folded_weight is the modern path's existing block factored out and called from both branches, and the test takes the weight_norm_api fixture and asserts not weight.is_inference() on all six convs. It is a patch rather than a commit because this branch lives in zhongc1211's fork, which I cannot push to; the mechanism above is measured, but I did not run the repo's own test file against it.

--- hifigan_orig.py	2026-09-22 09:31:47
+++ hifigan_new.py	2026-09-22 09:31:47
@@ -41,6 +41,24 @@
 
 def get_padding(kernel_size, dilation=1):
     return int((kernel_size * dilation - dilation) / 2)
+
+
+def _normalize_folded_weight(module: nn.Module, name: str = "weight") -> None:
+    """Leave a folded weight as a frozen Parameter that outlives the fold's grad context.
+
+    Both APIs leave something else behind when the fold runs under
+    ``no_grad``/``inference_mode``: the parametrization API registers a plain
+    (possibly inference) tensor, and the legacy hook registers a Parameter whose
+    data is an inference tensor and still requires grad. Rebuild it outside
+    inference_mode so callers in any grad mode get the same semantics.
+    """
+    weight = getattr(module, name)
+    if isinstance(weight, nn.Parameter) and not weight.is_inference():
+        return
+    with torch.inference_mode(False):
+        frozen = nn.Parameter(weight.detach().clone(), requires_grad=False)
+    delattr(module, name)
+    module.register_parameter(name, frozen)
 
 
 def _fold_weight_norm(module: nn.Module) -> int:
@@ -55,20 +73,12 @@
         # Removes every parametrization registered on "weight"; only
         # weight_norm is applied to the layers in this module.
         parametrize.remove_parametrizations(module, "weight", leave_parametrized=True)
-        weight = module.weight
-        if not isinstance(weight, nn.Parameter):
-            # Under no_grad/inference_mode, leave_parametrized registers a plain
-            # (possibly inference) tensor instead of a Parameter. Normalize so
-            # callers in any grad mode get the same semantics and the weight
-            # stays usable outside the fold's grad context.
-            with torch.inference_mode(False):
-                frozen = nn.Parameter(weight.detach().clone(), requires_grad=False)
-            del module.weight
-            module.register_parameter("weight", frozen)
+        _normalize_folded_weight(module)
         return 1
     for hook in list(module._forward_pre_hooks.values()):
         if isinstance(hook, WeightNorm):
             remove_legacy_weight_norm(module, name=hook.name)
+            _normalize_folded_weight(module, hook.name)
             return 1
     return 0
 
--- cosy_orig.py	2026-09-22 09:31:47
+++ cosy_new.py	2026-09-22 09:31:47
@@ -165,16 +165,24 @@
         torch.testing.assert_close(value, flow_state[name], rtol=0, atol=0)
 
 
-def test_fold_under_inference_mode_stays_usable_outside_it():
+def test_fold_under_inference_mode_stays_usable_outside_it(weight_norm_api):
     """Folding inside inference_mode must not leak inference tensors: the
     materialized weight stays a frozen Parameter and forward keeps working
-    in a normal/no_grad context (pre-hardening this raised RuntimeError)."""
+    in a normal/no_grad context (pre-hardening this raised RuntimeError).
+
+    Both APIs run here on purpose. They leave different things behind under
+    inference_mode - the parametrization API a plain (possibly inference) tensor,
+    the legacy hook a Parameter wrapping an inference tensor - so each needs its
+    own normalization before the weight is usable outside the fold's grad
+    context."""
     block = hifigan.ResBlock(channels=4).eval()
     x = torch.randn(1, 4, 12)
     with torch.inference_mode():
         assert block.remove_weight_norm() == 6
-    assert isinstance(block.convs1[0].weight, nn.Parameter)
-    assert not block.convs1[0].weight.requires_grad
+    for conv in (*block.convs1, *block.convs2):
+        assert isinstance(conv.weight, nn.Parameter)
+        assert not conv.weight.requires_grad
+        assert not conv.weight.is_inference()
     with torch.no_grad():
         out = block(x)
     assert torch.isfinite(out).all()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Correction and a real run: the diff I posted had placeholder file labels, and I have since checked this out (d3c6eb1) and run the file itself.

  • new test against the unpatched source: test_fold_under_inference_mode_stays_usable_outside_it[legacy] FAILED, 16 passed - the legacy variant is exactly the hole you pointed at.
  • new test against the patched source: 17 passed.

Re-posted with proper paths:

--- a/vllm_omni/model_executor/models/cosyvoice3/code2wav_core/hifigan.py
+++ b/vllm_omni/model_executor/models/cosyvoice3/code2wav_core/hifigan.py
@@ -41,6 +41,24 @@
 
 def get_padding(kernel_size, dilation=1):
     return int((kernel_size * dilation - dilation) / 2)
+
+
+def _normalize_folded_weight(module: nn.Module, name: str = "weight") -> None:
+    """Leave a folded weight as a frozen Parameter that outlives the fold's grad context.
+
+    Both APIs leave something else behind when the fold runs under
+    ``no_grad``/``inference_mode``: the parametrization API registers a plain
+    (possibly inference) tensor, and the legacy hook registers a Parameter whose
+    data is an inference tensor and still requires grad. Rebuild it outside
+    inference_mode so callers in any grad mode get the same semantics.
+    """
+    weight = getattr(module, name)
+    if isinstance(weight, nn.Parameter) and not weight.is_inference():
+        return
+    with torch.inference_mode(False):
+        frozen = nn.Parameter(weight.detach().clone(), requires_grad=False)
+    delattr(module, name)
+    module.register_parameter(name, frozen)
 
 
 def _fold_weight_norm(module: nn.Module) -> int:
@@ -55,20 +73,12 @@
         # Removes every parametrization registered on "weight"; only
         # weight_norm is applied to the layers in this module.
         parametrize.remove_parametrizations(module, "weight", leave_parametrized=True)
-        weight = module.weight
-        if not isinstance(weight, nn.Parameter):
-            # Under no_grad/inference_mode, leave_parametrized registers a plain
-            # (possibly inference) tensor instead of a Parameter. Normalize so
-            # callers in any grad mode get the same semantics and the weight
-            # stays usable outside the fold's grad context.
-            with torch.inference_mode(False):
-                frozen = nn.Parameter(weight.detach().clone(), requires_grad=False)
-            del module.weight
-            module.register_parameter("weight", frozen)
+        _normalize_folded_weight(module)
         return 1
     for hook in list(module._forward_pre_hooks.values()):
         if isinstance(hook, WeightNorm):
             remove_legacy_weight_norm(module, name=hook.name)
+            _normalize_folded_weight(module, hook.name)
             return 1
     return 0
 
--- a/tests/model_executor/models/cosyvoice3/test_cosyvoice3_hift_weight_norm.py
+++ b/tests/model_executor/models/cosyvoice3/test_cosyvoice3_hift_weight_norm.py
@@ -165,16 +165,24 @@
         torch.testing.assert_close(value, flow_state[name], rtol=0, atol=0)
 
 
-def test_fold_under_inference_mode_stays_usable_outside_it():
+def test_fold_under_inference_mode_stays_usable_outside_it(weight_norm_api):
     """Folding inside inference_mode must not leak inference tensors: the
     materialized weight stays a frozen Parameter and forward keeps working
-    in a normal/no_grad context (pre-hardening this raised RuntimeError)."""
+    in a normal/no_grad context (pre-hardening this raised RuntimeError).
+
+    Both APIs run here on purpose. They leave different things behind under
+    inference_mode - the parametrization API a plain (possibly inference) tensor,
+    the legacy hook a Parameter wrapping an inference tensor - so each needs its
+    own normalization before the weight is usable outside the fold's grad
+    context."""
     block = hifigan.ResBlock(channels=4).eval()
     x = torch.randn(1, 4, 12)
     with torch.inference_mode():
         assert block.remove_weight_norm() == 6
-    assert isinstance(block.convs1[0].weight, nn.Parameter)
-    assert not block.convs1[0].weight.requires_grad
+    for conv in (*block.convs1, *block.convs2):
+        assert isinstance(conv.weight, nn.Parameter)
+        assert not conv.weight.requires_grad
+        assert not conv.weight.is_inference()
     with torch.no_grad():
         out = block(x)
     assert torch.isfinite(out).all()

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

the legacy branch returns before the inference-tensor normalization logic.

Addressed in 2e021d07, using @0z5a's shared normalizer for both APIs. Ordinary non-inference Parameters keep their existing flags; plain or inference tensors are rebuilt outside inference_mode.

The regression now uses weight_norm_api across grad, no_grad and inference_mode. All six combinations check every folded convolution for not weight.is_inference() and complete forward/backward after leaving the fold context. There is also cold-loader coverage for an actual legacy-created vocoder, with no forward before loading.

The same source bytes passed all 26 focused CPU cases and the applicable pre-commit checks before this commit. Thanks @linyueqian for spotting the gap and @0z5a for the tested fix.

@0z5a

0z5a commented Sep 23, 2026

Copy link
Copy Markdown
Contributor

@timzsu @zhongc1211 I realized that my draft #7827 overlaps with this C4 work. I started it without confirming coordination on the workstream first; sorry for the duplicated effort.

Having compared the two PRs, I support moving C4 forward here in #7871. This PR keeps the five F0 predictor parametrizations for the separate C5/#7518 work, and I have already shared a tested fix for the legacy weight norm path under inference_mode in the review thread.

I will keep #7827 as a draft for now while we hear from you and settle the scope. @timzsu, does proceeding with #7871 for C4 and leaving F0 to C5 match your intent? If so, I can close #7827 after the coordination is clear. The #6870 C4 status row currently links #7871 but attributes it to me; that should credit @zhongc1211.

@timzsu

timzsu commented Sep 23, 2026

Copy link
Copy Markdown
Contributor

Hi, I have changed the RFC to correctly credit @zhongc1211 for this PR. I think it will be nice if you guys can collaborate on this feature.

Apply the shared normalizer suggested by @0z5a to both weight-norm
APIs so folding under inference_mode cannot leave inference Parameters.
Preserve existing ordinary Parameter flags.

Cover both APIs across grad contexts, cold loading into a legacy-created
vocoder, the loader's transition to eval, and multichunk finalization.
Keep the generator-only 77-layer scope and five F0 norms unchanged.

Related: vllm-project#7871
Signed-off-by: zack <huixindaddy@yahoo.com>
@zhongc1211 zhongc1211 changed the title [Perf][CosyVoice3] Fold frozen HiFT generator weight norms after checkpoint load [Model] CosyVoice3: fold HiFT generator weight norms after loading Sep 23, 2026
@zhongc1211

Copy link
Copy Markdown
Contributor Author

I think it will be nice if you guys can collaborate on this feature.

Thanks @timzsu for correcting the RFC attribution, and @0z5a for the tested legacy-path fix and for supporting #7871. That fix is now in 2e021d07, with expanded regressions.

The scope remains 77 generator norms, leaving the five F0 norms to C5/#7518. @0z5a, please flag anything else from #7827 that we should carry over.

@linyueqian, could you take another look at this head with that scope?

Materialize F0 weights on CPU in FP32 using the shared folding helpers,
while keeping generator and F0 folding as separate entry points.

Add joint C4/C5 loading and streaming regressions for modern and legacy
weight normalization, including strict loading and idempotence.

Signed-off-by: zack <huixindaddy@yahoo.com>
@vllm-omni-review-bot

Copy link
Copy Markdown

Omni ReviewBot triage note

Automated triage of commit 5c7a151f7b05 produced:

  • Priority: high. Prompt maintainer attention is suggested.

These are automated triage suggestions only — the final decision belongs to the maintainers.

@zhongc1211 zhongc1211 changed the title [Model] CosyVoice3: fold HiFT generator weight norms after loading [Model] CosyVoice3: fold generator and F0 weight norms after loading Sep 25, 2026
@zhongc1211

Copy link
Copy Markdown
Contributor Author

Two PRs for one workstream item should not both land

#7827 is now closed without merging. @0z5a's shared normalizer is in 2e021d07, with the legacy regressions @linyueqian requested.

The scope has changed since my previous update: 5c7a151f7b05 adds the five F0 folds here, alongside the 77 generator folds. F0 stays on CPU in FP32 and is moved there before folding. This does not include the C2 backend routing from #7518.

@timzsu @BeatSeat, please flag any conflict with carrying C5 here while keeping C2 separate. A later integration still needs to reconcile F0 placement and dtype before folding. @linyueqian, could you review the updated head, including the new F0/joint regression tests?

@BeatSeat

Copy link
Copy Markdown
Contributor

Thanks @zhongc1211! No conflict at all — fully support carrying C5 in #7871 alongside C4, while keeping #7518 strictly focused on C2 (GPU-resident F0 + pure-PyTorch HiFT). Once #7518 lands, #7871 can easily rebase on top. Thanks for driving this forward!

@Sy0307

Sy0307 commented Sep 27, 2026

Copy link
Copy Markdown
Collaborator

Testing it and it works well. Maybe we can merge it ASAP? Thx for nice work again @zhongc1211

@vllm-omni-review-bot

Copy link
Copy Markdown

Omni ReviewBot: no human activity for 7 days

@zhongc1211 this pull request has had no human commit, comment or review since 2026-09-27. Please confirm the current plan and next step. The author or a maintainer decides whether to change the PR state.

To keep it moving, any one of these is enough: push an update, reply to the open blocker, or post the current plan and timeline.

@BeatSeat

BeatSeat commented Oct 4, 2026

Copy link
Copy Markdown
Contributor

Should it be merged ASAP? @linyueqian

@zhongc1211 pls check the conflicts and rebase

@zhongc1211

Copy link
Copy Markdown
Contributor Author

Thanks @BeatSeat. This PR’s generator and F0 folding changes are already in main through #8224. Both regression test files are included too. #8224 also updated F0 folding to use the inference device.

So there doesn’t seem to be anything left to rebase. @Sy0307 @linyueqian, should we close #7871 as included in #8224, or is anything still missing?

@linyueqian

Copy link
Copy Markdown
Collaborator

Thanks for the implementation and regression coverage. This was included in #8224. I checked the folding helpers and both regression files against the merge and current main, and the folding code is retained. The F0 loader and two test assertions were adapted to fold on the inference device, and the tests later gained CPU thread limits. I didn't find any remaining work from this PR to rebase, so we can close it as incorporated in #8224.

@zhongc1211 zhongc1211 closed this Oct 7, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request

Projects

None yet

Development

Successfully merging this pull request may close these issues.

8 participants