Corrective Forcing

Corrective Forcing (CoF) is a unified post-training paradigm for diffusion- and flow-based generative speech enhancement that adapts pretrained models to self-generated rollouts and corrects their predictions at both rollout-state and local-transition levels.

Demo GitHub arXiv Hugging Face ModelScope


This repository hosts the official checkpoints of Corrective Forcing: Unified Post-Training for Diffusions and Flows in Generative Speech Enhancement. [Code is available here]

Corrective Forcing (CoF) is a unified post-training framework for diffusion- and flow-based generative speech enhancement. It addresses the training–inference mismatch between analytical path states used during pretraining and self-generated rollout states encountered during inference, where prediction and discretization errors can accumulate along discretized sampling trajectories. CoF adapts pretrained models directly to their own rollouts with corrective supervision at both the rollout-state and local-transition levels. A shared clean-speech prediction parameterization enables the same post-training objective to be applied across diffusion and flow formulations.

CoF consists of two complementary components:

  • Dynamic Rollout Correction (DRC): trains the model on self-generated rollout states under dynamic sampling schedules, correcting clean-speech predictions toward the ground truth while exposing the model to diverse inference discretizations.
  • Counterfactual Transition Consistency (CTC): regularizes local evolution under transition-induced state deviations by aligning predictions after factual and locally corrected counterfactual transitions constructed from the same rollout state.

We validate CoF on two representative generative speech enhancement formulations, covering Schrödinger bridge and flow matching:

  • SB-VE: a Schrödinger bridge formulation that models stochastic probability transport between degraded and clean speech using a variance-exploding diffusion process.
  • OT-CFM: an optimal-transport conditional flow-matching formulation that learns the vector field associated with a linear probability path, enabling deterministic transport from degraded toward clean speech.

Dataset

The checkpoints are trained on the following datasets:

Dataset Construction Training set Validation set Test set
Voicebank+Demand This dataset is a widely used, publicly available benchmark for single-channel speech denoising. Clean utterances are drawn from the VoiceBank corpus. Noise recordings consist of eight DEMAND noises and two synthetic noises. Each clean utterance is mixed with one of the ten noise recordings at SNRs of 15, 10, 5, and 0 dB for training and 17.5, 12.5, 7.5, and 2.5 dB for testing 10,802 utterances (≈9 h) from 28 speakers 770 utterances 824 utterances from 2 unseen speakers
WSJ0+WHAM This dataset is constructed for speech denoising under realistic ambient noise. Clean utterances are drawn from the WSJ0 corpus. Noise recordings are taken from the WHAM! corpus of real-world ambient noise. Each clean utterance is mixed with one of the noise recordings at an SNR uniformly drawn from −6 dB to 14 dB 25,550 mixtures (≈50 h) from si_tr_s 1,302 mixtures from si_et_05 2,412 mixtures from si_dt_05
WSJ0+Reverb This dataset is constructed for speech dereverberation. Clean utterances are drawn from the WSJ0 corpus, following the same data split as WSJ0+WHAM. Reverberant utterances are generated by convolving clean speech with synthetic room impulse responses simulated via the image-source method, with T60 uniformly sampled from [0.4, 1.0] s and room dimensions sampled from [5, 15] × [5, 15] × [2, 6] m. The corresponding clean targets are regenerated in a dry room with the same geometry and a fixed absorption coefficient of 0.99 same data split as WSJ0+WHAM si_et_05 si_dt_05

All audio is resampled to 16 kHz.

Checkpoint List

Each checkpoint lives under the path <dataset>/<formulation>/<stage>/ and consists of a self-describing weight file named <dataset>_<formulation>[_cof].safetensors (for example voicebank_sbve_cof.safetensors) plus a config.yml whose weights.default_test_model points at it. The pretraining stage (stage 1) is trained with teacher forcing, while the CoF stage (stage 2) is the post-trained result exported at the best-PESQ checkpoint.

The sampling solver is named per formulation and recorded both in each config.yml (formulation.sampling.solver) and in the metadata embedded in the weight file: SB-VE uses SB_SDE_Solver (canonical) and SB_ODE_Solver (probability flow), OT-CFM uses OTCFM_ODE_Solver. The solver recorded for each CoF checkpoint is the solver its rollouts were generated with and the recommended inference solver.

Dataset Formulation Stage Path SHA256
Voicebank+Demand SB-VE Stage 1 (pre-training) voicebank+demand/sb-ve/pretraining 9acd9bf9787dd3d32181e5c38078477b4e6440df67240908e34a7eec1f59dc71
Stage 2 (CoF) voicebank+demand/sb-ve/CoF dea4bf38b264dc299dd3ebd67745651804c91316d7e8e521570494f0246920b3
OT-CFM Stage 1 (pre-training) voicebank+demand/ot-cfm/pretraining ce50c32c5c973f4cf4703098fd584ca30cd94496711b72b015d772429437fd24
Stage 2 (CoF) voicebank+demand/ot-cfm/CoF 10a282f026dc324cc6385a68fed3b9326776e10fe663afe9627e84410590dab1
WSJ0+Reverb SB-VE Stage 1 (pre-training) wsj0+reverb/sb-ve/pretraining 5ba378e43a0d0c2ec4067d6240bfe7b60a48ee3c11221656a1132af0b39a65b1
Stage 2 (CoF) wsj0+reverb/sb-ve/CoF 5af5bd7259716a053676cc4180b5ace9815b5eff5010bcf2f1a6d42e306920d3
OT-CFM Stage 1 (pre-training) wsj0+reverb/ot-cfm/pretraining 94ed3e99b5f31bd8a711675ae6704f398f97a0f2f05049200aa4e259127a2ee4
Stage 2 (CoF) wsj0+reverb/ot-cfm/CoF 286905e37115a4e1853225ad74996b4ebfe174f8d28bb1b2b199883beda16b5b
WSJ0+WHAM SB-VE Stage 1 (pre-training) wsj0+wham/sb-ve/pretraining 5d243000f49bef4d0c9c054af9c39009ba6ba38a038fb7a1f377efca37ddac55
Stage 2 (CoF) wsj0+wham/sb-ve/CoF 41b9d432aabac4ffcb49caa1e8abb5bad6a36ccf0e7c2ab2161bcd5c9a2b37dc
OT-CFM Stage 1 (pre-training) wsj0+wham/ot-cfm/pretraining c10e0205deac5176a3f66e157ebe00b59153f9a8f59996475ebadae3bcc8cd45
Stage 2 (CoF) wsj0+wham/ot-cfm/CoF 0daf6370744cdb9b6c9488bfb4f0e217a2793c629090f5fbf8bdd1d164d837b9

Citation

If this project helps your research, please consider citing:

@misc{yao2026corrective,
  title={Corrective Forcing: Unified Post-Training for Diffusions and Flows in Generative Speech Enhancement},
  author={Yao, Qing and Gao, Lijian and Mao, Qirong},
  journal={arXiv preprint arXiv:2609.24651},
  year={2026}
}

License

Apache License 2.0.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Paper for Yorch233/CoF