From 25f709575ebf7632dd6c3f2bcedecef067e848c9 Mon Sep 17 00:00:00 2001 From: aryan5v Date: Thu, 17 Sep 2026 14:01:44 -0700 Subject: [PATCH 1/5] [feat]: consolidate CompactH3 training pipeline --- docs/quantization/h3_int8_affine.md | 290 +++ docs/quantization/h3_nvfp4.md | 443 +++++ docs/quantization/h3_w4a16.md | 366 ++++ docs/quantization/loader_quant_params.md | 87 + docs/training/attn_qat.md | 2 +- examples/compacth3/README.md | 54 + examples/distill/MiniMax-H3/distill_dmd.sh | 28 + examples/inference/minimax_h3/README.md | 63 + examples/inference/minimax_h3/h3_vsa_dmd.py | 624 +++++++ .../compacth3/release14b_recovery_wandb.yaml | 133 ++ .../compacth3/release14b_validation_five.json | 44 + .../minimax_h3/README.md | 45 + .../minimax_h3/qad_nvfp4_4call.yaml | 140 ++ .../minimax_h3/release20b_dmd2_v12_dense.yaml | 145 ++ .../release20b_validation_five.json | 42 + .../train/configs/fasth3_14b_recovery.yaml | 92 + .../train/configs/fasth3_base42_recovery.yaml | 91 + .../configs/fasth3_detail_band_recovery.yaml | 96 + .../fasth3_detail_band_recovery34.yaml | 101 ++ .../configs/fasth3_release_long_recovery.yaml | 91 + .../configs/overfit_minimax_h3_t2va.yaml | 4 +- .../comparison_five_prompts.json | 32 + .../quick_gate_prompts.json | 74 + .../release_sentinel_24.json | 26 + .../rescue_gate_prompts.json | 26 + .../fasth3_14b_2step_qad/showcase_prompt.json | 10 + fastvideo/attention/utils/flash_attn_cute.py | 146 +- fastvideo/configs/models/dits/minimax_h3.py | 3 +- .../dataset/parquet_dataset_map_style.py | 191 +- fastvideo/dataset/shape_bucket.py | 49 + fastvideo/dataset/validation_dataset.py | 11 +- fastvideo/fastvideo_args.py | 8 +- fastvideo/layers/quantization/__init__.py | 8 +- .../layers/quantization/int8_affine_config.py | 988 ++++++++++ fastvideo/layers/quantization/nvfp4_config.py | 581 +++++- .../layers/quantization/nvfp4_qat_config.py | 4 + fastvideo/layers/quantization/w4a16_config.py | 1099 ++++++++++++ fastvideo/models/dits/minimax_h3.py | 34 +- fastvideo/models/loader/component_loader.py | 40 +- fastvideo/models/loader/fsdp_load.py | 366 +++- fastvideo/models/loader/shard_cache.py | 373 ++++ .../schedulers/scheduling_minimax_h3.py | 7 +- .../minimax_h3/stages/minimax_h3_decoding.py | 16 +- .../minimax_h3/stages/minimax_h3_denoising.py | 92 +- fastvideo/pipelines/pipeline_batch_info.py | 4 + .../preprocess_minimax_h3_overfit.py | 49 +- .../preprocess_minimax_h3_text_only.py | 213 +++ .../tests/attention/test_compile_policy.py | 19 + .../test_flash_attn_cute_custom_op.py | 31 + .../test_vsa_h3_inference_metadata_parity.py | 177 ++ .../tests/attention/test_vsa_h3_metadata.py | 24 + .../test_exact_shape_bucket_sampler.py | 115 ++ .../dataset/test_parquet_dataset_map_style.py | 74 +- .../tests/dataset/test_validation_dataset.py | 33 + fastvideo/tests/loader/test_shard_cache.py | 220 +++ .../ops/quantization/test_allowlist_mirror.py | 43 + .../quantization/test_int8_affine_config.py | 339 ++++ .../ops/quantization/test_int8_dispatch.py | 70 + .../quantization/test_nvfp4_h3_prefixes.py | 161 ++ .../ops/quantization/test_nvfp4_sidecar.py | 334 ++++ .../test_quant_param_allowlist.py | 139 ++ .../test_quant_sidecar_roundtrip.py | 697 ++++++++ .../ops/quantization/test_w4a16_config.py | 462 +++++ .../tests/train/callbacks/test_callback.py | 1 + fastvideo/tests/train/callbacks/test_ema.py | 25 + .../test_latent_vis_shape_context.py | 50 + .../tests/train/callbacks/test_validation.py | 211 ++- .../test_validation_sampling_contract.py | 95 + .../train/fixtures/minimax_h3_dmd2_min.yaml | 61 + .../train/methods/test_dmd2_data_forcing.py | 463 +++++ .../test_dmd2_fake_score_loss_space.py | 136 ++ .../train/methods/test_dmd2_fastgen_parity.py | 229 +++ .../train/methods/test_dmd2_rollout_carry.py | 656 +++++++ .../methods/test_dmd2_timestep_bounds.py | 64 +- .../train/methods/test_dmd2_vsd_normalizer.py | 47 + .../train/methods/test_minimax_h3_dmd2.py | 1154 ++++++++++++ .../train/methods/test_minimax_h3_finetune.py | 50 +- .../tests/train/trainer/test_validation.py | 31 + .../tests/train/utils/test_checkpoint.py | 421 ++++- fastvideo/tests/train/utils/test_config.py | 65 +- .../train/utils/test_inference_checkpoint.py | 478 +++++ .../test_inference_checkpoint_distributed.py | 286 +++ .../test_moduleloader_attention_backend.py | 187 +- .../tests/train/utils/test_torch_compile.py | 246 +++ fastvideo/tests/train/utils/test_tracking.py | 136 ++ fastvideo/train/attn_qat/README.md | 2 +- fastvideo/train/callbacks/callback.py | 15 +- fastvideo/train/callbacks/ema.py | 3 + fastvideo/train/callbacks/latent_vis.py | 93 + fastvideo/train/callbacks/validation.py | 348 +++- .../train/entrypoint/dcp_to_diffusers.py | 100 +- fastvideo/train/entrypoint/train.py | 15 +- fastvideo/train/methods/base.py | 50 +- .../methods/distribution_matching/dmd2.py | 1222 ++++++++++++- fastvideo/train/models/base.py | 16 + fastvideo/train/models/minimax_h3/__init__.py | 4 + .../train/models/minimax_h3/minimax_h3.py | 284 ++- .../train/models/minimax_h3/minimax_h3_dmd.py | 522 ++++++ fastvideo/train/trainer.py | 56 + fastvideo/train/utils/checkpoint.py | 320 +++- fastvideo/train/utils/config.py | 37 + fastvideo/train/utils/dataloader.py | 3 +- fastvideo/train/utils/inference_checkpoint.py | 627 +++++++ fastvideo/train/utils/moduleloader.py | 66 +- fastvideo/train/utils/optimizer.py | 94 +- fastvideo/train/utils/tracking.py | 70 +- fastvideo/train/utils/training_config.py | 34 + mkdocs.yml | 5 + .../convert_minimax_h3_adaln_rank.py | 44 +- .../export_h3_dmd2_student.py | 202 +++ .../analysis/adaln/analyze_adaln_rank.py | 423 +++++ .../analysis/adaln/compare_parent_dmd2.py | 67 + .../analysis/adaln/run_adaln_rank.sh | 59 + .../analysis/adaln/summarize_adaln_rank.py | 56 + scripts/compacth3/analysis/adaln_lowrank.py | 1047 +++++++++++ .../analysis/checkpoint_sweep_metrics.py | 296 +++ ...val_corrected_dmd_inference_exports.sbatch | 133 ++ scripts/compacth3/eval_dmd_export_1400.sh | 133 ++ .../grade_dmd2_all_checkpoints.sbatch | 49 + scripts/compacth3/qad/gen_qad_configs.py | 124 ++ scripts/compacth3/qad/qad_checkpoint_gate.py | 162 ++ scripts/compacth3/qad/setup_qad.py | 147 ++ scripts/compacth3/qad/tune_qad_cadence.py | 30 + .../quantization/export_lane_int8.sh | 26 + .../quantization/export_lane_nvfp4.sh | 26 + .../quantization/export_lane_w4a16.sh | 26 + .../quantization/export_quant_dit.py | 213 +++ .../quantization/run_export_nvfp4.sh | 23 + .../resume_release20b_dmd2_paired_generic.sh | 70 + scripts/compacth3/run_eval_lane.sh | 89 + .../submit_release14b_hardened.sbatch | 102 ++ scripts/compacth3/sweep.sh | 62 + scripts/compacth3/sweep_prompts.json | 23 + .../admit_h3_fold_for_recovery.py | 63 + .../aggregate_minimax_h3_block_scores.py | 219 +++ scripts/fasth3_sprint/assemble_h3_stage.py | 24 + scripts/fasth3_sprint/audio_fidelity_gate.py | 225 +++ .../fasth3_sprint/audit_recovery_weights.py | 191 ++ .../build_h3_audio_stratified_prompt_index.py | 106 ++ .../fasth3_sprint/build_h3_recovery_config.py | 26 + scripts/fasth3_sprint/check_h3_base_gate.py | 17 + .../fasth3_sprint/extract_review_assets.py | 134 ++ .../fasth3_sprint/h3_serve_eval_prompts.json | 82 + .../h3_serve_eval_prompts_cool_480p.json | 42 + ...rve_eval_prompts_gateway_fashion_480p.json | 18 + scripts/fasth3_sprint/h3_stage_promotion.py | 43 + .../h3_taeh3_kitchen_prompt.json | 20 + .../fasth3_sprint/prepare_h3_prompt_pool.py | 195 ++ scripts/fasth3_sprint/prepare_h3_stage34.py | 54 + scripts/fasth3_sprint/run_baseline_matrix.py | 397 +++++ .../fasth3_sprint/run_recovery_comparison.sh | 188 ++ .../fasth3_sprint/run_recovery_diagnostics.sh | 86 + .../fasth3_sprint/score_minimax_h3_blocks.py | 486 +++++ scripts/fasth3_sprint/seams_for_block_map.py | 44 + scripts/fasth3_sprint/select_h3_block_map.py | 216 +++ .../fasth3_sprint/slurm_block_score.sbatch | 62 + .../slurm_block_score_aggregate.sbatch | 43 + .../slurm_h3_34_latest_five.sbatch | 52 + ..._h3_activation34_from42_recovery500.sbatch | 74 + .../slurm_h3_activation42_five.sbatch | 42 + .../slurm_h3_base42_recovery500.sbatch | 56 + .../slurm_h3_release_eval.sbatch | 96 + .../slurm_h3_release_long_phase.sbatch | 113 ++ .../slurm_prune_candidate.sbatch | 111 ++ scripts/fasth3_sprint/slurm_tests.sbatch | 53 + scripts/fasth3_sprint/split_h3_prompt_pool.py | 54 + scripts/fasth3_sprint/submit_block_score.sh | 34 + .../fasth3_sprint/submit_pruned_candidates.sh | 23 + .../fasth3_sprint/validate_h3_candidate.py | 221 +++ scripts/fasth3_sprint/verify_media.py | 28 + ...erify_recovery_export_prediction_parity.py | 392 ++++ scripts/fasth3_sprint/verify_speech_asr.py | 160 ++ .../minimax_h3_native_t2va/README.md | 170 ++ .../derive_filtered_dataset.py | 549 ++++++ .../derive_filtered_dataset_1tray.sbatch | 35 + .../encode_native_t2va_1tray.sbatch | 61 + .../minimax_h3_native_t2va/encode_worker.py | 560 ++++++ .../finalize_dataset.py | 612 +++++++ .../finalize_extension_1tray.sbatch | 24 + .../minimax_h3_native_t2va/freeze_sources.py | 1583 +++++++++++++++++ .../prepare_extension_1tray.sbatch | 41 + .../probe_native_t2va.sbatch | 36 + .../minimax_h3_native_t2va/schema_dry_run.py | 71 + .../seed_encoded_chunks.py | 429 +++++ .../minimax_h3_native_t2va/v10_sources.json | 65 + .../v10_sources_v2.json | 65 + .../validate_heldout_media.py | 80 + scripts/run_release20b_dmd2_v12_16gpu.sh | 208 +++ ...release20b_dmd2_v12_corrected_16gpu.sbatch | 50 + ...release20b_dmd2_v12_corrected_32gpu.sbatch | 43 + scripts/train/mfu_calc_minimax_h3.py | 302 ++++ tests/local_tests/minimax_h3/PORT_STATUS.md | 3 + tests/local_tests/minimax_h3/README.md | 3 + .../models/test_fsdp_load_mixed_dtype.py | 149 ++ .../preprocess/test_minimax_h3_native_t2va.py | 881 +++++++++ .../test_minimax_h3_native_t2va_filtered.py | 239 +++ 196 files changed, 33363 insertions(+), 465 deletions(-) create mode 100644 docs/quantization/h3_int8_affine.md create mode 100644 docs/quantization/h3_nvfp4.md create mode 100644 docs/quantization/h3_w4a16.md create mode 100644 docs/quantization/loader_quant_params.md create mode 100644 examples/compacth3/README.md create mode 100755 examples/distill/MiniMax-H3/distill_dmd.sh create mode 100644 examples/inference/minimax_h3/README.md create mode 100644 examples/inference/minimax_h3/h3_vsa_dmd.py create mode 100644 examples/train/configs/compacth3/release14b_recovery_wandb.yaml create mode 100644 examples/train/configs/compacth3/release14b_validation_five.json create mode 100644 examples/train/configs/distribution_matching/minimax_h3/README.md create mode 100644 examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call.yaml create mode 100644 examples/train/configs/distribution_matching/minimax_h3/release20b_dmd2_v12_dense.yaml create mode 100644 examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json create mode 100644 examples/train/configs/fasth3_14b_recovery.yaml create mode 100644 examples/train/configs/fasth3_base42_recovery.yaml create mode 100644 examples/train/configs/fasth3_detail_band_recovery.yaml create mode 100644 examples/train/configs/fasth3_detail_band_recovery34.yaml create mode 100644 examples/train/configs/fasth3_release_long_recovery.yaml create mode 100644 examples/training/fasth3_14b_2step_qad/comparison_five_prompts.json create mode 100644 examples/training/fasth3_14b_2step_qad/quick_gate_prompts.json create mode 100644 examples/training/fasth3_14b_2step_qad/release_sentinel_24.json create mode 100644 examples/training/fasth3_14b_2step_qad/rescue_gate_prompts.json create mode 100644 examples/training/fasth3_14b_2step_qad/showcase_prompt.json create mode 100644 fastvideo/dataset/shape_bucket.py create mode 100644 fastvideo/layers/quantization/int8_affine_config.py create mode 100644 fastvideo/layers/quantization/w4a16_config.py create mode 100644 fastvideo/models/loader/shard_cache.py create mode 100644 fastvideo/pipelines/preprocess/preprocess_minimax_h3_text_only.py create mode 100644 fastvideo/tests/attention/test_compile_policy.py create mode 100644 fastvideo/tests/attention/test_vsa_h3_inference_metadata_parity.py create mode 100644 fastvideo/tests/dataset/test_exact_shape_bucket_sampler.py create mode 100644 fastvideo/tests/dataset/test_validation_dataset.py create mode 100644 fastvideo/tests/loader/test_shard_cache.py create mode 100644 fastvideo/tests/ops/quantization/test_allowlist_mirror.py create mode 100644 fastvideo/tests/ops/quantization/test_int8_affine_config.py create mode 100644 fastvideo/tests/ops/quantization/test_int8_dispatch.py create mode 100644 fastvideo/tests/ops/quantization/test_nvfp4_h3_prefixes.py create mode 100644 fastvideo/tests/ops/quantization/test_nvfp4_sidecar.py create mode 100644 fastvideo/tests/ops/quantization/test_quant_param_allowlist.py create mode 100644 fastvideo/tests/ops/quantization/test_quant_sidecar_roundtrip.py create mode 100644 fastvideo/tests/ops/quantization/test_w4a16_config.py create mode 100644 fastvideo/tests/train/callbacks/test_latent_vis_shape_context.py create mode 100644 fastvideo/tests/train/callbacks/test_validation_sampling_contract.py create mode 100644 fastvideo/tests/train/fixtures/minimax_h3_dmd2_min.yaml create mode 100644 fastvideo/tests/train/methods/test_dmd2_data_forcing.py create mode 100644 fastvideo/tests/train/methods/test_dmd2_fake_score_loss_space.py create mode 100644 fastvideo/tests/train/methods/test_dmd2_fastgen_parity.py create mode 100644 fastvideo/tests/train/methods/test_dmd2_rollout_carry.py create mode 100644 fastvideo/tests/train/methods/test_dmd2_vsd_normalizer.py create mode 100644 fastvideo/tests/train/methods/test_minimax_h3_dmd2.py create mode 100644 fastvideo/tests/train/utils/test_inference_checkpoint.py create mode 100644 fastvideo/tests/train/utils/test_inference_checkpoint_distributed.py create mode 100644 fastvideo/tests/train/utils/test_torch_compile.py create mode 100644 fastvideo/tests/train/utils/test_tracking.py create mode 100644 fastvideo/train/callbacks/latent_vis.py create mode 100644 fastvideo/train/models/minimax_h3/minimax_h3_dmd.py create mode 100644 fastvideo/train/utils/inference_checkpoint.py create mode 100644 scripts/checkpoint_conversion/export_h3_dmd2_student.py create mode 100644 scripts/compacth3/analysis/adaln/analyze_adaln_rank.py create mode 100644 scripts/compacth3/analysis/adaln/compare_parent_dmd2.py create mode 100644 scripts/compacth3/analysis/adaln/run_adaln_rank.sh create mode 100644 scripts/compacth3/analysis/adaln/summarize_adaln_rank.py create mode 100644 scripts/compacth3/analysis/adaln_lowrank.py create mode 100644 scripts/compacth3/analysis/checkpoint_sweep_metrics.py create mode 100644 scripts/compacth3/eval_corrected_dmd_inference_exports.sbatch create mode 100644 scripts/compacth3/eval_dmd_export_1400.sh create mode 100644 scripts/compacth3/grade_dmd2_all_checkpoints.sbatch create mode 100644 scripts/compacth3/qad/gen_qad_configs.py create mode 100644 scripts/compacth3/qad/qad_checkpoint_gate.py create mode 100644 scripts/compacth3/qad/setup_qad.py create mode 100644 scripts/compacth3/qad/tune_qad_cadence.py create mode 100644 scripts/compacth3/quantization/export_lane_int8.sh create mode 100644 scripts/compacth3/quantization/export_lane_nvfp4.sh create mode 100644 scripts/compacth3/quantization/export_lane_w4a16.sh create mode 100644 scripts/compacth3/quantization/export_quant_dit.py create mode 100644 scripts/compacth3/quantization/run_export_nvfp4.sh create mode 100644 scripts/compacth3/resume_release20b_dmd2_paired_generic.sh create mode 100644 scripts/compacth3/run_eval_lane.sh create mode 100644 scripts/compacth3/submit_release14b_hardened.sbatch create mode 100644 scripts/compacth3/sweep.sh create mode 100644 scripts/compacth3/sweep_prompts.json create mode 100644 scripts/fasth3_sprint/admit_h3_fold_for_recovery.py create mode 100644 scripts/fasth3_sprint/aggregate_minimax_h3_block_scores.py create mode 100644 scripts/fasth3_sprint/assemble_h3_stage.py create mode 100644 scripts/fasth3_sprint/audio_fidelity_gate.py create mode 100644 scripts/fasth3_sprint/audit_recovery_weights.py create mode 100644 scripts/fasth3_sprint/build_h3_audio_stratified_prompt_index.py create mode 100644 scripts/fasth3_sprint/build_h3_recovery_config.py create mode 100644 scripts/fasth3_sprint/check_h3_base_gate.py create mode 100644 scripts/fasth3_sprint/extract_review_assets.py create mode 100644 scripts/fasth3_sprint/h3_serve_eval_prompts.json create mode 100644 scripts/fasth3_sprint/h3_serve_eval_prompts_cool_480p.json create mode 100644 scripts/fasth3_sprint/h3_serve_eval_prompts_gateway_fashion_480p.json create mode 100644 scripts/fasth3_sprint/h3_stage_promotion.py create mode 100644 scripts/fasth3_sprint/h3_taeh3_kitchen_prompt.json create mode 100644 scripts/fasth3_sprint/prepare_h3_prompt_pool.py create mode 100644 scripts/fasth3_sprint/prepare_h3_stage34.py create mode 100644 scripts/fasth3_sprint/run_baseline_matrix.py create mode 100644 scripts/fasth3_sprint/run_recovery_comparison.sh create mode 100644 scripts/fasth3_sprint/run_recovery_diagnostics.sh create mode 100644 scripts/fasth3_sprint/score_minimax_h3_blocks.py create mode 100644 scripts/fasth3_sprint/seams_for_block_map.py create mode 100644 scripts/fasth3_sprint/select_h3_block_map.py create mode 100644 scripts/fasth3_sprint/slurm_block_score.sbatch create mode 100644 scripts/fasth3_sprint/slurm_block_score_aggregate.sbatch create mode 100644 scripts/fasth3_sprint/slurm_h3_34_latest_five.sbatch create mode 100644 scripts/fasth3_sprint/slurm_h3_activation34_from42_recovery500.sbatch create mode 100644 scripts/fasth3_sprint/slurm_h3_activation42_five.sbatch create mode 100644 scripts/fasth3_sprint/slurm_h3_base42_recovery500.sbatch create mode 100644 scripts/fasth3_sprint/slurm_h3_release_eval.sbatch create mode 100644 scripts/fasth3_sprint/slurm_h3_release_long_phase.sbatch create mode 100644 scripts/fasth3_sprint/slurm_prune_candidate.sbatch create mode 100644 scripts/fasth3_sprint/slurm_tests.sbatch create mode 100755 scripts/fasth3_sprint/split_h3_prompt_pool.py create mode 100644 scripts/fasth3_sprint/submit_block_score.sh create mode 100644 scripts/fasth3_sprint/submit_pruned_candidates.sh create mode 100644 scripts/fasth3_sprint/validate_h3_candidate.py create mode 100644 scripts/fasth3_sprint/verify_media.py create mode 100755 scripts/fasth3_sprint/verify_recovery_export_prediction_parity.py create mode 100644 scripts/fasth3_sprint/verify_speech_asr.py create mode 100644 scripts/preprocess/minimax_h3_native_t2va/README.md create mode 100644 scripts/preprocess/minimax_h3_native_t2va/derive_filtered_dataset.py create mode 100644 scripts/preprocess/minimax_h3_native_t2va/derive_filtered_dataset_1tray.sbatch create mode 100644 scripts/preprocess/minimax_h3_native_t2va/encode_native_t2va_1tray.sbatch create mode 100644 scripts/preprocess/minimax_h3_native_t2va/encode_worker.py create mode 100644 scripts/preprocess/minimax_h3_native_t2va/finalize_dataset.py create mode 100644 scripts/preprocess/minimax_h3_native_t2va/finalize_extension_1tray.sbatch create mode 100644 scripts/preprocess/minimax_h3_native_t2va/freeze_sources.py create mode 100644 scripts/preprocess/minimax_h3_native_t2va/prepare_extension_1tray.sbatch create mode 100644 scripts/preprocess/minimax_h3_native_t2va/probe_native_t2va.sbatch create mode 100644 scripts/preprocess/minimax_h3_native_t2va/schema_dry_run.py create mode 100644 scripts/preprocess/minimax_h3_native_t2va/seed_encoded_chunks.py create mode 100644 scripts/preprocess/minimax_h3_native_t2va/v10_sources.json create mode 100644 scripts/preprocess/minimax_h3_native_t2va/v10_sources_v2.json create mode 100644 scripts/preprocess/minimax_h3_native_t2va/validate_heldout_media.py create mode 100755 scripts/run_release20b_dmd2_v12_16gpu.sh create mode 100644 scripts/submit_release20b_dmd2_v12_corrected_16gpu.sbatch create mode 100644 scripts/submit_release20b_dmd2_v12_corrected_32gpu.sbatch create mode 100644 scripts/train/mfu_calc_minimax_h3.py create mode 100644 tests/local_tests/preprocess/test_minimax_h3_native_t2va.py create mode 100644 tests/local_tests/preprocess/test_minimax_h3_native_t2va_filtered.py diff --git a/docs/quantization/h3_int8_affine.md b/docs/quantization/h3_int8_affine.md new file mode 100644 index 0000000000..788ae89acc --- /dev/null +++ b/docs/quantization/h3_int8_affine.md @@ -0,0 +1,290 @@ +# Affine group-64 INT8 for MiniMax-H3 (CUDA) + +Weight-only affine INT8 quantization for H3 DiT inference on CUDA, with the +group-64 / 8-bit affine math the Apple Silicon (MLX) deployment lane already +validates. Implemented in +`fastvideo/layers/quantization/int8_affine_config.py`. + +This page is written for someone who has none of the session context in which +the lane was built. It covers what the scheme is, how to turn it on, exactly +which H3 layers are and are not quantized, and what has (and has not) been +verified. + +--- + +## 1. What this is + +Affine quantization stores each weight as + +``` +w ≈ code * scale + bias +``` + +where `code` is an unsigned integer and `scale`/`bias` are shared by a **group +of 64 weights taken along the input (contraction) dimension**. For a linear +weight of shape `[out, in]` that means `in / 64` groups per output row. + +The quantizer is a per-group min/max affine quantizer, not a symmetric one: +it anchors at whichever endpoint of the group's range has the larger +magnitude, so the extreme weight in each group round-trips exactly. That +detail is not incidental — it is the behaviour of MLX's `mx.quantize(..., +mode="affine")`, and it is what the project's QAT pipeline was tuned against. + +### Weight-only, deliberately + +Activations are **not** quantized. `apply()` dequantizes the stored codes back +to the activation dtype and runs a normal bf16/fp32 GEMM. Consequences worth +knowing: + +- Accuracy is limited by the weight error alone (measured below), not by an + activation-error term nobody has characterised for H3. +- It needs no INT8 tensor cores and no custom kernel — it runs anywhere bf16 + does, including CPU. +- The memory/bandwidth win is real only if you also drop the bf16 weight (see + `retain_original_weight`). Compute is unaffected until a fused INT8 GEMM + lands; that is a follow-up, not a prerequisite. + +## 2. This is NOT the MLX INT8 QAT callback + +Two different things in this repo are "affine group-64 INT8". Do not conflate +them: + +| | MLX lane (`mlx_affine_qat.py`) | This lane (`int8_affine_config.py`) | +|---|---|---| +| Purpose | **Training-time** fake-quant callback | **Load-time** inference quantization | +| When it runs | Every forward, during QAT finetuning | Once, when the checkpoint is loaded | +| What it produces | A straight-through-estimate `w` for the optimizer to train against | Stored int8 codes + per-group scales/biases | +| Target runtime | Apple Silicon / MLX | CUDA (PyTorch) | +| Weight after | Full-precision master, unchanged | Quantized in place (bf16 copy optionally retained) | +| Gradients | Pass through to the master weight | None — inference only | + +They share the *quantizer math* (this module's `int8_affine_quantize` / +`int8_affine_dequantize` are bit-identical to `mlx_affine_quantize_reference` / +`mlx_affine_dequantize_reference` — verified on fp32, fp16 and bf16 inputs) +and nothing else. Nothing here writes a QAT checkpoint, and nothing here can +be used to *run* QAT: `apply()` deliberately falls back to the dense bf16 +weight whenever `torch.is_grad_enabled()`, so a training step never sees a +frozen dequantized copy. + +## 3. Turning it on for H3 + +The config is registered under the name `INT8Affine` (see +`fastvideo/layers/quantization/__init__.py`). + +**Option A — the verified H3 profile (recommended).** Build the config +explicitly, because the registry resolves a bare name through the no-argument +constructor and therefore cannot carry the H3-specific profile: + +```python +from fastvideo.layers.quantization.int8_affine_config import INT8AffineConfig + +fastvideo_args.transformer_quant = INT8AffineConfig.for_minimax_h3() +``` + +`transformer_quant` accepts a pre-built instance and pins it onto +`pipeline_config.dit_config.quant_config` in `FastVideoArgs._apply_transformer_quant`, +before the DiT is constructed — which is required, because linears attach their +`quant_method` during `__init__`. Setting +`pipeline_config.dit_config.quant_config` directly works too. + +**Option B — by name.** `transformer_quant: "INT8Affine"` (YAML) or +`--transformer-quant INT8Affine` builds `INT8AffineConfig()` with the generic +defaults. That default is safe on H3 (it selects the same attention/FFN GEMMs, +minus `adaln_proj.linear`, and the same exclusions apply), but it is the +model-agnostic profile rather than the H3-verified one. + +The load-time conversion is triggered from `_maybe_quantize_model` in +`fastvideo/models/loader/fsdp_load.py`. That function dispatches on an explicit +`isinstance` chain, so it needs a branch for `INT8AffineQuantizeMethod` calling +`convert_model_to_int8_affine(model)`. **If that branch is missing, inference +is still correct** — `apply()` converts lazily on first forward and logs a +warning naming this exact cause — but you will see the warning, and the +conversion happens inside the first forward instead of at load. + +Any BF16 checkpoint works unchanged; there are no pre-quantized weights to +produce and none are written back (the codes/scales are non-persistent +buffers, so they are re-derived on every load). + +## 4. Which H3 layers are quantized + +H3's `MiniMaxH3TransformerBlock` builds +`{prefix}.transformer_blocks.{i}.{attn, ff, adaln_proj}`, and +`MiniMaxH3TokenRefiner` builds +`{prefix}.token_refiner.refiner_blocks.{i}.{attn, ff}` (no `adaln_proj`). +With the defaults from `MiniMaxH3ArchConfig` (`prefix="minimax_h3"`, +`num_layers=50`, `num_refiner_layers=2`): + +**Quantized** (362 linears under `for_minimax_h3()`): + +| Prefix | Count | +|---|---| +| `minimax_h3.transformer_blocks.{0..49}.attn.to_q` / `to_k` / `to_v` / `to_out` | 200 | +| `minimax_h3.transformer_blocks.{0..49}.ff.fc_in` / `fc_out` | 100 | +| `minimax_h3.transformer_blocks.{0..49}.adaln_proj.linear` | 50 | +| `minimax_h3.token_refiner.refiner_blocks.{0..1}.attn.to_q` / `to_k` / `to_v` / `to_out` | 8 | +| `minimax_h3.token_refiner.refiner_blocks.{0..1}.ff.fc_in` / `fc_out` | 4 | + +**Excluded, and why:** + +| Prefix | Why | +|---|---| +| `...attn.to_gate_compress` | **Never quantize this.** See §5. | +| `minimax_h3.adaln_basis` | Global timestep-basis projector feeding every block's modulation. | +| `minimax_h3.proj_in`, `audio_proj_in`, `proj_out`, `audio_proj_out` | H3 pins these to fp32 (`_keep_in_fp32_modules`) to preserve input/output precision. | +| `minimax_h3.time_embedder.fc_in` / `fc_out` | Same fp32 set. | +| `minimax_h3.context_embedder` | Judgment call — see below. | +| norms, `rope`, `scale_shift_table`-style params | Not `LinearBase`; never candidates. | + +`context_embedder` is excluded **by default** although H3 does not pin it to +fp32. It is the text input projection, structurally the same kind of module as +`proj_in` / `audio_proj_in` (which H3 *does* keep in fp32), and quantizing the +text conditioning stream while leaving the video/audio input streams in fp32 is +an asymmetry nobody has validated. Pass `include_context_embedder=True` to opt +in once there is evidence either way. + +Note also that `adaln_proj.linear` **is** included (it is a real per-block +GEMM) while `adaln_basis` is not. They are different modules with similar +names; the exclusion list keys on `adaln_basis`, which is not a substring of +`adaln_proj.linear`. + +## 5. `attn.to_gate_compress` must never be quantized + +`MiniMaxH3Attention.to_gate_compress` is the VSA sparse-attention gate. Two +reasons it is special: + +1. Its output decides **sparse routing** — which tiles the sparse attention + attends to. That is a discrete decision. Quantization error there does not + perturb an activation by a fraction of a percent; it can flip a routing + decision outright, and the resulting error is not bounded by the weight + quantization error. +2. H3's own deploy path explicitly ignores it, and the released checkpoint + zero-initializes the gate, so the branch is exactly disabled until it is + finetuned. There is nothing to gain by quantizing it. + +The name matches none of the usual exclusion heuristics (it contains no +`norm`, no `scale_shift_table`, no `proj_*`), so it would be swept into any +broad suffix rule. It is therefore excluded by an **explicit, fail-closed deny +list**: + +- `INT8AffineConfig.exclude_substrings` is always the union of the hard-coded + never-quantize names and any caller-supplied list. The constructor can only + *add* exclusions. +- The deny check short-circuits before both the `layer_suffixes` and + `target_layers` paths, so even explicitly listing + `minimax_h3.transformer_blocks.7.attn.to_gate_compress` in `target_layers` + does not select it. + +`test_to_gate_compress_is_excluded_even_by_a_broad_allowlist` in +`fastvideo/tests/ops/quantization/test_int8_affine_config.py` is the +regression guard for this; it exercises both escape routes. + +## 6. Requirements + +- **Dependencies:** none beyond PyTorch. No `flashinfer`, no CUDA kernels, no + INT8 tensor cores. The module imports on a CPU-only host. +- **Hardware:** `get_min_capability()` returns 75 (Turing), matching + `AbsMaxFP8Config`. The compute path is an ordinary bf16/fp32 GEMM, so this + is a conservative floor rather than a real requirement. +- **Shape constraint:** `group_size` (default 64) must divide the weight's + input dimension. H3 satisfies this everywhere (`hidden_size=5376`, + `ffn_dim=14336`, and `2 * ffn_dim`, all divisible by 64, including under + tensor parallelism at the sizes used). A violation raises `ValueError` at + conversion rather than silently mis-grouping. +- **State dicts:** the quantized codes/scales/biases are non-persistent + buffers. Nothing about this config changes checkpoint contents; you cannot + save a "quantized H3" and must re-derive at load. +- **Training:** not supported by this config (see §2). Use a separate + `*_qat_train`-style method if a recovery run is needed. + +## 7. What is tested, and what is not + +**Tested** (`pytest fastvideo/tests/ops/quantization/test_int8_affine_config.py`, +16 tests, CPU-only, no GPU needed): + +- The module imports without CUDA or flashinfer. +- Quantizer parity: codes/scales/biases match an independently written + restatement of the MLX algorithm, and — separately verified outside the + suite — are **bit-identical** to `mlx_affine_quantize_reference` / + `mlx_affine_dequantize_reference` on fp32, fp16 and bf16 inputs. The + in-tree parity test skips (with a reason) while `mlx_affine_qat.py` is absent + from this worktree, and activates automatically once it lands. +- Round-trip error on `N(0,1)` weights, measured over seeds 0-5 with ~2x + headroom: `max|Δw| / max|w| ≤ 0.0055` and `rms|Δw| / max|w| ≤ 0.0013` for an + fp32 source; `≤ 0.0093` and `≤ 0.0021` for a bf16 source. +- Layer selection for real H3 names: the include set above is selected, and + `to_gate_compress`, `adaln_basis`, the fp32-pinned modules, + `context_embedder`, norms and `rope` are not — including under a hostile + broad suffix rule and via an explicit `target_layers` set. +- The enumerated H3 prefix set and the runtime suffix rule agree exactly. +- A real `ReplicatedLinear` mounts `INT8AffineQuantizeMethod` where expected + and `UnquantizedLinearMethod` on the gate; conversion + `apply()` runs on + CPU, and the purge path (`retain_original_weight=False`) works. + +**NOT tested — do not assume any of this works:** + +- **No model-level run.** Nothing here has been executed against a real H3 + checkpoint, on GPU or otherwise. The DiT has not been instantiated with this + config, and no forward has produced a video. +- **No quality evidence.** No SSIM, no VLM adherence check, no comparison to + the bf16 baseline. The measured error is *weight-level* round-trip error; + how it propagates through 50 blocks of a diffusion DiT is unmeasured. +- **No performance numbers.** Memory, throughput and load-time cost are all + unmeasured. In particular the dequantize-per-forward path in `apply()` has + never been timed; it may be a significant overhead at inference. +- **Load-time peak memory is unmeasured.** Conversion makes an fp32 copy of + each weight (`.detach().float().nan_to_num()`) one layer at a time. For + H3's largest targeted weight that is a transient ~0.6 GB on top of the + loaded model. It should be freed per layer, but this has not been profiled + on a real load. +- **No multi-GPU / FSDP / TP validation.** The conversion walks + `DTensor`-wrapped weights via `to_local()`, following `convert_model_to_nvfp4`, + but has only been exercised on single-process CPU tensors. Whether quantizing + a *shard* independently reproduces quantizing the whole weight depends on + which dimension the shard is taken along and on the tensor-parallel degree — + that has not been checked. If you enable this under FSDP or TP, verify that + the shard's input dimension is still divisible by `group_size` (the + conversion raises if it is not) and that scales agree with a single-process + conversion. +- **The loader-hook dispatch is not wired by this change.** See §3; the + sibling change to `_maybe_quantize_model` is required for the load-time + (rather than lazy) conversion path. +- **`torch.compile` interaction is untested.** The lazy fallback in `apply()` + mutates the layer on first call, which is not compile-friendly; the load-time + conversion path is the one to use under compile. + +## 8. Do you need a QAD / QAT recovery run for 4090 deployment? + +Stated plainly, because this is a stated future use: + +- **To run at all on a 4090: no.** The 4090 is sm89 with bf16 support and this + config's compute path is a plain bf16 GEMM, so it runs as-is. No recovery + run, no kernel build, no calibration data. +- **To preserve quality on a 4090: unknown, and this is the honest answer.** + This is post-training quantization with no recovery step. Weight-only PTQ at + 8 bits with group-64 typically holds up better than activation-quantized + schemes, but "typically" is not evidence, and nothing in this lane has been + evaluated end-to-end. Treat a recovery run as *contingent on an eval gate*, + not as a known requirement: run the existing quality gates against the bf16 + baseline first, and only if they regress does QAT/QAD become the next step. +- **If it does regress, a recovery run needs new code.** This config cannot be + used for training (§2), and no INT8 analogue of `nvfp4_qat_train_config.py` / + `fp8_qat_train_config.py` exists. The template is one of those files: a + `QuantizeMethodBase` that keeps a trainable master weight and + fake-quantizes with a straight-through estimator each forward. The + fake-quantization itself already exists here — `int8_affine_quantize` + + `int8_affine_dequantize` compose into exactly the STE the MLX lane's + `fake_quantize_mlx_affine` performs, minus the `simulate_dtype` cast (which + exists only because MLX loads checkpoints as fp16). +- **Also note:** the 4090 is a different target from the MLX lane, so a + recovery run would not be transferable work-for-work with the Apple Silicon + QAT effort — but they share the quantizer, so a checkpoint trained with the + MLX QAT callback is quantized by the *same* decisions this config makes. + +## 9. Files + +| Path | What | +|---|---| +| `fastvideo/layers/quantization/int8_affine_config.py` | `INT8AffineConfig`, `INT8AffineQuantizeMethod`, `convert_model_to_int8_affine`, `int8_affine_quantize` / `int8_affine_dequantize`, `minimax_h3_int8_affine_prefixes` | +| `fastvideo/tests/ops/quantization/test_int8_affine_config.py` | CPU-only unit tests | +| `fastvideo/layers/quantization/__init__.py` | Registry entry for the name `INT8Affine` (owned elsewhere) | +| `fastvideo/models/loader/fsdp_load.py` | `_maybe_quantize_model` dispatch (owned elsewhere) | diff --git a/docs/quantization/h3_nvfp4.md b/docs/quantization/h3_nvfp4.md new file mode 100644 index 0000000000..568ad236ca --- /dev/null +++ b/docs/quantization/h3_nvfp4.md @@ -0,0 +1,443 @@ +# NVFP4 quantization for MiniMax-H3 (Blackwell / sm100) + +Load-time NVFP4 weight quantization for the MiniMax-H3 joint audio-video DiT, +implemented in `fastvideo/layers/quantization/nvfp4_config.py`. + +This page is written for someone with none of the session context in which the +lane was built. It covers what NVFP4 is here, why the config used to be a +**silent no-op on H3**, how to turn it on, exactly which layers are quantized, +what is required to run it, and how to write/read a compact quantized +checkpoint. + +--- + +## 1. What NVFP4 is here + +NVFP4 is NVIDIA's block-scaled FP4 format: an **e2m1** 4-bit weight code, an +**e4m3** scale shared by every 16 weights along the input dimension, and a +single **fp32/bf16 global scale** per tensor. In this repo it is backed by +[`flashinfer`](https://github.com/flashinfer-ai/flashinfer) — +`nvfp4_quantize` for the conversion and `mm_fp4` for the GEMM. + +It is *not* generic FP4 / OCP-FP4 / MX-FP4 — hence the explicit name. The +scale layout used everywhere in this module is `SfLayout.layout_128x4` with +`do_shuffle=False`; rows are padded to a 128-row tile for the kernel and +narrowed back (`_nvfp4_quantize`). + +For a linear weight of shape `[out, in]`: + +| Tensor | dtype | shape | contents | +|---|---|---|---| +| `_nvfp4_weight` | `uint8` | `[out, ceil(in/2)]` | two e2m1 codes per byte, packed along K | +| `_nvfp4_weight_scale` | `uint8` | `[out, ceil(in/16)]` | e4m3 block-scale bit patterns (rows padded to a multiple of 128 when `out` is not; H3's dims are, so this does not apply here) | +| `_weight_global_sf` | `bfloat16` | scalar | `(448 * 6) / max\|W\|` | +| `_nvfp4_alpha` | `float32` | scalar | `1 / _weight_global_sf`, kept at fp32 precision | + +All four are registered as **`persistent=False`** buffers by +`convert_model_to_nvfp4`, so they never appear in a `state_dict`. That is +deliberate (they are derived state on the standard path) and is exactly why a +compact quantized checkpoint needs the sidecar described in §7. + +Weight storage drops from 16 bits/weight to **0.5625** (4 bits + 1 scale byte +per 16 weights) — about **3.6x** smaller than bf16. + +## 2. Why NVFP4 used to do nothing on H3 + +`NVFP4Config.get_quant_method` decides what to quantize by looking the layer's +module path up in a set. That set used to be a hardcoded module-level constant +of ~577 literal `ltx2.blocks.*` strings: + +```python +_LTX2_NVFP4_LINEAR_PREFIXES = frozenset(...) # 48 blocks x 12 suffixes + adaln +``` + +H3's linears are named `minimax_h3.transformer_blocks.*`, so **not one path +matched**, and: + +- no `NVFP4QuantizeMethod` was attached to any layer, +- `_maybe_quantize_model` in `fsdp_load.py` therefore found no NVFP4 layer and + skipped the conversion entirely, +- the model ran dense bf16 — with **no error, no warning, and no log line**. + +That is the worst possible failure shape: `transformer_quant="NVFP4"` looked +configured, and the only way to notice was to measure memory or read the code. + +The fix lifts the layer list into the config (the class docstring had been +asking for this): + +```python +NVFP4Config(layer_profile="refine", + retain_original_weights=None, + layer_prefixes=None, # None -> the LTX-2 set (unchanged default) + exclude_prefixes=None) # extra never-quantize patterns +``` + +`layer_prefixes=None` keeps the historical LTX-2 behaviour bit-for-bit, so +existing LTX-2 deployments are unaffected. Other models opt in explicitly. + +## 3. Turning it on for H3 + +**Option A — the H3 profile (use this one).** Build the config explicitly and +hand it to `FastVideoArgs`: + +```python +from fastvideo.layers.quantization.nvfp4_config import NVFP4Config + +fastvideo_args.transformer_quant = NVFP4Config.for_minimax_h3() +``` + +`FastVideoArgs._apply_transformer_quant` pins a pre-built instance onto +`pipeline_config.dit_config.quant_config` before the DiT is constructed, which +is required — linears attach their `quant_method` during `__init__`. Setting +`pipeline_config.dit_config.quant_config` directly works too and takes +precedence. + +Equivalent, if you prefer the raw constants: + +```python +from fastvideo.layers.quantization.nvfp4_config import ( + MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES, + MINIMAX_H3_NVFP4_LINEAR_PREFIXES, + NVFP4Config, +) + +config = NVFP4Config( + layer_prefixes=MINIMAX_H3_NVFP4_LINEAR_PREFIXES, + exclude_prefixes=MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES, +) +``` + +`for_minimax_h3(...)` also forwards keyword arguments, e.g. +`for_minimax_h3(retain_original_weights=True)`. + +**Option B — by name. This does not work for H3.** `transformer_quant: +"NVFP4"` (YAML) or `--transformer-quant NVFP4` resolves through +`get_quantization_config("NVFP4")()` — a no-argument constructor, so it gets +`layer_prefixes=None`, i.e. the LTX-2 set. On H3 that is the silent no-op from +§2. The registry has no way to carry a model-specific layer list; a +pre-built instance is the only way to express one. + +**How to confirm it actually engaged.** Three things must be true after the +model is loaded: + +1. The loader logs `Converting loaded model weights for NVFP4 linear layers`. +2. `convert_model_to_nvfp4` logs its retention receipt, e.g. + `NVFP4 weight purge receipt: purged 0 original bf16 weight tensors ...; retained 300`. + Expected on the standard multi-GPU path — see §6 on the FSDP retention. +3. Peak memory moves: +~7 GiB for the packed buffers, and −~25 GiB only if the + receipt reports purges (single-process loads). + +If none of those appear, the config did not cover the model's layer paths — +check `layer_prefixes` against `[name for name, _ in model.named_modules()]`. + +## 4. Which H3 layers are quantized + +With `MiniMaxH3ArchConfig` defaults (`prefix="minimax_h3"`, `num_layers=50`, +`hidden_size=5376`, `ffn_dim=14336`): + +**Quantized — 300 linears** (`MINIMAX_H3_NVFP4_LINEAR_PREFIXES`): + +| Prefix | Count | +|---|---| +| `minimax_h3.transformer_blocks.{0..49}.attn.to_q` / `to_k` / `to_v` / `to_out` | 200 | +| `minimax_h3.transformer_blocks.{0..49}.ff.fc_in` / `fc_out` | 100 | + +Those 300 linears hold ≈13.5B parameters: ≈25.1 GiB as bf16, ≈7.1 GiB as +packed NVFP4 (arithmetic from the arch constants, not a measurement). + +**Deliberately excluded:** + +| Prefix | Why | +|---|---| +| `...attn.to_gate_compress` | **Never quantize.** See §5. | +| `minimax_h3.token_refiner.refiner_blocks.*` | Text-stream refiner; small, and its attention sits outside the packed video/audio sequence the FP4 kernels are tuned for. Quantizing it is a follow-up, not a default. | +| `minimax_h3.transformer_blocks.*.adaln_proj.linear`, `norm_out.linear` | Per-block/per-forward modulation projections; small, and they feed the modulation path whose precision the distilled deploy is sensitive to. | +| `proj_in`, `audio_proj_in`, `proj_out`, `audio_proj_out`, `time_embedder`, `rope` | H3 pins these to fp32 via `MiniMaxH3Transformer3DModel._keep_in_fp32_modules` — quantizing them would fight the model's own dtype policy. | +| `context_embedder` | Judgment call: the text conditioning stream, structurally like the fp32-pinned input projections. | +| norms, `rope` tables, `scale_shift_table`-style params | Not `LinearBase`; never candidates. | + +The set is a **positive allowlist**, so extending it (say, adding the token +refiner) is a one-line change to `layer_prefixes` — nothing else in the module +needs to know. + +Note that `adaln_basis` (global timestep-basis projector) and +`adaln_proj.linear` (per-block modulation) are different modules with similar +names; only the latter is even a candidate, and it is excluded. + +## 5. `attn.to_gate_compress` must never be quantized + +`MiniMaxH3Attention.to_gate_compress` is the VSA (Video Sparse Attention) +compression gate. Two independent reasons: + +1. **Its output decides sparse routing** — which tiles the sparse attention + attends to. That is a discrete decision, so quantization error there is not + bounded by the weight-quantization error; it can flip a routing decision. +2. **H3's own deployment path ignores it**, and the released checkpoint + zero-initializes the gate, so the branch is exactly disabled until it is + finetuned. `MiniMaxH3Attention._gate_active()` tests the loaded weight once + and skips the branch entirely while it is structurally zero. There is + nothing to gain by quantizing it, and a nonzero quantized gate would defeat + the skip. + +The name matches no generic exclusion heuristic (no `norm`, no `scale_shift_table`, +no `proj_*`), so it would be swept up by any broad suffix rule. It is excluded +by **three independent mechanisms**, so no single mistake can quantize it: + +1. **It is absent from the allowlist.** `MINIMAX_H3_NVFP4_LINEAR_PREFIXES` + contains exactly the 300 paths of §4; the gate is not among them. +2. **`for_minimax_h3()` sets `exclude_prefixes`** to + `MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES` = `("attn.to_gate_compress",)`. +3. **`_ALWAYS_EXCLUDED_LINEAR_SUFFIXES` is unconditional.** The deny check runs + *first* in `NVFP4Config.is_nvfp4_linear_prefix`, before the allowlist, and no + caller can override it. Passing an allowlist that explicitly contains + `minimax_h3.transformer_blocks.0.attn.to_gate_compress` still returns + `False`. + +Exclusion entries match either a full module path or a trailing suffix, and +only at a dot boundary (`"ff.fc_in"` does not match `"cross_ff.fc_in"`). + +Regression guards: +`test_gate_is_excluded_even_if_a_caller_allowlists_it` and +`test_get_quant_method_attaches_for_h3_and_skips_the_gate` in +`fastvideo/tests/ops/quantization/test_nvfp4_h3_prefixes.py`. + +## 6. Requirements and caveats + +**Requirements** + +- `flashinfer-python` (validated set in + [Optimizations](../inference/optimizations.md) §flashinfer). The module + itself imports fine without it; only the quantize/GEMM ops raise, with + `Install with 'pip install flashinfer-python'`. Exception: restoring a + sidecar (§7) needs no flashinfer at all. +- **Blackwell / sm100.** `NVFP4Config.get_min_capability()` returns `100`. The + kernels need FP4 tensor cores. +- A bf16 checkpoint. No pre-quantized weights are required: conversion happens + at load, from the dense weights, every time. + +**Caveat: `x_global_sf` is hardcoded `1.0`** + +`NVFP4QuantizeMethod.x_global_sf` is `torch.tensor(1.0)` — a constant, never +derived from the data. It is used twice: + +- as the global scale when quantizing *activations* + (`quantize_input` / `apply`), +- as the denominator of the GEMM alpha: `alpha = layer._nvfp4_alpha / x_global_sf` + (or `1 / (x_global_sf * weight_global_sf)`). + +The weight side of the same math *is* data-derived +(`convert_model_to_nvfp4` computes `(448 * 6) / max|W|`), and the standalone +`fastvideo/layers/fp4linear.py` derives the activation side too +(`_global_sf` / `448.0 * 6.0 / maxabs`). So H3 (and LTX-2) currently quantize +activations against a fixed unit global scale instead of adapting to each +activation's dynamic range. This is the **first suspect** for any quality +regression from NVFP4, and it is a shared pre-existing property of this config, +not something the H3 enablement introduced. + +It has not been measured on H3. The experiment to run: compare `x_global_sf = +1.0` against a per-tensor `448*6/max|x|` on a fixed prompt/seed and score both +with the project's quality gates. Because `x_global_sf` is *not* part of a +saved checkpoint (§7) it must be identical on both sides of a save/load — if it +ever becomes data-derived, it has to be persisted alongside the weights. + +**Other caveats** + +- `layer_profile` (`"base"` / `"refine"`) is an LTX-2 streaming concept. H3 + layers are never "refine-only", so every H3 quantized layer is FP4 on every + step. +- No H3 quality evidence exists yet (no SSIM, no VLM adherence gate against the + bf16 baseline), and no performance numbers. See §9. +- **FSDP: the dense bf16 weights are retained, not purged.** `shard_model` runs + before `_maybe_quantize_model` in `maybe_load_fsdp_model`, so by conversion + time every `weight` is a `DTensor`; `convert_model_to_nvfp4` quantizes the + local shard via `to_local()` but skips the purge for DTensor weights + (per-shard resharding bookkeeping was never implemented). On the standard + multi-GPU path the receipt therefore reads `retained 300`, and enabling NVFP4 + *adds* the ~7 GiB of packed buffers rather than replacing the ~25 GiB of bf16 + weights. The bandwidth win at inference is real (the FP4 GEMM reads the + packed weight); the memory win is not, until that purge lands. This is a + pre-existing property of the purge policy, not of the H3 enablement. +- **FSDP/TP and sidecars:** conversion quantizes each rank's *local shard*, so + the global scale is per-shard. A sidecar saved from a sharded model stores + local shards and is only reloadable into an identically sharded model. + +## 7. Compact NVFP4 checkpoints + +**The problem.** The FP4 tensors are `persistent=False`, so they are not in a +`state_dict`: saving an H3 model writes ~25 GiB of dense bf16 weights for these +300 linears, and every load re-quantizes them from scratch (requiring +flashinfer and the time to run 300 quantizations). + +**The format.** A *sidecar* safetensors file holding the quantized tensors, +keyed by module path: + +``` +::_nvfp4_weight +::_nvfp4_weight_scale +::_weight_global_sf +::_nvfp4_alpha +``` + +Its safetensors `metadata` carries a JSON manifest under the key +`fastvideo_nvfp4`: + +```json +{ + "format": "fastvideo.nvfp4", + "version": 1, + "sf_layout": "layout_128x4", + "do_shuffle": false, + "block_size": 16, + "num_layers": 300, + "layers": {"minimax_h3.transformer_blocks.0.attn.to_q": [5376, 5376], "...": "..."}, + "quant_prefixes": {"minimax_h3.transformer_blocks.0.attn.to_q": "minimax_h3.transformer_blocks.0.attn.to_q"}, + "model_class": "MiniMaxH3Transformer3DModel" +} +``` + +`layers` records each linear's `[out, in]`, so a load can validate tensor +shapes (and report coverage) even when the bf16 weights are not present at all. +`quant_prefixes` records the prefix the layer was tagged with, which is how you +tell an H3-built sidecar from an LTX-2-built one. + +**Writing one:** + +```python +from fastvideo.layers.quantization.nvfp4_config import save_nvfp4_checkpoint + +receipt = save_nvfp4_checkpoint(model, "transformer.nvfp4.safetensors") +# {'num_layers': 300, 'num_tensors': 1200, 'quantized_bytes': ..., +# 'dense_bfloat16_bytes': ..., 'compression_ratio': 3.55...} +``` + +The model must already be converted (`convert_model_to_nvfp4` has run, which is +what the loader does at load time). The receipt, which is also logged, reports +both sizes so the win is visible without `ls -l`. + +**Reading one:** + +```python +from fastvideo.layers.quantization.nvfp4_config import load_nvfp4_checkpoint + +restored = load_nvfp4_checkpoint(model, "transformer.nvfp4.safetensors") +``` + +`load_nvfp4_checkpoint`: + +- registers the four buffers on every NVFP4-tagged linear, byte-for-byte as a + fresh conversion would (`test_load_restores_buffers_without_reconverting`), +- **never calls flashinfer** — a host that only serves a pre-quantized + checkpoint needs neither the kernels nor a GPU for this step, +- works when the dense `weight` is absent entirely (the compact case), +- validates the manifest (`format`, `version`, `sf_layout`, `do_shuffle`, + `block_size`) and every tensor shape, raising `ValueError` — a layout + mismatch is never downgraded, because mis-read nibbles are silent corruption, +- reports layer-set mismatches: `strict=True` (default) raises, `strict=False` + logs and restores the intersection, +- applies the same bf16-weight retention policy as the conversion + (`purge_dense_weights=True` by default). + +Conventional path for a checkpoint file or directory: +`nvfp4_sidecar_path_for(".../transformer.safetensors")` → +`.../transformer.nvfp4.safetensors`; for a directory → `/nvfp4.safetensors`. + +### What still needs a loader-side change (not done here) + +`load_nvfp4_checkpoint` is complete, tested, and callable today, but the loader +does not yet *dispatch* it — `fastvideo/models/loader/fsdp_load.py` is owned by +another change, so this work did not touch it. Two separate things are needed +for a compact release, and they are independent: + +1. **Skip the re-quantization.** In `_maybe_quantize_model`, where + `convert_model_to_nvfp4(model)` is called, branch on the sidecar's presence: + + ```python + from fastvideo.layers.quantization.nvfp4_config import ( + load_nvfp4_checkpoint, nvfp4_sidecar_path_for, + ) + + sidecar = nvfp4_sidecar_path_for(transformer_checkpoint_path) + if os.path.exists(sidecar): + load_nvfp4_checkpoint(model, sidecar) + else: + convert_model_to_nvfp4(model) + ``` + + `_maybe_quantize_model(model)` does not currently receive a path, so its + signature (or its call site) has to grow one. Without this, the sidecar + saves load *time* only if the caller invokes `load_nvfp4_checkpoint` + manually after the model is built — the automatic path still re-quantizes. + +2. **Drop the bf16 weights from the released file.** This is the half that + actually shrinks the *release*. The sidecar already omits them, but the + main checkpoint cannot: at load time the DiT is constructed with a real + `weight` `Parameter` (`NVFP4QuantizeMethod.create_weights` allocates it), so + `weight` is in `model.state_dict()`, and + `load_model_from_full_model_state_dict` treats any model key the checkpoint + does not provide as a new/unmapped parameter — it zero-initializes it and + raises `ValueError` unless the name is on the narrow allowlist described in + [Quantized Checkpoint Loading](loader_quant_params.md). Shipping a + checkpoint without the 300 NVFP4 `weight` tensors therefore requires that + check to accept `weight` on NVFP4-tagged linears *whose sidecar covers + them* — and requires `create_weights` (or the load path) to tolerate a + `weight` that is never filled. + + With both halves in place the transformer file for these 300 linears goes + from ≈25 GiB to ≈7 GiB. + +Until then, treat the sidecar as: a working format + restore path, a way to +serve pre-quantized weights without flashinfer, and a de-risked piece of the +compact-release work. + +## 8. Tests + +CPU-only, no flashinfer, no CUDA (the FP4 ops and +`NVFP4QuantizeMethod.__init__`'s cuda allocation are stubbed): + +```bash +pytest fastvideo/tests/ops/quantization/test_nvfp4_h3_prefixes.py -v +pytest fastvideo/tests/ops/quantization/test_nvfp4_sidecar.py -v +``` + +`test_nvfp4_h3_prefixes.py` (9 tests) covers the LTX-2 default (no +regression, 577 prefixes), the 300-linear H3 set, the gate exclusion including +the hostile-allowlist case, `from_config` round-tripping, the +import-without-flashinfer contract, and a real `ReplicatedLinear` end-to-end +attachment check. + +`test_nvfp4_sidecar.py` (14 tests) covers the save receipt and size win, +byte-identical restore, restore without flashinfer, restore into a model with +no dense weights, the retention policy, layer-set/layout/version/shape +mismatch handling, and the "no NVFP4 layers attached" error that names the +prefix-set cause. + +## 9. What is NOT verified + +- **No GPU run.** Nothing here has been executed on Blackwell, and the FP4 + kernels have not been exercised by this work. All tests stub the quantizer, + so they verify shapes, dtypes, layout bookkeeping, keying and policy — not + numerics. +- **No quality evidence.** No SSIM, no VLM adherence check, no comparison to + the bf16 baseline for H3. Combined with the `x_global_sf = 1.0` caveat (§6), + a quality regression is plausible and unmeasured. +- **No performance numbers.** Memory and throughput impact unmeasured; the + activation-quantize cost per call is unmeasured. +- **No FSDP/TP validation** of save/load. A sidecar holds local shards + (§6); round-tripping under `fully_shard` has not been tested. +- **No end-to-end loader dispatch** — see §7. +- **Checkpoint-wrapper caveat.** Sidecar keys are `named_modules()` FQNs. If a + model is saved with `checkpoint_wrapper`-style prefixes and loaded without + them (or vice versa), the keys will not match; `strict=True` will say so + rather than restoring partially. + +## 10. Files + +| Path | What | +|---|---| +| `fastvideo/layers/quantization/nvfp4_config.py` | `NVFP4Config` (+ `layer_prefixes` / `exclude_prefixes` / `for_minimax_h3`), `NVFP4QuantizeMethod`, `convert_model_to_nvfp4`, `save_nvfp4_checkpoint`, `load_nvfp4_checkpoint`, `read_nvfp4_sidecar_metadata`, `nvfp4_sidecar_path_for`, `MINIMAX_H3_NVFP4_LINEAR_PREFIXES` | +| `fastvideo/tests/ops/quantization/test_nvfp4_h3_prefixes.py` | CPU-only prefix-selection tests | +| `fastvideo/tests/ops/quantization/test_nvfp4_sidecar.py` | CPU-only serialization tests | +| `fastvideo/models/loader/fsdp_load.py` | `_maybe_quantize_model` dispatch (owned elsewhere; sidecar branch not yet wired — §7) | +| `fastvideo/layers/linear.py` | `ReplicatedLinear.__init__` is where `get_quant_method` is called | +| `docs/quantization/loader_quant_params.md` | Why unmapped parameters are a hard error at load | +| `docs/quantization/h3_int8_affine.md` | The sibling INT8 lane for the same model | diff --git a/docs/quantization/h3_w4a16.md b/docs/quantization/h3_w4a16.md new file mode 100644 index 0000000000..b01f3fb238 --- /dev/null +++ b/docs/quantization/h3_w4a16.md @@ -0,0 +1,366 @@ +# W4A16 (4-bit weight, 16-bit activation) for MiniMax-H3 + +Load-time, weight-only 4-bit quantization for the H3 joint audio-video DiT, +targeting Ada-generation consumer GPUs. Implemented in +[`fastvideo/layers/quantization/w4a16_config.py`](../../fastvideo/layers/quantization/w4a16_config.py). + +> **Read section 6 before quoting any number from this page.** There is no +> fused W4A16 GEMM in this repository or in its dependencies. What ships here +> is a *correctness reference*: real 4-bit weight storage, and a +> dequantize-then-dense-GEMM compute path. It is a memory lane, not a speed +> lane, until a kernel lands. + +--- + +## 1. What this is + +**W4A16** names exactly two choices: + +| | Precision | Where | +|---|---|---| +| Weights | 4-bit integer codes, per-group scale + zero point | stored on device, at rest | +| Activations | 16-bit float (bf16 / fp16) | untouched, every forward | + +The weights are quantized **group-wise**: consecutive runs of `group_size` +values along the contraction (input) axis share one `(scale, zero)` pair. The +fit is min/max affine — `scale = (max - min) / 15`, `zero = round(-min / scale)` +— so both endpoints of every group round-trip exactly, and the per-element +reconstruction error is bounded by half a quantizer step. + +Which linears get quantized is a **constructor field** (`target_layers`, an +explicit allowlist; or `layer_suffixes`, a generic suffix rule), never a +hardcoded model constant. That is deliberate: `NVFP4Config` hardcodes the LTX-2 +prefix list and its own docstring flags that as a wart. Do not repeat it here. + +### Storage math + +For a weight tensor with `K` input features per row: + +| Item | Bits per weight | +|---|---| +| 4-bit codes (two per byte, low nibble = lower K index) | 4 | +| per-group scale (fp32) | `32 / group_size` | +| per-group zero point (fp32) | `32 / group_size` | +| **Total at `group_size=64`** | **5.0** | +| **Total at `group_size=128`** | **4.5** | +| (BF16 baseline, for comparison) | 16 | + +So roughly **3.2x smaller than BF16** at `group_size=64`. For H3's 20B DiT that +is the difference between a checkpoint that fits a 24 GB card's weights and one +that does not. + +### Weight-only, deliberately + +There is no activation quantizer in this module and there should not be one +until a kernel exists. Quantizing activations without a fused GEMM means +quantize-dequantize on every forward — pure overhead for a numerical result +the dense GEMM would have produced anyway. Weight-only keeps the change to the +one thing that is unambiguously a win at rest: what is stored. + +--- + +## 2. Why this lane targets Ada, not Blackwell + +The relevant hardware split: + +| Generation | Example SKUs | Native FP4 tensor cores | Native FP8 | W4A16 | +|---|---|---|---|---| +| Ada (sm_89) | RTX 4090 24 GB, RTX 6000 Ada 48 GB | **No** | Yes | via INT4/INT8 tensor cores (needs a fused kernel) | +| Hopper (sm_90) | H100 | No | Yes | via INT4 tensor cores | +| Blackwell (sm_100/sm_120) | RTX 5090, RTX PRO 6000 | Yes (NVFP4) | Yes | possible, but NVFP4 is the better answer there | + +Ada has **no native NVFP4 path**. The `NVFP4` lane in this repository is +Blackwell: it is backed by FlashInfer's `mm_fp4`, it declares +`get_min_capability() == 100`, and it should not be presented as a 4090 speed +claim. On a 4090, the low-bit options that actually apply are INT8, FP8, and +W4A16 — and of those three, W4A16 is the one that shrinks storage the most. + +This matches the release plan's hardware matrix +(`H3_REALTIME_RELEASE_PLAN_2026-09-16.md`): *"RTX 4090 and RTX 6000 Ada do not +have native NVFP4 Tensor Core support. Their low-bit release work should +prioritize INT8, FP8, and W4A16 candidates that have real kernels on Ada."* + +The last clause is the operative one. This module supplies the W4A16 *schema +and wiring*; see section 6 for the kernel gap. + +--- + +## 3. Turning it on for H3 + +```python +from fastvideo.layers.quantization.w4a16_config import W4A16Config + +fastvideo_args.transformer_quant = W4A16Config.for_minimax_h3() +``` + +or by registry name once the registry entry is in place (section 8): + +```python +fastvideo_args.transformer_quant = "W4A16" +``` + +`transformer_quant` is pinned onto `pipeline_config.dit_config.quant_config` in +`FastVideoArgs._apply_transformer_quant`, so the DiT constructs its linears with +the config attached; the loader then materializes the 4-bit buffers from the +freshly loaded BF16 weights. **No pre-quantized checkpoint is required** — this +quantizes a plain BF16 checkpoint at load time, exactly like `NVFP4` and +`INT8Affine`. + +### Constructor options + +| Field | Default | Meaning | +|---|---|---| +| `group_size` | `64` | Values per `(scale, zero)` group along the input axis | +| `bits` | `4` | `4` (packed two-per-byte) or `8` (one per byte) | +| `target_layers` | `None` | Explicit allowlist of full module paths; takes precedence over `layer_suffixes` | +| `layer_suffixes` | generic attn/FFN set | Dot-boundary suffix rule, for models with no enumerable list | +| `exclude_substrings` | `()` | **Additional** never-quantize substrings — can only widen the deny list | +| `retain_original_weight` | `True` | Keep the dense bf16 `layer.weight` after conversion | + +| Method | Meaning | +|---|---| +| `W4A16Config()` | Generic attn/FFN suffix set, no model-specific list | +| `W4A16Config.for_minimax_h3()` | The 362-linear H3 profile below | + +`group_size` must divide the input dim of every targeted linear. The released H3 +config (hidden 5376, attention inner 7168, ffn 14336, adaln 2688) works for 32, +64 and 128; 64 is the default, matching the sibling INT8 lane. A layer whose +input dim does not divide cleanly is **skipped with a warning and runs dense** +rather than failing model construction. + +--- + +## 4. Which H3 layers are quantized + +`for_minimax_h3()` targets **362 linears**, derived from H3's real module names +by `minimax_h3_w4a16_prefixes()` (a function, not a literal — the profile is +generated, so an architecture change that breaks it fails in a test rather than +silently in production). + +| Scope | Suffix | Linears | +|---|---|---| +| `transformer_blocks.{0..49}` | `attn.to_q` | 50 | +| | `attn.to_k` | 50 | +| | `attn.to_v` | 50 | +| | `attn.to_out` | 50 | +| | `ff.fc_in` | 50 | +| | `ff.fc_out` | 50 | +| | `adaln_proj.linear` | 50 | +| `token_refiner.refiner_blocks.{0,1}` | `attn.{to_q,to_k,to_v,to_out}` | 8 | +| | `ff.{fc_in,fc_out}` | 4 | +| | **Total** | **362** | + +This is the bulk of H3's parameter count: attention QKV/output projections, the +SwiGLU FFN, the per-block AdaLN modulation projection, and the two text-refiner +blocks. + +### Excluded, and why + +| Excluded | Reason | +|---|---| +| `attn.to_gate_compress` | H3's VSA sparse-attention gate — see section 5 | +| `proj_in`, `proj_out` | In `MiniMaxH3Transformer3DModel._keep_in_fp32_modules` | +| `audio_proj_in`, `audio_proj_out` | Same fp32 keep set | +| `time_embedder` | Same fp32 keep set | +| `adaln_basis` | Global timestep-basis projector feeding every block; not in the target list | +| `context_embedder` | Text input projection — a different kind of module from the ffn/attn GEMMs, and quantizing the text stream while the video/audio input streams stay fp32 is an unvalidated asymmetry | +| `norm_out.linear` | Small, once per forward | + +The fp32-pinned modules are excluded **twice over** — they are absent from the +allowlist *and* named in the deny list — so a future widening of the allowlist +cannot silently reach them. + +--- + +## 5. `attn.to_gate_compress` must never be quantized + +`attn.to_gate_compress` is the compression gate on H3's VSA (VIDEO_SPARSE_ATTN_H3) +sparse-attention branch. It is not an ordinary linear, for two independent +reasons: + +1. **Its output steers a discrete decision.** The gate's output feeds the sparse + attention's tile selection. A quantized gate does not merely perturb the + output by a small amount — it can change *which tiles* the sparse attention + reads. That is a routing change, not a numerical one, and there is no + tolerance bound that makes it safe. +2. **H3 probes it structurally.** `MiniMaxH3Attention._gate_active()` tests the + loaded weight once (`bool((weight != 0).any())`) to skip a guaranteed-zero + branch. The gate is zero-initialized in the released checkpoint, so this + skip is exact and saves a full GEMM plus a third more all-to-all traffic per + layer. A dequantized weight is a different object than the one being probed. + +The name matches none of the usual `norm` / `embedder` / `scale_shift` exclusion +heuristics, so it is named explicitly in `_NEVER_QUANTIZE_SUBSTRINGS`. + +**How the exclusion is enforced** — it is fail-closed, not advisory: + +```python +self.exclude_substrings = tuple(_NEVER_QUANTIZE_SUBSTRINGS) + tuple(exclude_substrings or ()) +``` + +`_NEVER_QUANTIZE_SUBSTRINGS` is always present and `exclude_substrings` can +only *add*, never remove. `is_target_layer()` checks the deny list **before** +`target_layers`, so a caller who names the gate in their allowlist still gets +`False`. The layer then receives `UnquantizedLinearMethod` from +`LinearBase.__init__`'s `None` fallback and runs dense. + +This is verified against the real module, not just the string predicate: +`test_h3_gate_module_is_built_and_left_dense` forces VSA backend resolution, +constructs the actual `MiniMaxH3Attention`, and asserts +`to_gate_compress.quant_method` is `UnquantizedLinearMethod` while `to_q` and +`to_out` carry `W4A16QuantizeMethod`. + +--- + +## 6. Requirements, and what is NOT validated + +### Hardware + +| Requirement | Value | +|---|---| +| Minimum declared capability | 75 (Turing) — the reference path is a plain 16-bit GEMM, so nothing 4-bit is required *to load* | +| Intended deployment target | **Ada sm_89** — RTX 4090 24 GB, RTX 6000 Ada 48 GB | +| CUDA | Not required for the config itself; the reference path is pure PyTorch | + +### Kernel status — the honest version + +**There is no W4A16 GEMM kernel in this repository or in its installed +dependencies.** Concretely: + +- `fastvideo-kernel` ships **INT8** GEMM (`csrc/turbodiffusion/gemm/gemm.cu` + is `int8_gemm` over `cutlass::NumericConverter`), FP4 attention + for sm_100 / sm_120, and block-sparse attention. **No 4-bit weight GEMM on any + architecture**, Ada included. +- The `int4` matches in `csrc/turbodiffusion` are the CUDA 16-byte vector type + (`int4` / `uint4`), not 4-bit quantization. The same is true of the `uint4` + matches in the sm_100 block-sparse attention kernel. +- `autoawq`, `auto_gptq`, `gptqmodel`, `marlin`, `bitsandbytes`, `vllm`, + `flashinfer` and `compressed_tensors` are **not installed** in this + environment. +- The full-tree search for `w4a16` / `W4A16` across this worktree and the + sibling `FastVideo-*` trees returns only prose in + `H3_REALTIME_RELEASE_PLAN_2026-09-16.md` — no implementation anywhere. + +The only W4A16-adjacent code that already existed in the tree is vestigial: +`fastvideo/layers/linear.py` lists `"AWQMarlinLinearMethod"`, `"MarlinLinearMethod"`, +`"GPTQMarlinLinearMethod"`, `"HQQMarlinMethod"` and friends in +`WEIGHT_LOADER_V2_SUPPORTED`, and `fastvideo/models/parameter.py` documents +`PackedvLLMParameter` as *"Parameter for model weights which are packed on +disk. Example: GPTQ Marlin weights are int4 or int8, packed into int32."* +Both are inherited vLLM surface area: the strings are a name allowlist with no +implementation behind them, and `PackedvLLMParameter` / `PackedColumnParameter` +have **zero** construction sites outside `linear.py` and `parameter.py`. Neither +is a usable path, and this module does not depend on them. + +### What the reference path actually does + +`W4A16QuantizeMethod.apply()` dequantizes the full weight to the activation +dtype and calls `F.linear`. Consequences, stated plainly: + +- **Correctness**: exact. Applying the layer produces bit-identical output to + explicitly dequantizing the stored codes and running a dense GEMM (asserted in + the test suite). The only error is the quantization error itself, which is + bounded by half a quantizer step per element. +- **Memory**: the stored weight is genuinely ~3.2x smaller than BF16. But + dequantizing per forward means a transient dense weight exists during each + layer's forward. Steady-state savings are real; peak transient is not the + 4-bit number. +- **Speed**: **slower than BF16.** Every forward pays a full dequantize (a + multiply-add over the whole weight) on top of the same 16-bit GEMM BF16 would + have run. It does not use 4-bit tensor cores at all. +- **No memory/speed claim is validated on hardware.** Nothing in this lane has + been run on a 4090 or an RTX 6000 Ada. + +### What is tested + +`pytest fastvideo/tests/ops/quantization/test_w4a16_config.py` — 29 CPU-only +tests, no CUDA, no flashinfer, no kernel: + +- config imports without CUDA deps; `get_name`, capability, dtypes, `from_config` +- group-wise 4-bit round-trip within the derived half-step bound, at + `group_size` 32/64/128 and across bf16/fp16/fp32 +- exact round-trip for a group landing on the code grid +- packed layout: 4 bits/weight, codes are half the element count, low-nibble-first + order verified structurally, NaN does not poison a group +- H3 selection: 362 targets, the intended linears included, `to_gate_compress` + excluded **including on the real `MiniMaxH3Attention` module**, fp32-pinned + modules excluded, the deny list cannot be widened away, dot-boundary suffix + matching +- load-time conversion: non-persistent buffers do not leak into `state_dict`, + `apply` matches the documented reference exactly, lazy conversion off the + loader hook, the grad-enabled forward stays dense, opt-in weight purging, + bias handling, untargeted layers untouched + +### What is NOT tested or validated + +- **No GPU run of any kind.** No 4090, no RTX 6000 Ada, no Blackwell, no SLURM. +- **No end-to-end H3 generation.** Nothing here has produced a video. The + honest 4-bit question — *how much does group-64 W4A16 degrade H3's output?* — + is unmeasured. Expect it to need a QAD/QAT recovery pass, as the INT8 lane's + doc discusses for its own scheme. +- **No performance measurement.** There is no benchmark, and the reference path + is not a performance path by construction. +- **No fused kernel, no multi-GPU / sequence-parallel exercise.** +- **No distillation or finetune recovery.** + +--- + +## 7. Relationship to the other precision lanes + +| Lane | Registry name | Target hardware | Weight bits | Compute path | Status | +|---|---|---|---|---|---| +| BF16 | — | anything | 16 | dense | baseline | +| **W4A16** (this) | `W4A16` | **Ada sm_89** (4090, RTX 6000 Ada) | **4** | dequantize + dense 16-bit GEMM | **reference only — no kernel** | +| INT8 affine | `INT8Affine` | Ada / Turing+ (sm_75+) | 8 | dequantize + dense GEMM | reference path, real INT8 GEMM exists in `fastvideo-kernel` | +| FP8 / AbsMax FP8 | `FP8`, `AbsMaxFP8` | Ada+ (sm_89 has FP8 tensor cores) | 8 | FP8 GEMM | implemented | +| NVFP4 | `NVFP4` | **Blackwell sm_100+** | 4 | FlashInfer `mm_fp4` | implemented, real kernel | + +Reading the table: W4A16 and INT8 affine share a *shape* (load-time conversion, +weight-only, dequantize-then-GEMM reference) and differ in the arithmetic. W4A16 +vs NVFP4 is a hardware split, not a preference: NVFP4 is the Blackwell answer and +requires sm_100; W4A16 is the Ada answer and is the only 4-bit option there. If +you are deploying on a 5090, use NVFP4. If you are deploying on a 4090 or an +RTX 6000 Ada, W4A16 is the 4-bit lane — with the kernel caveat above. + +Do not confuse **RTX 6000 Ada** (sm_89, Ada) with **RTX PRO 6000 Blackwell** +(sm_120, Blackwell). Record the exact SKU when reporting this lane's results. + +--- + +## 8. Files + +| File | Contents | +|---|---| +| `fastvideo/layers/quantization/w4a16_config.py` | `W4A16Config`, `W4A16QuantizeMethod`, `convert_model_to_w4a16`, `w4a16_quantize` / `w4a16_dequantize`, `minimax_h3_w4a16_prefixes` | +| `fastvideo/tests/ops/quantization/test_w4a16_config.py` | 29 CPU-only unit tests | +| `fastvideo/layers/quantization/__init__.py` | Registry entry for the name `W4A16` (**owned elsewhere** — see below) | +| `fastvideo/models/loader/fsdp_load.py` | Loader dispatch to `convert_model_to_w4a16` (**owned elsewhere**) | + +### Wiring the lane in (two edits outside this module) + +The config is complete and self-sufficient, but two files owned by other work +need a branch for the load-time conversion to run eagerly instead of lazily: + +**1. `fastvideo/layers/quantization/__init__.py`** — add `"W4A16"` to the +`QuantizationMethods` literal, import `W4A16Config` in `get_quantization_config`'s +lazy import block, and add `"W4A16": W4A16Config` to `method_to_config`. + +**2. `fastvideo/models/loader/fsdp_load.py`** — in `_maybe_quantize_model`, add +the import and the `isinstance` branch, following the existing pattern: + +```python +from fastvideo.layers.quantization.w4a16_config import ( + W4A16QuantizeMethod, + convert_model_to_w4a16, +) +... +if isinstance(qm, W4A16QuantizeMethod): + convert_model_to_w4a16(model) + return +``` + +Without edit 2, inference is still **correct** — `apply()` converts lazily on +first forward and logs a warning — but the conversion happens later and the +warning is noise on every load. Without edit 1, `transformer_quant = "W4A16"` +raises `Invalid quantization method`; passing a `W4A16Config` instance directly +works either way. diff --git a/docs/quantization/loader_quant_params.md b/docs/quantization/loader_quant_params.md new file mode 100644 index 0000000000..e7e0f21936 --- /dev/null +++ b/docs/quantization/loader_quant_params.md @@ -0,0 +1,87 @@ +# Quantized models and the loader's zero-init allowlist + +Read this before adding a quantization config that registers new parameters. +Setting `engine.quantization.transformer_quant` to such a config and then +loading a checkpoint otherwise fails with: + +``` +ERROR fsdp_load.py Unsupported new parameter: transformer_blocks.0.attn.to_out.scale_input. +Allowed patterns: ['gate_compress', 'proj_l'] +``` + +## What the check actually guards + +`load_model_from_full_model_state_dict` in +`fastvideo/models/loader/fsdp_load.py` builds the sharded state dict from the +checkpoint, then computes + +```python +unused_keys = set(model.state_dict().keys()) - set(sharded_sd.keys()) +``` + +i.e. *parameters the instantiated model has but the checkpoint did not +provide*. Every such key is a potential silent failure: the loader +zero-initializes it, so a name-mapping bug, a renamed layer, or a checkpoint +from a different architecture would produce a model full of zeros that trains +and generates garbage instead of failing at load. + +So the loop raises `ValueError` for any unmapped key **except** a small +allowlist of names that are expected to be absent from every checkpoint. Note +that `strict=False` does *not* disable this — that flag only governs +checkpoint keys missing from the model, not model keys missing from the +checkpoint. + +Genuine mismatch detection is the point of the check, so the allowlist stays +narrow: `weight`, `bias`, and every other real model parameter still raise. + +## Why quantization configs register new parameters + +A quantized linear method (`create_weights`) builds its own scale tensors +alongside the packed weight. The checkpoint holds only the original +`weight`/`bias`, so the scales are always "new parameters" from the loader's +point of view. Examples: + +| Config | Registered by `create_weights` | In `state_dict()`? | +|---|---|---| +| `AbsMaxFP8` (`fastvideo/layers/quantization/absmax_fp8.py`) | `weight`, `scale_weight`, `scale_input` | yes | +| `INT8Affine` (`fastvideo/layers/quantization/int8_affine_config.py`) | `weight` only — codes/scales/biases are `persistent=False` buffers | no | +| `NVFP4` / `FP8` | `weight` plus `persistent=False` buffers | no | + +`persistent=False` buffers never enter `state_dict()`, so they never reach the +allowlist. Only configs that call `register_parameter` for scale tensors do. + +## Admitted names + +`ALLOWED_NEW_PARAM_PATTERNS` (module level in +`fastvideo/models/loader/fsdp_load.py`, checked by `is_allowed_new_param`) is a +substring match, and currently admits: + +| Pattern | Source | +|---|---| +| `gate_compress` | VSA gate tensor built by the attention backend | +| `proj_l` | SLA projection built by the attention backend | +| `scale_weight` | `AbsMaxFP8` per-tensor / per-merged-partition weight scale | +| `scale_input` | `AbsMaxFP8` per-tensor input (activation) scale | + +The two `scale_*` entries match the parameter leaf names the quantization +layer actually registers; they are not a generic `scale` wildcard, so a real +parameter that merely contains "scale" is still rejected. + +## Adding a new quantization config + +1. Build the model with the config enabled and read the `Unsupported new + parameter` error — it names the exact FQN. +2. If the offenders are scale/zero-point tensors registered with + `register_parameter`, add their **leaf name** (e.g. `scale_weight`) to + `ALLOWED_NEW_PARAM_PATTERNS` in `fastvideo/models/loader/fsdp_load.py`. + Add the specific names only; do not add a bare `scale` or similar token. +3. Prefer `register_buffer(..., persistent=False)` for values recomputed at + load time — such buffers need no allowlist entry at all. +4. Update `fastvideo/tests/ops/quantization/test_quant_param_allowlist.py`, + which asserts both that quant scales are accepted and that an unmapped real + weight still raises. + +`fastvideo/models/loader/shard_cache.py` keeps a mirrored +`_ALLOWED_NEW_PARAM_PATTERNS` tuple used to validate cache manifests. A +mismatch there never fails a run — it only disables the shard cache — but a +new quant config should update it too so the cache stays usable. diff --git a/docs/training/attn_qat.md b/docs/training/attn_qat.md index 8437da9cbd..1703e0399c 100644 --- a/docs/training/attn_qat.md +++ b/docs/training/attn_qat.md @@ -105,7 +105,7 @@ The migrated recipe preserves these behaviors: |---|---| | Student fake-quantized attention | `models.student.attention_backend: ATTN_QAT_TRAIN` | | Teacher and critic full-precision attention | Role-local `FLASH_ATTN` | -| Generator update every five critic steps | `method.generator_update_interval: 5` | +| Four critic-only steps, then one student-only step | `method.generator_update_interval: 5` | | Three-step rollout | `method.dmd_denoising_steps: [1000, 757, 522]` | | Score timestep range | `method.min_timestep_ratio: 0.02`, `max_timestep_ratio: 0.98` | | Legacy guidance `cond + 2(cond - uncond)` | Standard CFG scale `3.0` | diff --git a/examples/compacth3/README.md b/examples/compacth3/README.md new file mode 100644 index 0000000000..4d206954be --- /dev/null +++ b/examples/compacth3/README.md @@ -0,0 +1,54 @@ +# CompactH3 code index + +This branch consolidates the runnable code for the CompactH3 42-block and +34-block model lines. It intentionally excludes checkpoints, generated media, +W&B data, logs, paper drafts, and internal handoff reports. + +## Carried-forward recipes + +1. **Block selection:** score MiniMax-H3 blocks and select the activation-based + map. The uniform pruning experiments remain reproducible through the shared + scoring/selection tools, but are not release recipes. +2. **42-block recovery:** recover the activation-selected student in stages + (500 updates, a 200-update continuation, then selection at update 750). +3. **42-block DMD2:** initialize from the recovered update-750 parent and use + base H3 for both the frozen teacher and the trainable fake-score critic. The + selected corrected four-call checkpoint is update 1400. +4. **34-block recovery:** derive the 34-block student from the recovered + 42-block line and continue dense teacher-state recovery. No 34-block DMD2 + checkpoint is promoted yet. +5. **Quantization:** export BF16, INT8-affine, NVFP4, and W4A16 variants. The + NVFP4 QAD recipe keeps the audio projections out of FP4 and validates audio + and video separately. + +## Where things live + +- Block scoring, map selection, folding, recovery checks, and AV gates: + `scripts/fasth3_sprint/` +- Recovery configs: + `examples/train/configs/fasth3_*.yaml` +- Current 34-block recovery config: + `examples/train/configs/compacth3/release14b_recovery_wandb.yaml` +- Corrected four-call DMD2 and QAD configs: + `examples/train/configs/distribution_matching/minimax_h3/` +- Corrected DMD2 launch/export scripts: + `scripts/run_release20b_dmd2_v12_16gpu.sh`, + `scripts/submit_release20b_dmd2_v12_corrected_*.sbatch`, and + `scripts/checkpoint_conversion/export_h3_dmd2_student.py` +- Hardened recovery, checkpoint evaluation, QAD, and quantized export tools: + `scripts/compacth3/` +- AdaLN rank and checkpoint-sweep analysis: + `scripts/compacth3/analysis/` + +The checked-in SLURM launchers retain the cluster paths used by the runs so +their behavior is auditable. Override the root/checkpoint variables when +deploying elsewhere; no credentials are stored in this repository. + +## Code provenance + +The consolidation is applied as a squash on current upstream `main`. Recovery +utilities come from the verified recovery snapshot (`2e7fa15f`), corrected +DMD2 from the audited release snapshot (`c63ae5b`), and quantization/QAD support +from the quant-support snapshot (`236cd131`). Newer upstream H3 inference, +FastH3 V2 scheduling, LoRA, TAEH3, MXFP8, and multi-device loading changes are +retained. diff --git a/examples/distill/MiniMax-H3/distill_dmd.sh b/examples/distill/MiniMax-H3/distill_dmd.sh new file mode 100755 index 0000000000..9a863b2aaf --- /dev/null +++ b/examples/distill/MiniMax-H3/distill_dmd.sh @@ -0,0 +1,28 @@ +#!/usr/bin/env bash +# Direct torchrun wrapper for the current MiniMax-H3 DMD2 config. +# Use examples/train/slurm/dmd2_32xgb200.sbatch for production allocation. + +set -euo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +REPO_ROOT="$(cd "${SCRIPT_DIR}/../../.." && pwd)" +cd "${REPO_ROOT}" + +export MASTER_PORT="${MASTER_PORT:-29513}" +export FASTVIDEO_FA4="${FASTVIDEO_FA4:-1}" + +export NUM_GPUS="${NUM_GPUS:-4}" +WORLD_SIZE="${NUM_GPUS}" +SP_SIZE="${SP_SIZE:-1}" +HSDP_REPLICATE="${HSDP_REPLICATE:-1}" +HSDP_SHARD="${HSDP_SHARD:-${WORLD_SIZE}}" +CONFIG="${CONFIG:-examples/train/configs/distribution_matching/minimax_h3/dmd2_sp1_fsdp40_vidprom_v6.yaml}" +OUTPUT_DIR="${OUTPUT_DIR:-outputs/minimax_h3_dmd2_local}" + +exec bash examples/train/run.sh "${CONFIG}" \ + --training.distributed.num_gpus "${WORLD_SIZE}" \ + --training.distributed.sp_size "${SP_SIZE}" \ + --training.distributed.hsdp_replicate_dim "${HSDP_REPLICATE}" \ + --training.distributed.hsdp_shard_dim "${HSDP_SHARD}" \ + --training.checkpoint.output_dir "${OUTPUT_DIR}" \ + "$@" diff --git a/examples/inference/minimax_h3/README.md b/examples/inference/minimax_h3/README.md new file mode 100644 index 0000000000..0994e1d9dd --- /dev/null +++ b/examples/inference/minimax_h3/README.md @@ -0,0 +1,63 @@ +# MiniMax-H3 inference examples + +Basic single-request H3 examples live in `examples/inference/basic/` +(`basic_minimax_h3_t2v.py`, `basic_minimax_h3_fl2va.py`, +`basic_minimax_h3_ref2va.py`). This directory holds H3-specific benchmark +tooling. + +## `h3_vsa_dmd.py` — VSA-H3 vs dense attention, few-step DMD inference + +Benchmarks 3-step (DMD-style) H3 T2VA inference under two attention +backends and prints a latency/speedup table: + +- `dense` — `FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN` with FA4 + (`FASTVIDEO_FA4=1`). If the flash-attn package is not installed the + FLASH_ATTN request falls back to Torch SDPA (the worker log prints + "Using Torch SDPA backend"); the baseline is then SDPA, not FA4. +- `vsa` — `FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3` at + `--sparsity` (default 0.9), applied at generator boot through + `FastVideoArgs.VSA_sparsity` (`pipeline.experimental`). + `--vsa-tile-size {64,256}` (default 256) flows the same way + (`FastVideoArgs.VSA_tile_size`). At tile 256, `--vsa-kernel triton` + (default, no optional dependencies) uses the 256-to-64 expansion path + and `cutedsl` opts into the FA4 CuTe 256-tile forward (requires the + optional FA4 CuTe build, `flash_attn.cute`); at tile 64 the forward is + always the native 64-token Triton kernel and `--vsa-kernel` is ignored. +- `microbench` — model-free per-attention-layer proxy on the exact packed + H3 sequence geometry (dense FA4/SDPA vs the full `MiniMaxH3VSAImpl` + tile/pool/top-k/kernel/untile path). Useful standalone, and as the + speedup proxy when the full VSA pipeline leg is unavailable. + +Each mode boots its own generator in a fresh subprocess (the backend env +var is resolved at boot), runs `--warmup` untimed request(s), then times +`--num-prompts` requests with fixed seeds shared across modes so the +per-mode videos can be eyeballed against each other. Model-load time is +reported separately from per-request latency. A crash in one mode is +contained: its signature is saved to `//crash_signature.txt` +and the remaining modes still report. + +```bash +FASTVIDEO_FA4=1 python examples/inference/minimax_h3/h3_vsa_dmd.py \ + --model-path /path/to/MiniMax-H3 \ + --prompts-json /path/to/validation.json \ + --num-prompts 4 \ + --output-dir outputs/h3_vsa_dmd \ + --modes dense,vsa,microbench \ + --dmd-steps 1000,667,333 \ + --num-gpus 4 +``` + +`--prompts-json` expects `{"data": [{"caption": ...}]}`; without it a +built-in prompt set is used. Results land in +`//results.json`, per-mode videos in `//`, and +an aggregate `summary.json` plus a final table on stdout. + +### Caveat: dense-trained checkpoints under VSA + +With the base (dense-trained) H3 checkpoint this benchmark measures SPEED +only. The base model was never trained under VSA top-k masks, so at 90% +sparsity output-quality parity is not expected — judge quality with a +VSA-trained (sparse-student) DMD checkpoint. The 3-step DMD ladder applied +to the base checkpoint is likewise a latency proxy for a distilled +student, not a quality reference: real few-step quality requires a DMD +student checkpoint. diff --git a/examples/inference/minimax_h3/h3_vsa_dmd.py b/examples/inference/minimax_h3/h3_vsa_dmd.py new file mode 100644 index 0000000000..6440e47171 --- /dev/null +++ b/examples/inference/minimax_h3/h3_vsa_dmd.py @@ -0,0 +1,624 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Benchmark MiniMax-H3 few-step (DMD-style) inference: VSA-H3 vs dense attention. + +Runs the same prompts (same seeds) through the full T2VA pipeline once per +attention mode and reports per-request end-to-end latency plus the +denoising-stage time (``FASTVIDEO_STAGE_LOGGING=1``): + +- ``dense``: FLASH_ATTN with the FA4 CuTe kernels (``FASTVIDEO_FA4=1``). + When the flash-attn package is not installed, the FLASH_ATTN request + falls back to Torch SDPA (the worker log prints "Using Torch SDPA + backend") — the reported baseline is then SDPA, not FA4. +- ``vsa``: VIDEO_SPARSE_ATTN_H3 at ``--sparsity`` (default 0.9). The + sparsity is applied at generator boot via ``FastVideoArgs.VSA_sparsity`` + (``pipeline.experimental``); the H3 denoising stage builds per-step VSA + metadata from it. ``--vsa-tile-size`` (default 256) flows the same way + (``FastVideoArgs.VSA_tile_size``) and selects the tile geometry: at 256, + ``--vsa-kernel triton`` (default, no optional deps) uses the 256-to-64 + expansion path while ``cutedsl`` opts into the FA4 CuTe 256-tile forward + (requires the optional FA4 CuTe build, ``flash_attn.cute``); at 64 the + block map is already at kernel granularity, so the forward always runs + the native 64-token Triton kernel and ``--vsa-kernel`` does not apply. +- ``microbench``: model-free attention-layer microbenchmark on the exact + packed H3 sequence geometry of the requested video shape. Times + ``block_sparse_attn_256_bshd`` (Triton and, when importable, the FA4 CuTe + path) through the real ``MiniMaxH3VSAImpl`` tile/pool/top-k/untile path + against dense flash attention and torch SDPA. Use it as the speedup proxy + when the full VSA pipeline leg is unavailable. + +The attention backend is resolved at generator boot, so each mode runs in a +fresh subprocess (one generator boot per mode). This also isolates the modes +from each other's CUDA state: a crash in one leg still leaves the other +legs' numbers and the final table intact. + +Example (one 4-GPU node): + + FASTVIDEO_FA4=1 python examples/inference/minimax_h3/h3_vsa_dmd.py \\ + --model-path /path/to/MiniMax-H3 \\ + --prompts-json validation.json --num-prompts 4 \\ + --output-dir outputs/h3_vsa_dmd --modes dense,vsa,microbench + +Caveat — sparse-trained students: with the base (dense-trained) checkpoint +this benchmark measures SPEED only. The base model was never trained under +VSA top-k masks, so at 90% sparsity output-quality parity is not expected; +the per-mode videos are written for eyeballing, but judge quality with a +VSA-trained DMD student checkpoint. Likewise the 3-step DMD ladder applied +to the base checkpoint is a latency proxy for a distilled student, not a +quality reference. +""" + +from __future__ import annotations + +import argparse +import json +import os +import signal +import statistics +import subprocess +import sys +import threading +import time +from collections import deque +from pathlib import Path + +GENERATION_MODES = ("dense", "vsa") +ALL_MODES = GENERATION_MODES + ("microbench",) + +# Environment applied in the worker subprocess BEFORE importing fastvideo. +# The backend env var is folded into FastVideoArgs at generator boot. +MODE_ENV: dict[str, dict[str, str]] = { + "dense": { + "FASTVIDEO_ATTENTION_BACKEND": "FLASH_ATTN", + "FASTVIDEO_FA4": "1", + }, + "vsa": { + # Layers that do not support VSA-H3 (e.g. the token refiner) fall + # back to flash attention, so FA4 stays enabled here too. + "FASTVIDEO_ATTENTION_BACKEND": "VIDEO_SPARSE_ATTN_H3", + "FASTVIDEO_FA4": "1", + }, + "microbench": { + "FASTVIDEO_FA4": "1", + }, +} + +DEFAULT_PROMPTS = [ + "A cinematic drone shot over coastal cliffs at sunrise, golden light, gentle ocean waves, ultra detailed.", + "A barista pours latte art in a warm cafe, steam rising, shallow depth of field, soft morning light.", + "A red fox trots across fresh snow between pine trees, breath visible in the cold air, tracking shot.", + "Neon-lit rain-soaked city street at night, reflections on wet asphalt, pedestrians with umbrellas.", +] + +WARMUP_SEED = 999 +FIRST_SEED = 1000 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) + parser.add_argument("--model-path", required=True, help="Full modular MiniMax-H3 pipeline directory") + parser.add_argument("--prompts-json", default=None, help='Optional {"data": [{"caption": ...}]} prompt file') + parser.add_argument("--num-prompts", type=int, default=4, help="Timed requests per mode") + parser.add_argument("--sparsity", type=float, default=0.9, help="VSA sparsity for the vsa/microbench modes") + parser.add_argument("--output-dir", default="outputs/h3_vsa_dmd") + parser.add_argument("--modes", default="dense,vsa", help=f"Comma-separated subset of {ALL_MODES}") + parser.add_argument("--dmd-steps", default="1000,667,333", help="FASTVIDEO_DMD_DENOISING_STEPS ladder") + parser.add_argument("--dense-native-steps", + type=int, + default=None, + help="Run the dense mode on the scheduler's NATIVE n-step schedule instead of the " + "DMD ladder (FASTVIDEO_DMD_DENOISING_STEPS is removed for that mode). E.g. 50 turns " + "the dense leg into the teacher-style 50-step baseline, so the table compares " + "50-step dense against the few-step DMD VSA leg. Note the H3 scheduler builds an " + "n-point sigma grid ending at 0, i.e. n-1 transformer forwards") + parser.add_argument("--num-gpus", type=int, default=4) + parser.add_argument("--height", type=int, default=768) + parser.add_argument("--width", type=int, default=1344) + parser.add_argument("--num-frames", type=int, default=124) + parser.add_argument("--warmup", type=int, default=1, help="Untimed warm-up requests per mode") + parser.add_argument("--vsa-kernel", + choices=("triton", "cutedsl"), + default="triton", + help="VSA-256 kernel path; cutedsl needs the optional FA4 CuTe build " + "(ignored at --vsa-tile-size 64, which is native-Triton only)") + parser.add_argument("--vsa-tile-size", + type=int, + choices=(64, 256), + default=256, + help="VSA-H3 tile size in tokens, plumbed like sparsity via " + "FastVideoArgs.VSA_tile_size; 64 runs the native Triton block-sparse forward") + parser.add_argument("--mode-timeout", type=int, default=5400, help="Hard per-mode timeout in seconds") + parser.add_argument("--microbench-text-tokens", type=int, default=300, help="Assumed text prefix length") + parser.add_argument("--microbench-heads", default="14,56", help="Per-GPU head counts to microbench") + parser.add_argument("--_worker", choices=ALL_MODES, default=None, help=argparse.SUPPRESS) + return parser.parse_args() + + +def load_prompts(args: argparse.Namespace) -> list[str]: + if args.prompts_json: + with open(args.prompts_json) as handle: + rows = json.load(handle)["data"] + prompts = [row["caption"] for row in rows[:args.num_prompts]] + else: + prompts = [DEFAULT_PROMPTS[i % len(DEFAULT_PROMPTS)] for i in range(args.num_prompts)] + if len(prompts) < args.num_prompts: + raise ValueError(f"Requested {args.num_prompts} prompts but only {len(prompts)} available.") + return prompts + + +def dense_native_steps(mode: str, args: argparse.Namespace) -> int | None: + return args.dense_native_steps if (mode == "dense" and args.dense_native_steps) else None + + +def apply_worker_env(mode: str, args: argparse.Namespace) -> None: + """Set the mode's environment. Must run before any fastvideo import.""" + env = dict(MODE_ENV[mode]) + if dense_native_steps(mode, args): + # Teacher-style leg: the scheduler's own n-step schedule, no DMD + # ladder. The env var may leak in from the launch environment, so + # remove it explicitly (the H3 denoising stage reads it as a + # fallback when the pipeline config has no dmd_denoising_steps). + os.environ.pop("FASTVIDEO_DMD_DENOISING_STEPS", None) + # Per-step DMD_DEBUG stat lines double as in-log proof that all n + # native steps actually execute (two small latent stats per step — + # negligible next to a transformer forward at video resolutions). + env["FASTVIDEO_DMD_DEBUG_STATS"] = "1" + else: + env["FASTVIDEO_DMD_DENOISING_STEPS"] = args.dmd_steps + env["FASTVIDEO_STAGE_LOGGING"] = "1" # per-stage timings on the result object + env["FASTVIDEO_VSA_CUTEDSL"] = "1" if args.vsa_kernel == "cutedsl" else "0" + os.environ.update(env) + + +def denoise_seconds(result) -> float | None: + stages = getattr(getattr(result, "logging_info", None), "stages", None) + if not stages: + return None + for stage_name, metrics in stages.items(): + if "denois" in stage_name.lower(): + execution_time = metrics.get("execution_time") + if execution_time is not None: + return float(execution_time) + return None + + +def run_generation_worker(args: argparse.Namespace) -> int: + mode = args._worker + apply_worker_env(mode, args) + mode_dir = Path(args.output_dir) / mode + mode_dir.mkdir(parents=True, exist_ok=True) + prompts = load_prompts(args) + dmd_steps = [int(step) for step in args.dmd_steps.split(",") if step.strip()] + native_steps = dense_native_steps(mode, args) + num_inference_steps = native_steps or len(dmd_steps) + + from fastvideo import VideoGenerator + from fastvideo.api import ( + EngineConfig, + GenerationRequest, + GeneratorConfig, + OffloadConfig, + OutputConfig, + ParallelismConfig, + PipelineSelection, + SamplingConfig, + ) + + experimental: dict[str, float] = {} + if mode == "vsa": + # Boot-time run-level sparsity: the H3 denoising stage reads + # fastvideo_args.VSA_sparsity (mirrored onto ForwardBatch.VSA_sparsity + # per request) when building the per-step VSA metadata. The tile size + # rides the same experimental->FastVideoArgs path. + experimental["VSA_sparsity"] = args.sparsity + experimental["VSA_tile_size"] = args.vsa_tile_size + + schedule = (f"native {native_steps}-step schedule" if native_steps else f"dmd_steps={dmd_steps}") + print(f"[{mode}] booting generator (backend={os.environ['FASTVIDEO_ATTENTION_BACKEND']}, " + f"sparsity={experimental.get('VSA_sparsity', 0.0)}, " + f"tile={experimental.get('VSA_tile_size', '-')}, {schedule})", + flush=True) + boot_start = time.perf_counter() + generator = VideoGenerator.from_config( + GeneratorConfig( + model_path=args.model_path, + engine=EngineConfig( + num_gpus=args.num_gpus, + use_fsdp_inference=args.num_gpus > 1, + parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus), + offload=OffloadConfig( + dit=False, + dit_layerwise=False, + text_encoder=True, + vae=True, + pin_cpu_memory=False, + ), + ), + pipeline=PipelineSelection(experimental=experimental), + )) + load_time = time.perf_counter() - boot_start + print(f"[{mode}] generator ready in {load_time:.1f}s (model load, excluded from timings)", flush=True) + + def build_request(prompt: str, seed: int, output_name: str) -> GenerationRequest: + return GenerationRequest( + prompt=prompt, + negative_prompt="", + sampling=SamplingConfig( + height=args.height, + width=args.width, + num_frames=args.num_frames, + fps=24, + num_inference_steps=num_inference_steps, + guidance_scale=1.0, + batch_cfg=False, + seed=seed, + ), + output=OutputConfig( + output_path=str(mode_dir / output_name), + save_video=True, + return_frames=False, + ), + ) + + records: list[dict] = [] + try: + for warmup_index in range(args.warmup): + print(f"[{mode}] warm-up {warmup_index} (untimed; absorbs kernel JIT/autotune)", flush=True) + warmup_start = time.perf_counter() + generator.generate(build_request(prompts[0], WARMUP_SEED, f"warmup{warmup_index}.mp4")) + print(f"[{mode}] warm-up {warmup_index} done in {time.perf_counter() - warmup_start:.1f}s", flush=True) + + for index, prompt in enumerate(prompts): + request = build_request(prompt, FIRST_SEED + index, f"prompt{index:02d}.mp4") + request_start = time.perf_counter() + result = generator.generate(request) + e2e_seconds = time.perf_counter() - request_start + record = { + "prompt_index": index, + "seed": FIRST_SEED + index, + "e2e_seconds": e2e_seconds, + "generation_seconds": result.generation_time, + "denoise_seconds": denoise_seconds(result), + "video_path": result.video_path, + "peak_memory_mb": result.peak_memory_mb, + } + records.append(record) + parts = [f"[{mode}] {index:02d} e2e={e2e_seconds:.1f}s"] + if record["generation_seconds"] is not None: + parts.append(f"gen={record['generation_seconds']:.1f}s") + if record["denoise_seconds"] is not None: + parts.append(f"denoise={record['denoise_seconds']:.1f}s") + parts.append(f"-> {result.video_path}") + print(" ".join(parts), flush=True) + finally: + payload = { + "mode": mode, + "sparsity": args.sparsity if mode == "vsa" else 0.0, + "vsa_kernel": args.vsa_kernel if mode == "vsa" else None, + "vsa_tile_size": args.vsa_tile_size if mode == "vsa" else None, + "dmd_steps": None if native_steps else dmd_steps, + "num_inference_steps": num_inference_steps, + "shape": [args.height, args.width, args.num_frames], + "num_gpus": args.num_gpus, + "load_seconds": load_time, + "requests": records, + } + (mode_dir / "results.json").write_text(json.dumps(payload, indent=2)) + generator.shutdown() + return 0 + + +def _time_cuda_call(fn, warmup: int = 3, iters: int = 10) -> float: + import torch + for _ in range(warmup): + fn() + torch.cuda.synchronize() + start = time.perf_counter() + for _ in range(iters): + fn() + torch.cuda.synchronize() + return (time.perf_counter() - start) / iters + + +def run_microbench_worker(args: argparse.Namespace) -> int: + """Per-attention-layer proxy on the exact packed H3 geometry, single GPU. + + Measures one self-attention layer's compute (no sequence-parallel + all-to-all, which is identical across backends): dense flash attention + and SDPA on the true packed length vs the full VSA-H3 path + (tile scatter + fp32 tile pooling + top-k mask + block-sparse kernel + + untile) on the 256-padded tile buffer. + """ + mode = args._worker + apply_worker_env(mode, args) + + import torch + + from fastvideo.attention.backends.video_sparse_attn_h3 import (MiniMaxH3VSAImpl, MiniMaxH3VSAMetadataBuilder) + from fastvideo.pipelines.basic.minimax_h3.packing import (MINIMAX_H3_AUDIO_CHANNELS, audio_latent_num_frames, + video_latent_num_frames) + + device = torch.device("cuda:0") + torch.manual_seed(0) + head_dim = 128 + patch_size = (1, 2, 2) + spatial_ratio = 16 # H3 video VAE spatial compression + latent_frames = video_latent_num_frames(args.num_frames) + latent_height, latent_width = args.height // spatial_ratio, args.width // spatial_ratio + n_text = args.microbench_text_tokens + n_cond = 0 # T2V: no keyframe conditioning rows + n_audio = audio_latent_num_frames(args.num_frames) * MINIMAX_H3_AUDIO_CHANNELS + n_video = ((latent_frames // patch_size[0]) * (latent_height // patch_size[1]) * (latent_width // patch_size[2])) + seq_len = n_text + n_cond + n_audio + n_video + print(f"[microbench] packed H3 sequence for {args.height}x{args.width}x{args.num_frames}: " + f"text={n_text} (assumed) + cond={n_cond} + audio={n_audio} + video={n_video} = {seq_len} rows, " + f"head_dim={head_dim}", + flush=True) + + builder = MiniMaxH3VSAMetadataBuilder() + metadata_by_sparsity = { + sparsity: builder.build( + current_timestep=0, + raw_latent_shape=(latent_frames, latent_height, latent_width), + patch_size=patch_size, + VSA_sparsity=sparsity, + prefix_segments=(n_text, n_cond, n_audio), + device=device, + exempt=True, + tile_size=args.vsa_tile_size, + ) + for sparsity in (args.sparsity, 0.0) + } + reference_metadata = metadata_by_sparsity[args.sparsity] + print(f"[microbench] tiles: prefix={reference_metadata.num_prefix_tiles} " + f"video={reference_metadata.num_video_tiles} tile_elems={reference_metadata.tile_elems} " + f"padded_len={int(reference_metadata.variable_block_sizes.numel()) * reference_metadata.tile_elems}", + flush=True) + impl = MiniMaxH3VSAImpl(num_heads=0, head_size=head_dim, causal=False, softmax_scale=1.0, prefix="blocks.0.attn") + + flash_attn_func = None + fa_version = None + try: + from fastvideo.attention.utils import flash_attn_default + flash_attn_func = flash_attn_default.flash_attn_func + fa_version = flash_attn_default.fa_version + except Exception as error: # noqa: BLE001 - report and continue with SDPA only + print(f"[microbench] flash attention unavailable ({error}); dense rows fall back to SDPA only", flush=True) + + rows: list[dict] = [] + head_counts = [int(h) for h in args.microbench_heads.split(",") if h.strip()] + for num_heads in head_counts: + note = "per-rank slice of the sp=4 run" if num_heads == 14 else "full model on one GPU" + print(f"[microbench] heads={num_heads} ({note})", flush=True) + qkv = torch.randn(3, seq_len, num_heads, head_dim, device=device, dtype=torch.bfloat16) + query, key, value = (t.contiguous() for t in qkv.unbind(0)) + + def record(name: str, milliseconds: float, num_heads: int = num_heads) -> None: + rows.append({"heads": num_heads, "name": name, "ms": milliseconds}) + print(f"[microbench] {name:<34} {milliseconds:9.3f} ms/layer-call", flush=True) + + if flash_attn_func is not None: + def dense_flash(q=query[None], k=key[None], v=value[None]): + out = flash_attn_func(q, k, v) + return out[0] if isinstance(out, tuple) else out + + record(f"dense flash (FA{fa_version})", _time_cuda_call(dense_flash) * 1e3) + + def dense_sdpa(q=query[None], k=key[None], v=value[None]): + return torch.nn.functional.scaled_dot_product_attention( + q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)) + + record("dense torch SDPA", _time_cuda_call(dense_sdpa) * 1e3) + + kernel_choices = ["triton"] + if args.vsa_kernel == "cutedsl" and args.vsa_tile_size != 64: + # Tile 64 has no CuTe route (native Triton only) — a "cutedsl" + # row there would just re-measure the Triton path mislabeled. + kernel_choices.insert(0, "cutedsl") + for kernel in kernel_choices: + os.environ["FASTVIDEO_VSA_CUTEDSL"] = "1" if kernel == "cutedsl" else "0" + for sparsity, metadata in metadata_by_sparsity.items(): + def vsa_layer(qkv=qkv, metadata=metadata): + tiled = impl.preprocess_qkv(qkv, metadata) + q, k, v = tiled.chunk(3, dim=0) + out = impl.forward(q, k, v, None, metadata) + return impl.postprocess_output(out, metadata) + + label = f"VSA-H3 {kernel} t{args.vsa_tile_size} sparsity={sparsity:.2f}" + try: + record(label, _time_cuda_call(vsa_layer) * 1e3) + except Exception as error: # noqa: BLE001 - a kernel path may be uninstalled + print(f"[microbench] {label:<34} FAILED: {error}", flush=True) + rows.append({"heads": num_heads, "name": label, "ms": None, "error": str(error)}) + + mode_dir = Path(args.output_dir) / mode + mode_dir.mkdir(parents=True, exist_ok=True) + payload = { + "mode": mode, + "sparsity": args.sparsity, + "vsa_tile_size": args.vsa_tile_size, + "geometry": { + "seq_len": seq_len, + "text": n_text, + "cond": n_cond, + "audio": n_audio, + "video": n_video, + "head_dim": head_dim, + }, + "note": ("per-layer self-attention compute only, single GPU, excludes the sequence-parallel " + "all-to-all (identical across backends); heads=14 matches one rank of the 4-GPU sp run"), + "rows": rows, + } + (mode_dir / "results.json").write_text(json.dumps(payload, indent=2)) + return 0 + + +def run_mode_subprocess(mode: str, args: argparse.Namespace) -> dict: + """Run one mode in a fresh interpreter; stream output and survive crashes.""" + mode_dir = Path(args.output_dir) / mode + mode_dir.mkdir(parents=True, exist_ok=True) + command = [sys.executable, os.path.abspath(__file__), "--_worker", mode] + for key, value in vars(args).items(): + if key in ("_worker",) or value is None: + continue + command.extend([f"--{key.replace('_', '-')}", str(value)]) + child_env = dict(os.environ, PYTHONUNBUFFERED="1") + + print(f"\n=== mode {mode}: launching worker ===", flush=True) + start = time.perf_counter() + process = subprocess.Popen( + command, + env=child_env, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + start_new_session=True, + ) + + def kill_group() -> None: + print(f"=== mode {mode}: timeout after {args.mode_timeout}s, killing process group ===", flush=True) + try: + os.killpg(process.pid, signal.SIGKILL) + except ProcessLookupError: + pass + + watchdog = threading.Timer(args.mode_timeout, kill_group) + watchdog.start() + tail: deque[str] = deque(maxlen=60) + log_path = mode_dir / "worker.log" + with open(log_path, "w") as log_file: + assert process.stdout is not None + for line in process.stdout: + print(line, end="", flush=True) + log_file.write(line) + tail.append(line.rstrip("\n")) + return_code = process.wait() + watchdog.cancel() + elapsed = time.perf_counter() - start + + results_path = mode_dir / "results.json" + results = json.loads(results_path.read_text()) if results_path.exists() else None + status = {"mode": mode, "return_code": return_code, "elapsed_seconds": elapsed, "results": results} + if return_code != 0 or results is None: + signature = [line for line in tail if any(token in line for token in + ("Error", "error", "Traceback", "CUDA", "NCCL", "Signal", + "Segmentation", "terminate", "Killed"))] + status["crash_signature"] = signature[-15:] or list(tail)[-15:] + (mode_dir / "crash_signature.txt").write_text("\n".join(status["crash_signature"]) + "\n") + print(f"=== mode {mode}: FAILED (rc={return_code}); signature saved to {mode_dir / 'crash_signature.txt'} ===", + flush=True) + else: + print(f"=== mode {mode}: completed in {elapsed:.0f}s ===", flush=True) + return status + + +def _mode_stats(status: dict) -> dict | None: + results = status.get("results") + if not results or not results.get("requests"): + return None + requests = results["requests"] + generation = [r["generation_seconds"] for r in requests if r.get("generation_seconds") is not None] + denoise = [r["denoise_seconds"] for r in requests if r.get("denoise_seconds") is not None] + steps = results.get("num_inference_steps") + # Native schedule (dmd_steps is None): the H3 scheduler turns n inference + # steps into an n-point sigma grid ending at 0 = n-1 transformer forwards. + # The DMD ladder runs exactly one forward per ladder entry. + native = steps is not None and results.get("dmd_steps") is None + forwards = (steps - 1) if (native and steps > 1) else steps + mean_denoise = statistics.mean(denoise) if denoise else None + return { + "n": len(requests), + "steps": steps, + "load": results.get("load_seconds"), + "e2e": statistics.mean(r["e2e_seconds"] for r in requests), + "gen": statistics.mean(generation) if generation else None, + "denoise": mean_denoise, + "denoise_per_step": (mean_denoise / forwards) if mean_denoise is not None and forwards else None, + } + + +def _fmt(value: float | None, width: int, decimals: int = 1) -> str: + return f"{value:>{width}.{decimals}f}" if value is not None else f"{'-':>{width}}" + + +def _speedup(dense: dict | None, row: dict, metric: str) -> str: + if dense is None or dense.get(metric) is None or row.get(metric) in (None, 0): + return "-" + return f"{dense[metric] / row[metric]:.2f}x" + + +def summarize(statuses: list[dict], args: argparse.Namespace) -> None: + stats = {status["mode"]: _mode_stats(status) for status in statuses if status["mode"] in GENERATION_MODES} + dense = stats.get("dense") + + if args.dense_native_steps: + title = "H3 inference: 50-step-style dense baseline vs few-step DMD VSA" + schedule = (f"dense: native {args.dense_native_steps}-step schedule; " + f"vsa: dmd_steps={args.dmd_steps}") + else: + title = "H3 DMD 3-step inference: attention backend benchmark" + schedule = f"dmd_steps={args.dmd_steps}" + print(f"\n================ {title} ================") + vsa_kernel = "triton" if args.vsa_tile_size == 64 else args.vsa_kernel + print(f"shape={args.height}x{args.width}x{args.num_frames} gpus={args.num_gpus} " + f"{schedule} vsa sparsity={args.sparsity} (tile {args.vsa_tile_size}, {vsa_kernel})") + header = (f"{'mode':<12} {'n':>3} {'steps':>6} {'load(s)':>9} {'mean e2e(s)':>12} {'mean gen(s)':>12} " + f"{'mean denoise(s)':>16} {'denoise/step(s)':>16} {'e2e speedup':>12} {'denoise speedup':>16}") + print(header) + print("-" * len(header)) + for mode in ("dense", "vsa"): + label = f"vsa@{args.sparsity:.2f}" if mode == "vsa" else mode + row = stats.get(mode) + if row is None: + if any(status["mode"] == mode for status in statuses): + print(f"{label:<12} {'-':>3} {'-':>6} {'-':>9} {'CRASHED':>12} {'-':>12} {'-':>16} {'-':>16} " + f"{'-':>12} {'-':>16}") + continue + steps_text = str(row["steps"]) if row.get("steps") else "-" + print(f"{label:<12} {row['n']:>3} {steps_text:>6} {_fmt(row['load'], 9)} {_fmt(row['e2e'], 12)} " + f"{_fmt(row['gen'], 12)} {_fmt(row['denoise'], 16)} {_fmt(row['denoise_per_step'], 16, 2)} " + f"{_speedup(dense, row, 'e2e'):>12} {_speedup(dense, row, 'denoise'):>16}") + + micro = next((status for status in statuses if status["mode"] == "microbench"), None) + if micro and micro.get("results"): + results = micro["results"] + geometry = results["geometry"] + print(f"\nAttention-layer microbench (seq={geometry['seq_len']} rows: text {geometry['text']} + " + f"audio {geometry['audio']} + video {geometry['video']}; {results['note']}):") + for row in results["rows"]: + timing = f"{row['ms']:.3f} ms" if row.get("ms") is not None else f"FAILED: {row.get('error')}" + print(f" heads={row['heads']:>2} {row['name']:<34} {timing}") + if args.dense_native_steps: + print(f"\nNote: the native {args.dense_native_steps}-step schedule is a " + f"{args.dense_native_steps}-point sigma grid = {args.dense_native_steps - 1} transformer " + f"forwards; denoise/step divides by forwards ({args.dense_native_steps - 1} dense, " + f"{len([s for s in args.dmd_steps.split(',') if s.strip()])} vsa).") + print("\nNote: base checkpoint is dense-trained; the vsa leg measures speed, not quality parity.") + + +def main() -> None: + args = parse_args() + if args._worker in GENERATION_MODES: + sys.exit(run_generation_worker(args)) + if args._worker == "microbench": + sys.exit(run_microbench_worker(args)) + + modes = [mode.strip() for mode in args.modes.split(",") if mode.strip()] + unknown = sorted(set(modes) - set(ALL_MODES)) + if unknown: + raise ValueError(f"Unknown modes {unknown}; choose from {ALL_MODES}.") + + out_root = Path(args.output_dir) + out_root.mkdir(parents=True, exist_ok=True) + prompts = load_prompts(args) + (out_root / "prompts.txt").write_text("\n\n".join(f"[{i:02d}] {p}" for i, p in enumerate(prompts))) + + statuses = [run_mode_subprocess(mode, args) for mode in modes] + summarize(statuses, args) + summary_path = out_root / "summary.json" + summary_path.write_text(json.dumps(statuses, indent=2)) + print(f"\nPer-mode outputs and summary under: {out_root}") + sys.exit(0 if all(status["return_code"] == 0 and status["results"] is not None for status in statuses) else 2) + + +if __name__ == "__main__": + main() diff --git a/examples/train/configs/compacth3/release14b_recovery_wandb.yaml b/examples/train/configs/compacth3/release14b_recovery_wandb.yaml new file mode 100644 index 0000000000..aedc2dc058 --- /dev/null +++ b/examples/train/configs/compacth3/release14b_recovery_wandb.yaml @@ -0,0 +1,133 @@ +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/release-candidates/14b-backbone-r16-job7048-v1 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_base_recovery.MiniMaxH3BaseRecoveryMethod + denoising_weight: 0.0 + teacher_velocity_weight: 1.0 + video_velocity_weight: 1.0 + audio_velocity_weight: 4.0 + audio_seam_weight: 2.0 + feature_weight: 1.0 + teacher_grid_points: 50 + feature_local_block_indices: + - 4 + - 6 + - 7 + - 8 + - 9 + - 11 + - 12 + - 13 + - 29 + - 30 + - 31 + - 32 + modality_energy_floor: 0.001 + low_sigma_interval_fraction: 0.5 + low_sigma_interval_count: 12 + # Keep teacher-state inputs during recovery. The earlier 0.25 student-state + # exposure run amplified the folded student's audio and motion errors. + student_state_probability: 0.0 + modality_grad_probe_every: 0 + allow_prompt_only: true +training: + distributed: + num_gpus: 16 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 4 + hsdp_shard_dim: 4 + pin_cpu_memory: true + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/data/user-prompts-20260906/index-v3-audio-stratified + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260910 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 1.0e-05 + betas: + - 0.9 + - 0.999 + weight_decay: 0.0 + lr_scheduler: cosine + lr_warmup_steps: 100 + loop: + max_train_steps: 4000 + gradient_accumulation_steps: 4 + checkpoint: + output_dir: /path/bound/by/launcher + training_state_checkpointing_steps: 100 + use_cpu_process_group: true + checkpoints_total_limit: 3 + preserve_every_steps: 0 + preserve_steps: + - 100 + - 250 + - 500 + - 650 + - 700 + - 750 + - 800 + - 850 + - 900 + - 950 + - 1000 + - 1100 + - 1200 + - 1300 + - 1400 + - 1500 + - 2000 + - 4000 + - 5000 + - 6000 + resume_from_checkpoint: '' + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: folded-release-long-recovery + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + dit_precision: fp32 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 + validation: + _target_: fastvideo.train.callbacks.validation.ValidationCallback + pipeline_target: fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.MiniMaxH3Pipeline + dataset_file: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/job-scripts/release14b_validation_five.json + every_steps: 100 + run_at_start: false + sampling_steps: [50] + guidance_scale: 1.0 + num_frames: 124 + num_videos_per_prompt: 1 + use_validation_media_conditioning: false + offload_training_state: true + text_encoder_cpu_offload: true + vae_cpu_offload: true +pipeline: {} diff --git a/examples/train/configs/compacth3/release14b_validation_five.json b/examples/train/configs/compacth3/release14b_validation_five.json new file mode 100644 index 0000000000..8f8f06046e --- /dev/null +++ b/examples/train/configs/compacth3/release14b_validation_five.json @@ -0,0 +1,44 @@ +{ + "data": [ + { + "id": "speech_exact_chef", + "caption": "(S1) In a bright home kitchen, a chef looks straight at the camera and says [English] Fold the eggs gently and taste before you salt. A pot simmers behind her with soft bubbling and no music.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "music_vinyl_orbit", + "caption": "Close-up of a vinyl record spinning on a turntable in a dim listening room; warm analog jazz with double bass and brushed drums fills the room while the camera slowly orbits the platter.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "sfx_rain_tin_roof", + "caption": "Heavy rain hammers a tin roof over a porch swing; each droplet burst is crisp and close while thunder rolls in the distance and the wooden swing creaks in stereo.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "motion_kitesurf_carve", + "caption": "A kite surfer carves hard across choppy bay water while the camera dives alongside; spray hisses off the board edge, the sail flaps and snaps in the wind, and gulls cry overhead.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "multishot_train_diner", + "caption": "Two-shot scene: (shot 1) commuter train doors slide open with a pneumatic hiss and a station chime; (shot 2) cut to a diner interior where a waitress calls out an order while a grill sizzles.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + } + ] +} diff --git a/examples/train/configs/distribution_matching/minimax_h3/README.md b/examples/train/configs/distribution_matching/minimax_h3/README.md new file mode 100644 index 0000000000..df3e41f7ed --- /dev/null +++ b/examples/train/configs/distribution_matching/minimax_h3/README.md @@ -0,0 +1,45 @@ +# MiniMax-H3 DMD2 distillation configs + +Agent handoff: [`HANDOFF-h3-dmd2-vsa.md`](../../../../../HANDOFF-h3-dmd2-vsa.md). + +Few-step DMD2 distillation of the joint video/audio MiniMax-H3 transformer: +the v10 launch candidate is data-only over native-shape T2VA latents; earlier +recipes use a carried backward-simulation walk over the 4-step grid, optionally +mixed per batch with data-forced training (v9). The recipe history, comparison +scope, open issues, and launch preflight are tracked in +[`h3_dmd.md`](h3_dmd.md). + +SFT configs live in +[`../../fine_tuning/minimax_h3/`](../../fine_tuning/minimax_h3/). The production +Slurm launcher is +[`../../../slurm/dmd2_32xgb200.sbatch`](../../../slurm/dmd2_32xgb200.sbatch). + +## Files + +| File | Purpose | +|---|---| +| `dmd2_sp4_fsdp32_v12_datafree_mixed_dense_fa4.yaml` | V10.5-style data-free native-shape ablation on 32 GPUs/SP4 with a dense FA4 student; starts a fresh FastGen-aligned lineage. | +| `dmd2_sp1_fsdp64_v10_dataonly_mixed_vsa64.yaml` | **Launch candidate.** Data-only FastGen regime over five native-shape T2VA sources after filtering resolutions represented by fewer than 10 frozen videos, global batch 64 on 64 GPUs with full-world FSDP, regional compile for dense roles, and an eager VSA student. | +| `dmd2_sp1_fsdp64_v10_maxshape_gate_vsa64.yaml` | Two-step, 64-GPU critic/student capacity gate over the isolated `1760x768-362f` bucket; not a training lineage. | +| `dmd2_sp1_fsdp32_v10_dataonly_mixed_vsa64.yaml` | Historical 32-GPU job-2960 recipe; failed on the `1760x768-362f` critic backward and must not reuse the fsdp64 output namespace. | +| `dmd2_sp1_fsdp40_nuva_v9_dataforce_vsa64.yaml` | Previous v8 + per-batch data-forcing experiment over mixed prompt/latent data, batch 128 (accum 4). | +| `dmd2_sp1_fsdp40_vidprom_v8_bwdsim_vsa64.yaml` | Carry-only FastGen-parity recipe (data-free backward simulation, VSA-64 student). | +| `dmd2_sp1_fsdp40_vidprom_v7_vsa90.yaml` | Pre-parity 256-tile VSA recipe. | +| `dmd2_sp1_fsdp40_vidprom_v6.yaml` | Dense-student recipe: SP=1/full-shard, text-only simulate rollout, exclusive 4:1 cadence, FP32 compute boundaries. | +| `dmd2_sp1_fsdp40_vidprom.yaml` | Earlier hyperparameters under the current implementation. | +| `validation_wan64_h3.json` | Held-out prompts in H3's three-field format. | + +## Launch + +```bash +# Topology is derived from the allocation. +sbatch --nodes=10 examples/train/slurm/dmd2_32xgb200.sbatch + +# Select the config explicitly when a submit helper sets CONFIG. +CONFIG=examples/train/configs/distribution_matching/minimax_h3/dmd2_sp1_fsdp40_vidprom_v6.yaml \ + sbatch --nodes=10 examples/train/slurm/dmd2_32xgb200.sbatch +``` + +`FASTVIDEO_FA4=1` selects FA4 inside roles configured with `FLASH_ATTN`; the +launcher exports it. Do not set `FASTVIDEO_ATTENTION_BACKEND` globally because +attention backends are configured per role. diff --git a/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call.yaml b/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call.yaml new file mode 100644 index 0000000000..1d9708b315 --- /dev/null +++ b/examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call.yaml @@ -0,0 +1,140 @@ +# NVFP4 QAD for the 42-block 4-call DMD2 student. +# +# student : checkpoint-1400 of the CORRECTED re-run (job-paired8972-8975-4000-v3) +# teacher : frozen base H3 (base-h3-teacher-complete-v1) +# ladder : 4 calls -- dmd_denoising_steps [999, 749, 500, 250] +# decode : NOT taeh3. The decoder is downstream of the DiT, so keeping it out +# lets one QAD serve both the taeh3 preview and the full-VAE release. +# quant : nvfp4_qat_train -- FP4 forward, full-precision backward (STE). +# No weight conversion, so FSDP sharding/checkpointing stay dense-identical. +# +# AUDIO PROTECTION -- read before changing anything: +# * modality_loss_weights is carried over from the parent run unchanged, so the +# QAD does not silently rebalance video against audio. Upweight 'audio' here +# if the audio A/B regresses. +# * audio_proj_in / audio_proj_out are NOT in the FP4 target list (they match no +# DEFAULT_FP4_LAYERS entry), so they stay bf16. Do not add them. +# * The PR that added the NVFP4 encoder found the fully-quantized variant LOST THE +# VOICE TRACK. Audio is the first thing low precision breaks -- gate every QAD +# checkpoint on speech intelligibility, not on a combined scalar. +# * The validation panel below includes speech and music prompts on purpose. +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3/checkpoint-1400 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + quant_config: nvfp4_qat_train + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/release-candidates/base-h3-teacher-complete-v1 + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method + rollout_mode: simulate + rollout_carry: true + rollout_carry_slots: 8 + rollout_sample_type: ode + generator_update_interval: 5 + real_score_guidance_scale: 1.0 + dmd_denoising_steps: + - 999 + - 749 + - 500 + - 250 + min_timestep_ratio: 0.001 + max_timestep_ratio: 0.999 + score_timestep_shift: 2.4 + score_timestep_warp_max: 0.999 + score_timestep_continuous: true + fake_score_loss_space: x0 + modality_loss_weights: + video: 1.0 + audio: 1.0 + dmd_denom_floor_ratio: 0.05 + dmd_grad_cap: 100.0 + cfg_uncond: + text: zero + fake_score_learning_rate: 2.0e-06 + fake_score_betas: + - 0.9 + - 0.999 + fake_score_lr_scheduler: constant +training: + distributed: + num_gpus: 32 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 32 + pin_cpu_memory: true + data: + data_path: + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_50k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_5s_768p/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_mixed_res_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_fastgen_vidprom_150k/data + preprocessed_data_type: text_only + native_shape_bucketing: false + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 42 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 2.0e-06 + betas: + - 0.9 + - 0.999 + weight_decay: 0.01 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 200 + gradient_accumulation_steps: 8 + checkpoint: + output_dir: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-qad-nvfp4-4call-v1 + resume_from_checkpoint: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3/checkpoint-1400 + training_state_checkpointing_steps: 25 + require_complete_training_checkpoint: true + checkpointing_start_step: 1400 + checkpoints_total_limit: 12 + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: release20b-dmd2-qad-nvfp4-4call-v1 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 + validation: + _target_: fastvideo.train.callbacks.validation.ValidationCallback + pipeline_target: fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.MiniMaxH3Pipeline + dataset_file: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/code/release20b-dmd2-v12-corrected-v17/examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json + every_steps: 200 + run_at_start: false + sampling_steps: + - 4 + guidance_scale: 1.0 + use_record_dimensions: true + max_record_num_frames: 124 + num_videos_per_prompt: 1 + use_validation_media_conditioning: false + offload_training_state: true + unload_pipeline_after_validation: true + text_encoder_cpu_offload: true + vae_cpu_offload: true +model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + enable_torch_compile: false +dit_precision: fp32 +vsa: null diff --git a/examples/train/configs/distribution_matching/minimax_h3/release20b_dmd2_v12_dense.yaml b/examples/train/configs/distribution_matching/minimax_h3/release20b_dmd2_v12_dense.yaml new file mode 100644 index 0000000000..55b3c70f56 --- /dev/null +++ b/examples/train/configs/distribution_matching/minimax_h3/release20b_dmd2_v12_dense.yaml @@ -0,0 +1,145 @@ +# Release-20B adaptation of the successful MiniMax-H3 V12 dense DMD2 recipe. +# The launcher replaces both __SELECTED_PARENT__ values with the checkpoint +# selected by the matched-seed checkpoint-300/350/400/450 quality gate. +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: __SELECTED_PARENT__ + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + # The frozen teacher has no optimizer masters. FSDP already evaluates it + # in BF16, so BF16 construction is numerically identical at forward time + # and avoids retaining an unnecessary FP32 copy on every rank. + construction_precision: bf16 + disable_custom_init_weights: true + attention_backend: TORCH_SDPA + critic: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + # V12's fake-score model starts from the full base H3 distribution. A + # folded student critic is cheaper, but was the largest recipe drift in + # the rejected 2750-phase lineage. + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: true + disable_custom_init_weights: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + +method: + _target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method + rollout_mode: simulate + rollout_carry: true + rollout_carry_slots: 8 + rollout_sample_type: ode + generator_update_interval: 5 + real_score_guidance_scale: 1.0 + dmd_denoising_steps: [999, 749, 500, 250] + min_timestep_ratio: 0.001 + max_timestep_ratio: 0.999 + score_timestep_shift: 2.4 + score_timestep_warp_max: 0.999 + score_timestep_continuous: true + fake_score_loss_space: x0 + modality_loss_weights: + video: 1.0 + audio: 1.0 + cfg_uncond: + text: zero + fake_score_learning_rate: 2.0e-6 + fake_score_betas: [0.9, 0.999] + fake_score_lr_scheduler: constant + +training: + distributed: + num_gpus: 16 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 16 + pin_cpu_memory: true + data: + # Reuse the durable prompt-only corpus from the successful V12 run. The + # previous custom index disappeared between resumes, leaving an empty + # directory and killing job 8945. + data_path: + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_50k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_5s_768p/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_720_mixed_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_video_nuva_10k_mixed_res_len/data + - /mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/h3_t2av_fastgen_vidprom_150k/data + preprocessed_data_type: text_only + native_shape_bucketing: false + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 42 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 2.0e-6 + betas: [0.9, 0.999] + weight_decay: 0.01 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 1000 + # The runner overrides this to eight on the original 32-GPU V12 topology, + # preserving the recipe's global batch of 64. + gradient_accumulation_steps: 16 + checkpoint: + output_dir: __OUTPUT_DIR__ + resume_from_checkpoint: "" + save_inference_checkpoint_on_validation: true + inference_checkpoint_role: student + inference_checkpoint_dtype: bfloat16 + training_state_checkpointing_steps: 100 + require_complete_training_checkpoint: true + checkpointing_start_step: 100 + # Preserve intermediate peaks for the AV gate instead of assuming the + # last phase is the best phase. + # Smoke plus forty 100-phase gates through phase 4,000. Keep the complete + # trajectory so a middle quality peak cannot be pruned before AV review. + checkpoints_total_limit: 50 + tracker: + trackers: [wandb] + project_name: fasth3-14b-2step-qad-sprint + run_name: release20b-dmd2-v12-dense + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + enable_torch_compile: false + dit_precision: fp32 + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 + validation: + _target_: fastvideo.train.callbacks.validation.ValidationCallback + pipeline_target: fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline.MiniMaxH3Pipeline + dataset_file: examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json + every_steps: 100 + run_at_start: true + sampling_steps: [4] + guidance_scale: 1.0 + use_record_dimensions: true + max_record_num_frames: 124 + num_videos_per_prompt: 1 + use_validation_media_conditioning: false + offload_training_state: true + text_encoder_cpu_offload: true + vae_cpu_offload: true + +pipeline: + dit_config: + uniform_parameter_dtype: false diff --git a/examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json b/examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json new file mode 100644 index 0000000000..d3258fa757 --- /dev/null +++ b/examples/train/configs/distribution_matching/minimax_h3/release20b_validation_five.json @@ -0,0 +1,42 @@ +[ + { + "id": "speech_exact_chef", + "caption": "(S1) In a bright home kitchen, a chef looks straight at the camera and says [English] Fold the eggs gently and taste before you salt. A pot simmers behind her with soft bubbling and no music.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "music_vinyl_orbit", + "caption": "Close-up of a vinyl record spinning on a turntable in a dim listening room; warm analog jazz with double bass and brushed drums fills the room while the camera slowly orbits the platter.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "sfx_rain_tin_roof", + "caption": "Heavy rain hammers a tin roof over a porch swing; each droplet burst is crisp and close while thunder rolls in the distance and the wooden swing creaks in stereo.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "motion_kitesurf_carve", + "caption": "A kite surfer carves hard across choppy bay water while the camera dives alongside; spray hisses off the board edge, the sail flaps and snaps in the wind, and gulls cry overhead.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + }, + { + "id": "multishot_train_diner", + "caption": "Two-shot scene: (shot 1) commuter train doors slide open with a pneumatic hiss and a station chime; (shot 2) cut to a diner interior where a waitress calls out an order while a grill sizzles.", + "height": 480, + "width": 832, + "num_frames": 124, + "fps": 24 + } +] diff --git a/examples/train/configs/fasth3_14b_recovery.yaml b/examples/train/configs/fasth3_14b_recovery.yaml new file mode 100644 index 0000000000..4684e7bdd0 --- /dev/null +++ b/examples/train/configs/fasth3_14b_recovery.yaml @@ -0,0 +1,92 @@ +# FastH3 20-block four-call BF16 recovery. The SLURM launcher overrides the +# candidate, student attention backend, GPU topology, target step, and output. +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/checkpoints/h18-candidates/dense-activation + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA + +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_recovery.MiniMaxH3RecoveryMethod + modality_energy_floor: 1.0e-3 + denoising_weight: 1.0 + teacher_velocity_weight: 1.0 + feature_weight: 0.01 + feature_local_block_indices: [4, 9, 14, 19] + +training: + distributed: + num_gpus: 16 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 16 + pin_cpu_memory: true + + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-33b-20260806/data/h3_corpus + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260829 + num_latent_t: 37 + num_height: 512 + num_width: 896 + # 37 H3 video latents represent 124 frames and require 207 audio latents. + # Existing 37/200 artifacts must be recached; the model rejects them. + num_frames: 124 + + optimizer: + learning_rate: 1.0e-6 + betas: [0.9, 0.999] + weight_decay: 0.0 + lr_scheduler: constant + lr_warmup_steps: 0 + + loop: + max_train_steps: 200 + gradient_accumulation_steps: 1 + + checkpoint: + output_dir: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/h36-recovery/dense-activation + training_state_checkpointing_steps: 100 + checkpoints_total_limit: 2 + preserve_every_steps: 100 + preserve_steps: [200, 400, 600, 800, 1000] + resume_from_checkpoint: "" + + tracker: + trackers: [wandb] + project_name: fasth3-14b-2step-qad-sprint + run_name: h36-recovery-dense-activation-step200 + + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + + dit_precision: bf16 + +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 + ema: + _target_: fastvideo.train.callbacks.ema.EMACallback + decay: 0.9999 + start_iter: 0 + +pipeline: {} diff --git a/examples/train/configs/fasth3_base42_recovery.yaml b/examples/train/configs/fasth3_base42_recovery.yaml new file mode 100644 index 0000000000..cf2bbea2f3 --- /dev/null +++ b/examples/train/configs/fasth3_base42_recovery.yaml @@ -0,0 +1,91 @@ +# Paired full-schedule recovery; prompt-only zero latents are not valid targets. +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/checkpoints/base42-uniform-v1 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_base_recovery.MiniMaxH3BaseRecoveryMethod + denoising_weight: 1.0 + teacher_velocity_weight: 1.0 + feature_weight: 1.0 + teacher_grid_points: 50 + feature_local_block_indices: + - 3 + - 8 + - 13 + - 18 + - 24 + - 29 + - 34 + - 39 + modality_energy_floor: 0.001 +training: + distributed: + num_gpus: 4 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 4 + pin_cpu_memory: true + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/data/day1-mask-split/train + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260829 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 3.0e-05 + betas: + - 0.9 + - 0.999 + weight_decay: 0.0 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 2 + gradient_accumulation_steps: 1 + checkpoint: + output_dir: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/base42-preflight + training_state_checkpointing_steps: 2 + use_cpu_process_group: true + checkpoints_total_limit: 8 + preserve_every_steps: 25 + preserve_steps: + - 2 + - 25 + - 50 + - 100 + - 200 + resume_from_checkpoint: '' + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: base42-full-schedule-preflight + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + dit_precision: fp32 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 +pipeline: {} diff --git a/examples/train/configs/fasth3_detail_band_recovery.yaml b/examples/train/configs/fasth3_detail_band_recovery.yaml new file mode 100644 index 0000000000..39f2839962 --- /dev/null +++ b/examples/train/configs/fasth3_detail_band_recovery.yaml @@ -0,0 +1,96 @@ +# Detail-band recovery (audit 2026-09-08): 300 updates from the preserved +# activation-42 step-500 export with (a) low-sigma-biased interval sampling so +# the detail-forming steps of the shift-12/shift-3 grid are actually trained, +# (b) audio up-weighted to match its true gradient share, (c) the real paired +# denoising branch kept on as the ground-truth anchor. Arm B overrides +# method.student_state_probability to 0.25 from the launcher. +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/preserved/activation42-step500-job6878/export-500 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_base_recovery.MiniMaxH3BaseRecoveryMethod + denoising_weight: 1.0 + teacher_velocity_weight: 1.0 + video_velocity_weight: 1.0 + audio_velocity_weight: 4.0 + audio_seam_weight: 2.0 + feature_weight: 1.0 + teacher_grid_points: 50 + feature_local_block_indices: + - 7 + - 11 + - 13 + - 14 + modality_energy_floor: 0.001 + low_sigma_interval_fraction: 0.5 + low_sigma_interval_count: 12 + student_state_probability: 0.0 + modality_grad_probe_every: 0 + allow_prompt_only: false +training: + distributed: + num_gpus: 4 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 4 + pin_cpu_memory: true + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/data/day1-mask-split/train + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260829 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 3.0e-05 + betas: + - 0.9 + - 0.999 + weight_decay: 0.0 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 300 + gradient_accumulation_steps: 1 + checkpoint: + output_dir: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/base42-detail-band/job-placeholder + training_state_checkpointing_steps: 150 + use_cpu_process_group: true + checkpoints_total_limit: 3 + preserve_every_steps: 150 + preserve_steps: + - 300 + resume_from_checkpoint: '' + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: base42-detail-band + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + dit_precision: fp32 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 +pipeline: {} diff --git a/examples/train/configs/fasth3_detail_band_recovery34.yaml b/examples/train/configs/fasth3_detail_band_recovery34.yaml new file mode 100644 index 0000000000..58f543bd4b --- /dev/null +++ b/examples/train/configs/fasth3_detail_band_recovery34.yaml @@ -0,0 +1,101 @@ +# Detail-band recovery, 34-block release-candidate track (audit 2026-09-08): +# 300 updates from the job-6972 export (activation-34 step-0 + 200 prompt-58k) +# with (a) low-sigma-biased interval sampling so +# the detail-forming steps of the shift-12/shift-3 grid are actually trained, +# (b) audio up-weighted to match its true gradient share, (c) the real paired +# denoising branch kept on as the ground-truth anchor. Arm B overrides +# method.student_state_probability to 0.25 from the launcher. +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/activation34-prompt58k-recovery/job-6972/export-200 + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_base_recovery.MiniMaxH3BaseRecoveryMethod + denoising_weight: 1.0 + teacher_velocity_weight: 1.0 + video_velocity_weight: 1.0 + audio_velocity_weight: 4.0 + audio_seam_weight: 2.0 + feature_weight: 1.0 + teacher_grid_points: 50 + feature_local_block_indices: + - 4 + - 5 + - 6 + - 7 + - 8 + - 9 + - 10 + - 15 + modality_energy_floor: 0.001 + low_sigma_interval_fraction: 0.5 + low_sigma_interval_count: 12 + student_state_probability: 0.0 + modality_grad_probe_every: 0 + allow_prompt_only: false +training: + distributed: + num_gpus: 4 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 4 + pin_cpu_memory: true + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/data/day1-mask-split/train + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260829 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 3.0e-05 + betas: + - 0.9 + - 0.999 + weight_decay: 0.0 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 300 + gradient_accumulation_steps: 1 + checkpoint: + output_dir: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/activation34-detail-band/job-placeholder + training_state_checkpointing_steps: 150 + use_cpu_process_group: true + checkpoints_total_limit: 3 + preserve_every_steps: 150 + preserve_steps: + - 300 + resume_from_checkpoint: '' + tracker: + trackers: + - wandb + project_name: fasth3-14b-2step-qad-sprint + run_name: activation34-detail-band + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + dit_precision: fp32 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 +pipeline: {} diff --git a/examples/train/configs/fasth3_release_long_recovery.yaml b/examples/train/configs/fasth3_release_long_recovery.yaml new file mode 100644 index 0000000000..bfddb1d046 --- /dev/null +++ b/examples/train/configs/fasth3_release_long_recovery.yaml @@ -0,0 +1,91 @@ +# Long folded-backbone recovery. The launcher binds the exact selected student +# and pruning seams. The production launcher uses three DP replicas x four +# accumulation rounds (global effective batch 12) over the audio-stratified +# 58k prompt corpus; the in-file topology remains a 16-GPU reference maximum. +# Keep high-step recovery on Base-teacher prefix states. The controlled 7045 / +# 7046 comparison showed that 25% student prefixes from a half-recovered 6972 +# parent crushed audio brightness and voicing. Few-call inference-state +# exposure is handled later by DMD2/PDD, after the backbone clears quality +# gates, rather than teaching this recovery run to imitate its own errors. +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /path/bound/by/launcher + trainable: true + enable_gradient_checkpointing_type: full + attention_backend: TORCH_SDPA + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3Model + init_from: /mnt/nfs/vlm-aryan/hf-cache/hub/models--MiniMaxAI--MiniMax-H3/snapshots/42ed227ee7df40d41602854ae760620d6eb651fe + trainable: false + disable_custom_init_weights: true + attention_backend: TORCH_SDPA +method: + _target_: fastvideo.train.methods.knowledge_distillation.minimax_h3_base_recovery.MiniMaxH3BaseRecoveryMethod + denoising_weight: 0.0 + teacher_velocity_weight: 1.0 + video_velocity_weight: 1.0 + audio_velocity_weight: 4.0 + audio_seam_weight: 2.0 + feature_weight: 1.0 + teacher_grid_points: 50 + feature_local_block_indices: [] + modality_energy_floor: 0.001 + low_sigma_interval_fraction: 0.5 + low_sigma_interval_count: 12 + student_state_probability: 0.0 + modality_grad_probe_every: 0 + allow_prompt_only: true +training: + distributed: + num_gpus: 16 + sp_size: 4 + tp_size: 1 + hsdp_replicate_dim: 4 + hsdp_shard_dim: 4 + pin_cpu_memory: true + data: + data_path: /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/data/user-prompts-20260906/index-v3-audio-stratified + preprocessed_data_type: t2va + dataloader_num_workers: 0 + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 20260910 + num_latent_t: 37 + num_height: 480 + num_width: 832 + num_frames: 124 + optimizer: + learning_rate: 1.0e-05 + betas: [0.9, 0.999] + weight_decay: 0.0 + lr_scheduler: cosine + lr_warmup_steps: 100 + loop: + max_train_steps: 4000 + gradient_accumulation_steps: 4 + checkpoint: + output_dir: /path/bound/by/launcher + training_state_checkpointing_steps: 50 + use_cpu_process_group: true + checkpoints_total_limit: 3 + preserve_every_steps: 0 + preserve_steps: [100, 250, 500, 1000, 2000, 4000, 5000, 6000] + resume_from_checkpoint: '' + tracker: + trackers: [wandb] + project_name: fasth3-14b-2step-qad-sprint + run_name: folded-release-long-recovery + model: + precondition_outputs: false + enable_gradient_checkpointing_type: full + vsa: + sparsity: 0.0 + tile_size: 64 + cache_tile_buf: false + dit_precision: fp32 +callbacks: + grad_clip: + _target_: fastvideo.train.callbacks.grad_clip.GradNormClipCallback + max_grad_norm: 1.0 +pipeline: {} diff --git a/examples/train/configs/overfit_minimax_h3_t2va.yaml b/examples/train/configs/overfit_minimax_h3_t2va.yaml index ac29b37826..0a8df59486 100644 --- a/examples/train/configs/overfit_minimax_h3_t2va.yaml +++ b/examples/train/configs/overfit_minimax_h3_t2va.yaml @@ -83,4 +83,6 @@ callbacks: text_encoder_cpu_offload: true vae_cpu_offload: true -pipeline: {} +pipeline: + dit_config: + uniform_parameter_dtype: true diff --git a/examples/training/fasth3_14b_2step_qad/comparison_five_prompts.json b/examples/training/fasth3_14b_2step_qad/comparison_five_prompts.json new file mode 100644 index 0000000000..e972fd9dd3 --- /dev/null +++ b/examples/training/fasth3_14b_2step_qad/comparison_five_prompts.json @@ -0,0 +1,32 @@ +[ + { + "id": "speech_exact_presenter", + "category": "exact_speech", + "prompt": "(S1) In a quiet recording studio, a presenter looks directly into the camera and says [English] Fast video and clear audio arrive together. The camera is locked off, the face is well lit, and no music plays.", + "expected_audio": "Exact intelligible English sentence with visible lip motion and no music." + }, + { + "id": "motorcycle_tracking", + "category": "large_motion", + "prompt": "A motorcycle accelerates along a wet coastal road while the camera tracks beside it at speed. Water sprays from the tires, scenery moves rapidly, and the engine pitch rises with acceleration in stereo.", + "expected_audio": "Engine pitch and tire spray synchronized to acceleration and motion." + }, + { + "id": "mechanical_press", + "category": "visible_mechanical_sound", + "prompt": "Inside a clean workshop, a metal stamping press descends once onto a small steel plate. The visible impact produces one sharp metallic clang followed by a short machine hiss. Static three-quarter camera view.", + "expected_audio": "One impact clang and a short hiss synchronized to the press." + }, + { + "id": "two_shot_transition", + "category": "two_shot_transition", + "prompt": "A two-shot sequence: first, a wide view of a steam train entering a rural station with rhythmic wheel sounds; then a clean cut to a close view of the conductor opening a carriage door as steam hisses. Audio perspective changes naturally after the cut.", + "expected_audio": "Wheel rhythm before the cut and close steam hiss after it, with no desynchronization." + }, + { + "id": "glass_water_closeup", + "category": "fine_motion_and_materials", + "prompt": "A continuous close-up in a bright kitchen: a person slowly pours clear water from a glass pitcher into an empty drinking glass on a wooden counter. Fingers grip the handle naturally. The stream splashes and bubbles, the water level rises steadily, and sunlight refracts through the glass. Locked camera, no cuts, no music.", + "expected_audio": "Natural continuous pouring and splashing synchronized with the visible water stream, fading when pouring stops." + } +] diff --git a/examples/training/fasth3_14b_2step_qad/quick_gate_prompts.json b/examples/training/fasth3_14b_2step_qad/quick_gate_prompts.json new file mode 100644 index 0000000000..601f9c2055 --- /dev/null +++ b/examples/training/fasth3_14b_2step_qad/quick_gate_prompts.json @@ -0,0 +1,74 @@ +[ + { + "id": "speech_exact_presenter", + "category": "exact_speech", + "prompt": "(S1) In a quiet recording studio, a presenter looks directly into the camera and says [English] Fast video and clear audio arrive together. The camera is locked off, the face is well lit, and no music plays.", + "expected_audio": "Exact intelligible English sentence with visible lip motion and no music." + }, + { + "id": "speech_exact_barista", + "category": "exact_speech", + "prompt": "(S1) A barista behind a cafe counter smiles and says [English] Your coffee is ready by the window. Cups clink softly in the background while the camera holds a close medium shot.", + "expected_audio": "Exact intelligible English sentence, subtle cup sounds, and synchronized lips." + }, + { + "id": "dog_bark_visible", + "category": "animal_sound", + "prompt": "A golden retriever stands beside a red garden gate, looks toward the camera, and gives two distinct barks. Its mouth and chest movement visibly match each bark; birds remain faint in the distance.", + "expected_audio": "Two dog barks synchronized to the visible dog." + }, + { + "id": "mechanical_press", + "category": "visible_mechanical_sound", + "prompt": "Inside a clean workshop, a metal stamping press descends once onto a small steel plate. The visible impact produces one sharp metallic clang followed by a short machine hiss. Static three-quarter camera view.", + "expected_audio": "One impact clang and a short hiss synchronized to the press." + }, + { + "id": "violin_duet", + "category": "music", + "prompt": "Two violinists perform a gentle chamber duet on a small wooden stage. Their bow strokes are clearly visible and the stereo violin music follows the motion, with quiet room ambience and no speech.", + "expected_audio": "Coherent stereo violin duet synchronized to bow motion." + }, + { + "id": "silent_snow", + "category": "near_silence", + "prompt": "A wide locked shot of fresh snow falling over an empty field at dawn. Nothing moves except soft snowflakes and distant tree branches. The scene is nearly silent, with only a very faint winter breeze and no music or speech.", + "expected_audio": "Near silence without hiss, line noise, speech, or music." + }, + { + "id": "closeup_woman", + "category": "human_closeup", + "prompt": "Close-up portrait of a woman in warm window light listening thoughtfully, blinking naturally, then taking a quiet breath. Fine skin and eye detail, shallow depth of field, soft indoor room tone, no speech.", + "expected_audio": "Subtle room tone and breath without synthetic speech." + }, + { + "id": "closeup_man_laugh", + "category": "human_closeup", + "prompt": "Close-up portrait of an older man outdoors who breaks into a brief natural laugh. His eyes, cheeks, mouth, and shoulders move consistently; the laugh is intelligible and synchronized, with light park ambience.", + "expected_audio": "Brief synchronized natural laugh with park ambience." + }, + { + "id": "motorcycle_tracking", + "category": "large_motion", + "prompt": "A motorcycle accelerates along a wet coastal road while the camera tracks beside it at speed. Water sprays from the tires, scenery moves rapidly, and the engine pitch rises with acceleration in stereo.", + "expected_audio": "Engine pitch and tire spray synchronized to acceleration and motion." + }, + { + "id": "basketball_pan", + "category": "large_motion", + "prompt": "An athlete sprints across an indoor basketball court, catches a fast pass, and completes a powerful dunk as the camera pans quickly. Sneakers squeak, the rim rattles at impact, and the crowd reacts once.", + "expected_audio": "Squeaks, one rim impact, and crowd reaction synchronized to the action." + }, + { + "id": "two_shot_transition", + "category": "two_shot_transition", + "prompt": "A two-shot sequence: first, a wide view of a steam train entering a rural station with rhythmic wheel sounds; then a clean cut to a close view of the conductor opening a carriage door as steam hisses. Audio perspective changes naturally after the cut.", + "expected_audio": "Wheel rhythm before the cut and close steam hiss after it, with no desynchronization." + }, + { + "id": "onscreen_text", + "category": "onscreen_text", + "prompt": "A clean product-demo shot of a small electronic sign on a desk. The display clearly reads FAST H3 in large white capital letters while a hand presses one button and a single soft confirmation beep sounds. Locked camera, neutral background.", + "expected_audio": "One confirmation beep synchronized to the button press; on-screen text should read FAST H3." + } +] diff --git a/examples/training/fasth3_14b_2step_qad/release_sentinel_24.json b/examples/training/fasth3_14b_2step_qad/release_sentinel_24.json new file mode 100644 index 0000000000..2c37b372e0 --- /dev/null +++ b/examples/training/fasth3_14b_2step_qad/release_sentinel_24.json @@ -0,0 +1,26 @@ +[ + {"id":"speech_en_exact","category":"speech","prompt":"A well-lit presenter faces the locked camera in a silent studio and says [English] Fast video and clear audio arrive together. No music or background noise.","expected_audio":"Exact English sentence, clean voice, synchronized lips."}, + {"id":"speech_es_exact","category":"speech","prompt":"A woman at a quiet kitchen table looks into the camera and says [Spanish] El vaso está lleno de agua fría. No music plays.","expected_audio":"Exact intelligible Spanish sentence with synchronized lips."}, + {"id":"speech_zh_exact","category":"speech","prompt":"A man stands in a quiet library and says [Chinese] 今天的天气非常晴朗。 Static close shot, no music.","expected_audio":"Exact intelligible Mandarin sentence with synchronized lips."}, + {"id":"speech_ja_exact","category":"speech","prompt":"A woman at a train platform faces the camera and says [Japanese] 次の電車は三時に来ます。 The distant station remains quiet, with no music.","expected_audio":"Exact intelligible Japanese sentence with synchronized lips."}, + {"id":"duet_singing","category":"music","prompt":"Two singers perform a gentle acoustic duet on a small stage, trading one line each and then harmonizing. One guitar accompanies them; the audience stays quiet.","expected_audio":"Two distinct singing voices, stable harmony, and clean acoustic guitar."}, + {"id":"solo_violin","category":"music","prompt":"Close view of a violinist performing a slow lyrical melody in a dry rehearsal room. Bow direction and note attacks remain visible. No other instruments.","expected_audio":"Natural solo violin whose attacks follow the visible bow changes."}, + {"id":"drum_pattern","category":"music","prompt":"A drummer plays four measured hits: kick, snare, kick, then cymbal, with a clear pause between each. Locked front camera, no backing track.","expected_audio":"Exactly four distinct percussion events aligned with the visible strikes."}, + {"id":"piano_scale","category":"music","prompt":"Overhead close-up of two hands playing a rising eight-note piano scale, then stopping with both hands lifted. No speech and no room chatter.","expected_audio":"Eight clean rising piano notes that stop when the hands lift."}, + {"id":"paper_foley","category":"foley","prompt":"Extreme close-up of hands slowly folding crisp paper twice and tearing it once along the crease. Quiet room, no speech, no music.","expected_audio":"Two soft folds followed by one crisp synchronized tear."}, + {"id":"vegetable_chop","category":"foley","prompt":"Close-up of a chef making six evenly spaced knife cuts through a carrot on a wooden board, then placing the knife down. No music.","expected_audio":"Six dry synchronized chopping impacts and one softer knife placement."}, + {"id":"gravel_steps","category":"foley","prompt":"Low tracking shot of boots taking five slow steps across loose gravel and stopping. Wind is very faint; no speech or music.","expected_audio":"Five granular footstep crunches aligned to heel contact, then silence."}, + {"id":"zipper_fabric","category":"foley","prompt":"Close-up of a person slowly zipping a canvas backpack, tightening one strap, and setting it on a table. Quiet indoor room.","expected_audio":"Continuous zipper texture, brief fabric pull, and one soft table thump."}, + {"id":"balloon_pop","category":"av_sync","prompt":"A red balloon floats motionless. A visible pin touches it at exactly the middle of the shot and the balloon bursts once. Static camera, no music.","expected_audio":"One sharp pop exactly at visible rupture, with silence before and after."}, + {"id":"basketball_bounces","category":"av_sync","prompt":"Side view of a basketball dropped onto a gym floor. It bounces exactly three times, each bounce lower than the last, then rolls away.","expected_audio":"Three decreasing synchronized bounces followed by a quiet rolling sound."}, + {"id":"door_latch","category":"av_sync","prompt":"A hand turns a brass handle, opens a wooden door, and closes it until the latch clicks. The camera stays on the handle; no speech or music.","expected_audio":"Handle turn, hinge movement, closing thump, and final click aligned to motion."}, + {"id":"firework_single","category":"av_sync","prompt":"Night skyline with one firework launching, bursting once into a blue circle, and fading. No crowd and no background music.","expected_audio":"Launch whistle followed by one delayed boom, with no extra explosions."}, + {"id":"quiet_portrait","category":"silence_noise","prompt":"A silent locked portrait of a sleeping cat in a sunlit room. Only subtle breathing and curtain movement; explicitly no speech, music, buzzing, or hiss.","expected_audio":"Near-silence without synthetic hiss or unexpected events."}, + {"id":"snow_field","category":"silence_noise","prompt":"Wide static view of fresh snow falling over an empty field at dawn. No people, vehicles, animals, speech, or music.","expected_audio":"Very quiet natural ambience without voices, tones, or crackle."}, + {"id":"rain_window","category":"general_audio","prompt":"Continuous close-up of rain striking a window while distant traffic lights blur outside. The rainfall remains steady and no one speaks.","expected_audio":"Stable natural rain texture with faint distant traffic and no artifacts."}, + {"id":"ocean_waves","category":"general_audio","prompt":"Wide sunset beach view as three waves reach the shore in succession and foam recedes. Slow tripod pan, no people and no music.","expected_audio":"Three broad wave surges synchronized to shore contact, then receding foam."}, + {"id":"motorcycle_tracking","category":"motion","prompt":"A motorcycle accelerates along a wet coastal road while the camera tracks beside it at speed. Water sprays from the tires and the engine pitch rises in stereo.","expected_audio":"Engine pitch and tire spray synchronized to acceleration and motion."}, + {"id":"dog_frisbee","category":"motion","prompt":"A dog sprints across grass, leaps to catch a flying frisbee, lands, and runs back toward the camera in one continuous tracking shot.","expected_audio":"Footfalls, one brief jump effort, landing impact, and natural outdoor ambience."}, + {"id":"train_two_shot","category":"multishot","prompt":"First, a wide view of a steam train entering a rural station with rhythmic wheel sounds; then a clean cut to the conductor opening a carriage door as steam hisses.","expected_audio":"Wheel rhythm before the cut and close steam hiss after it, with perspective change."}, + {"id":"cafe_three_shot","category":"multishot","prompt":"Three-shot sequence in one cafe: espresso pours into a cup, milk is steamed, then the finished cup is set beside a customer. Each cut changes audio perspective naturally.","expected_audio":"Pour, steam, and cup placement in the correct shots without carryover artifacts."} +] diff --git a/examples/training/fasth3_14b_2step_qad/rescue_gate_prompts.json b/examples/training/fasth3_14b_2step_qad/rescue_gate_prompts.json new file mode 100644 index 0000000000..b8fd589f1d --- /dev/null +++ b/examples/training/fasth3_14b_2step_qad/rescue_gate_prompts.json @@ -0,0 +1,26 @@ +[ + { + "id": "speech_exact_presenter", + "category": "exact_speech", + "prompt": "(S1) In a quiet recording studio, a presenter looks directly into the camera and says [English] Fast video and clear audio arrive together. The camera is locked off, the face is well lit, and no music plays.", + "expected_audio": "Exact intelligible English sentence with visible lip motion and no music." + }, + { + "id": "motorcycle_tracking", + "category": "large_motion", + "prompt": "A motorcycle accelerates along a wet coastal road while the camera tracks beside it at speed. Water sprays from the tires, scenery moves rapidly, and the engine pitch rises with acceleration in stereo.", + "expected_audio": "Engine pitch and tire spray synchronized to acceleration and motion." + }, + { + "id": "mechanical_press", + "category": "visible_mechanical_sound", + "prompt": "Inside a clean workshop, a metal stamping press descends once onto a small steel plate. The visible impact produces one sharp metallic clang followed by a short machine hiss. Static three-quarter camera view.", + "expected_audio": "One impact clang and a short hiss synchronized to the press." + }, + { + "id": "two_shot_transition", + "category": "two_shot_transition", + "prompt": "A two-shot sequence: first, a wide view of a steam train entering a rural station with rhythmic wheel sounds; then a clean cut to a close view of the conductor opening a carriage door as steam hisses. Audio perspective changes naturally after the cut.", + "expected_audio": "Wheel rhythm before the cut and close steam hiss after it, with no desynchronization." + } +] diff --git a/examples/training/fasth3_14b_2step_qad/showcase_prompt.json b/examples/training/fasth3_14b_2step_qad/showcase_prompt.json new file mode 100644 index 0000000000..d1c217cefc --- /dev/null +++ b/examples/training/fasth3_14b_2step_qad/showcase_prompt.json @@ -0,0 +1,10 @@ +[ + { + "id": "spellblade_stand_back", + "category": "showcase_motion_speech_sync", + "expected_audio": "The English line 'Stand back!' is intelligible and synchronized, with storm wind, stone impacts, blade motion, and magical crackle matching visible events.", + "prompt": "integrated_multimodal_description: [Shot 1] A 16:9 single-take motion-graphics design in a vivid 2D anime action-fantasy style, with crisp cel-shaded planes, layered vector-like shapes, luminous energy streaks, and no readable text, opens on one lone sky-blue-haired spellblade standing on a shattered moonstone bridge above a violet storm. The spellblade pivots, raises a glowing crescent blade, and releases a high-speed spiral of gold-and-cyan energy that tears through incoming shadow shards; (S1) shouts [English]Stand back! while bracing against the magical recoil, with a sharp gasp and strained breath audible. After the action begins, the camera makes one rapid, smooth forward push-in toward the spellblade's determined face as the energy spiral fills the frame, maintaining the same continuous take through 00:05.000. Keep the spellblade as the only principal character, with no additional characters, logos, captions, subtitles, or readable signage.\n\noverall_soundscape: Violet storm wind roars around the bridge while stone fragments clatter, magical energy crackles, and the blade slices through the air. The spellblade's sharp gasp, strained breath, and shouted line remain clearly synchronized with the action.\n\nnon_diegetic_music: N/A", + "source_id": "t2va-2026082050-000102", + "source_seed": 2026082050 + } +] diff --git a/fastvideo/attention/utils/flash_attn_cute.py b/fastvideo/attention/utils/flash_attn_cute.py index ca38539924..ea9bb3c5c7 100644 --- a/fastvideo/attention/utils/flash_attn_cute.py +++ b/fastvideo/attention/utils/flash_attn_cute.py @@ -108,6 +108,10 @@ def _flash_attn_cute_forward( softcap=0.0, num_splits=1, pack_gqa=None, + # AOTAutograd executes its compiled forward with detached primals, so + # FA4 cannot infer from ``requires_grad`` that backward will need LSE. + # Keep the auxiliary tensor explicit in the custom-op contract. + return_lse=True, )[:2] return out, lse @@ -132,11 +136,60 @@ def _flash_attn_cute_setup_context(ctx: torch.autograd.function.FunctionCtx, inp q, k, v, softmax_scale, causal, deterministic = inputs out, lse = output ctx.save_for_backward(q, k, v, out, lse) + ctx.mark_non_differentiable(lse) ctx.softmax_scale = softmax_scale ctx.causal = causal ctx.deterministic = deterministic +@torch.library.custom_op( + "fastvideo::_flash_attn_cute_backward", + mutates_args=(), + device_types="cuda", +) +def _flash_attn_cute_backward_op( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + out: torch.Tensor, + grad_out: torch.Tensor, + lse: torch.Tensor, + softmax_scale: float | None, + causal: bool, + deterministic: bool, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + return _flash_attn_bwd( + q, + k, + v, + out, + grad_out, + lse, + softmax_scale=softmax_scale, + causal=causal, + softcap=0.0, + window_size_left=None, + window_size_right=None, + deterministic=deterministic, + ) + + +@torch.library.register_fake("fastvideo::_flash_attn_cute_backward") +def _flash_attn_cute_backward_fake( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + out: torch.Tensor, + grad_out: torch.Tensor, + lse: torch.Tensor, + softmax_scale: float | None, + causal: bool, + deterministic: bool, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + del out, grad_out, lse, softmax_scale, causal, deterministic + return torch.empty_like(q), torch.empty_like(k), torch.empty_like(v) + + def _flash_attn_cute_backward( ctx: torch.autograd.function.FunctionCtx, grad_out: torch.Tensor, @@ -144,19 +197,16 @@ def _flash_attn_cute_backward( ): del grad_lse q, k, v, out, lse = ctx.saved_tensors - dq, dk, dv = _flash_attn_bwd( + dq, dk, dv = torch.ops.fastvideo._flash_attn_cute_backward( q, k, v, out, grad_out, lse, - softmax_scale=ctx.softmax_scale, - causal=ctx.causal, - softcap=0.0, - window_size_left=None, - window_size_right=None, - deterministic=ctx.deterministic, + ctx.softmax_scale, + ctx.causal, + ctx.deterministic, ) return dq, dk, dv, None, None, None @@ -200,6 +250,7 @@ def _flash_attn_cute_varlen_forward( softcap=0.0, num_splits=1, pack_gqa=None, + return_lse=True, )[:2] return out, lse @@ -241,6 +292,7 @@ def _flash_attn_cute_varlen_setup_context(ctx: torch.autograd.function.FunctionC ) = inputs out, lse = output ctx.save_for_backward(q, k, v, out, lse, cu_seqlens_q, cu_seqlens_k) + ctx.mark_non_differentiable(lse) ctx.max_seqlen_q = max_seqlen_q ctx.max_seqlen_k = max_seqlen_k ctx.softmax_scale = softmax_scale @@ -248,6 +300,67 @@ def _flash_attn_cute_varlen_setup_context(ctx: torch.autograd.function.FunctionC ctx.deterministic = deterministic +@torch.library.custom_op( + "fastvideo::_flash_attn_cute_varlen_backward", + mutates_args=(), + device_types="cuda", +) +def _flash_attn_cute_varlen_backward_op( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + out: torch.Tensor, + grad_out: torch.Tensor, + lse: torch.Tensor, + cu_seqlens_q: torch.Tensor, + cu_seqlens_k: torch.Tensor, + max_seqlen_q: int, + max_seqlen_k: int, + softmax_scale: float | None, + causal: bool, + deterministic: bool, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + return _flash_attn_bwd( + q, + k, + v, + out, + grad_out, + lse, + softmax_scale=softmax_scale, + causal=causal, + softcap=0.0, + window_size_left=None, + window_size_right=None, + cu_seqlens_q=cu_seqlens_q, + cu_seqlens_k=cu_seqlens_k, + max_seqlen_q=max_seqlen_q, + max_seqlen_k=max_seqlen_k, + deterministic=deterministic, + ) + + +@torch.library.register_fake("fastvideo::_flash_attn_cute_varlen_backward") +def _flash_attn_cute_varlen_backward_fake( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + out: torch.Tensor, + grad_out: torch.Tensor, + lse: torch.Tensor, + cu_seqlens_q: torch.Tensor, + cu_seqlens_k: torch.Tensor, + max_seqlen_q: int, + max_seqlen_k: int, + softmax_scale: float | None, + causal: bool, + deterministic: bool, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + del out, grad_out, lse, cu_seqlens_q, cu_seqlens_k + del max_seqlen_q, max_seqlen_k, softmax_scale, causal, deterministic + return torch.empty_like(q), torch.empty_like(k), torch.empty_like(v) + + def _flash_attn_cute_varlen_backward( ctx: torch.autograd.function.FunctionCtx, grad_out: torch.Tensor, @@ -255,23 +368,20 @@ def _flash_attn_cute_varlen_backward( ): del grad_lse q, k, v, out, lse, cu_seqlens_q, cu_seqlens_k = ctx.saved_tensors - dq, dk, dv = _flash_attn_bwd( + dq, dk, dv = torch.ops.fastvideo._flash_attn_cute_varlen_backward( q, k, v, out, grad_out, lse, - softmax_scale=ctx.softmax_scale, - causal=ctx.causal, - softcap=0.0, - window_size_left=None, - window_size_right=None, - cu_seqlens_q=cu_seqlens_q, - cu_seqlens_k=cu_seqlens_k, - max_seqlen_q=ctx.max_seqlen_q, - max_seqlen_k=ctx.max_seqlen_k, - deterministic=ctx.deterministic, + cu_seqlens_q, + cu_seqlens_k, + ctx.max_seqlen_q, + ctx.max_seqlen_k, + ctx.softmax_scale, + ctx.causal, + ctx.deterministic, ) return dq, dk, dv, None, None, None, None, None, None, None diff --git a/fastvideo/configs/models/dits/minimax_h3.py b/fastvideo/configs/models/dits/minimax_h3.py index 464f75a699..4abed9ebf6 100644 --- a/fastvideo/configs/models/dits/minimax_h3.py +++ b/fastvideo/configs/models/dits/minimax_h3.py @@ -82,6 +82,5 @@ class MiniMaxH3Config(DiTConfig): arch_config: MiniMaxH3ArchConfig = field(default_factory=MiniMaxH3ArchConfig) prefix: str = "minimax_h3" - # FastVideo's Fully Sharded Data Parallel (FSDP) loading path requires one - # parameter dtype, while H3 inference keeps boundary projections in FP32. + # Disable model-selected FP32 compute groups when uniform precision is required. uniform_parameter_dtype: bool = False diff --git a/fastvideo/dataset/parquet_dataset_map_style.py b/fastvideo/dataset/parquet_dataset_map_style.py index f99937f75e..4d09493d71 100644 --- a/fastvideo/dataset/parquet_dataset_map_style.py +++ b/fastvideo/dataset/parquet_dataset_map_style.py @@ -2,7 +2,9 @@ import os import pickle import random +from collections import defaultdict from collections.abc import Sequence +from pathlib import Path from typing import Any import pyarrow as pa @@ -15,6 +17,7 @@ from torchdata.stateful_dataloader import StatefulDataLoader from fastvideo.platforms import current_platform +from fastvideo.dataset.shape_bucket import parse_video_shape_bucket_id from fastvideo.dataset.utils import collate_rows_from_parquet_schema from fastvideo.distributed import (get_sp_world_size, get_world_group, get_world_rank, get_world_size) from fastvideo.logger import init_logger @@ -37,6 +40,7 @@ def __init__( drop_last: bool = True, drop_first_row: bool = False, seed: int = 0, + sample_bucket_ids: Sequence[str] | None = None, ): self.batch_size = batch_size self.dataset_size = dataset_size @@ -47,30 +51,40 @@ def __init__( self.sp_world_size = sp_world_size # ── epoch-level RNG ──────────────────────────────────────────────── + if batch_size <= 0 or num_sp_groups <= 0 or sp_world_size <= 0: + raise ValueError("batch_size, num_sp_groups, and sp_world_size must be positive") + rng = torch.Generator().manual_seed(self.seed) - # Create a random permutation of all indices - global_indices = torch.randperm(self.dataset_size, generator=rng) - - if drop_first_row: - # drop 0 in global_indices - global_indices = global_indices[global_indices != 0] - self.dataset_size = self.dataset_size - 1 - - if self.drop_last: - # For drop_last=True, we: - # 1. Ensure total samples is divisible by (batch_size * num_sp_groups) - # 2. This guarantees each SP group gets same number of complete batches - # 3. Prevents uneven batch sizes across SP groups at end of epoch - num_batches = self.dataset_size // self.batch_size - num_global_batches = num_batches // self.num_sp_groups - global_indices = global_indices[:num_global_batches * self.num_sp_groups * self.batch_size] - else: - if self.dataset_size % (self.num_sp_groups * self.batch_size) != 0: - # add more indices to make it divisible by (batch_size * num_sp_groups) + if sample_bucket_ids is None: + # Legacy behavior: one permutation over the complete dataset. + global_indices = torch.randperm(self.dataset_size, generator=rng) + if drop_first_row: + global_indices = global_indices[global_indices != 0] + self.dataset_size -= 1 + + if self.drop_last: + num_batches = self.dataset_size // self.batch_size + num_global_batches = num_batches // self.num_sp_groups + global_indices = global_indices[:num_global_batches * self.num_sp_groups * self.batch_size] + elif self.dataset_size % (self.num_sp_groups * self.batch_size) != 0: padding_size = self.num_sp_groups * self.batch_size - (self.dataset_size % (self.num_sp_groups * self.batch_size)) logger.info("Padding the dataset from %d to %d", self.dataset_size, self.dataset_size + padding_size) global_indices = torch.cat([global_indices, global_indices[:padding_size]]) + self.bucket_schedule: tuple[str, ...] | None = None + self.bucket_padding: dict[str, int] | None = None + self.num_padded_samples = 0 + else: + if len(sample_bucket_ids) != self.dataset_size: + raise ValueError("sample_bucket_ids must contain one identifier per dataset row, got " + f"{len(sample_bucket_ids)} for dataset_size={self.dataset_size}") + global_indices, self.bucket_schedule = self._bucketed_global_schedule( + sample_bucket_ids, + rng=rng, + drop_first_row=drop_first_row, + ) + if drop_first_row: + self.dataset_size -= 1 # shard the indices to each sp group ith_sp_group = self.global_rank // self.sp_world_size @@ -78,6 +92,70 @@ def __init__( self.sp_group_local_indices = sp_group_local_indices logger.info("Dataset size for each sp group: %d", len(sp_group_local_indices)) + def _bucketed_global_schedule( + self, + sample_bucket_ids: Sequence[str], + *, + rng: torch.Generator, + drop_first_row: bool, + ) -> tuple[torch.Tensor, tuple[str, ...]]: + """Build same-shape rounds shared by every data-parallel group. + + One round contains ``num_sp_groups * batch_size`` rows from exactly + one bucket. Strided DP sharding below gives every group a distinct + local batch while preserving the same bucket at that microstep. SP + ranks map to the same group and therefore receive identical indices. + """ + by_bucket: dict[str, list[int]] = defaultdict(list) + for index, bucket_id in enumerate(sample_bucket_ids): + if drop_first_row and index == 0: + continue + parse_video_shape_bucket_id(bucket_id) + by_bucket[bucket_id].append(index) + + samples_per_round = self.num_sp_groups * self.batch_size + rounds: list[torch.Tensor] = [] + round_bucket_ids: list[str] = [] + bucket_padding: dict[str, int] = {} + padded = 0 + for bucket_id in sorted(by_bucket): + bucket_indices = torch.tensor(by_bucket[bucket_id], dtype=torch.long) + bucket_indices = bucket_indices[torch.randperm(len(bucket_indices), generator=rng)] + remainder = len(bucket_indices) % samples_per_round + if remainder: + # Native bucketing must retain every frozen row, including a + # rare bucket smaller than one global microbatch. Repeat only + # within that bucket; legacy drop_last behavior remains in the + # unbucketed branch above. + padding_size = samples_per_round - remainder + repeats = (padding_size + len(bucket_indices) - 1) // len(bucket_indices) + padding = bucket_indices.repeat(repeats)[:padding_size] + bucket_indices = torch.cat((bucket_indices, padding)) + padded += padding_size + bucket_padding[bucket_id] = padding_size + logger.info( + "Exact-shape bucket %s has %d row(s); repeated %d row(s) to fill global microbatches of %d", + bucket_id, + len(by_bucket[bucket_id]), + padding_size, + samples_per_round, + ) + bucket_rounds = list(bucket_indices.reshape(-1, samples_per_round).unbind(0)) + rounds.extend(bucket_rounds) + round_bucket_ids.extend([bucket_id] * len(bucket_rounds)) + + if not rounds: + raise ValueError("Exact-shape bucketing requires at least one dataset row") + round_order = torch.randperm(len(rounds), generator=rng).tolist() + self.bucket_padding = bucket_padding + self.num_padded_samples = padded + if padded: + logger.info("Exact-shape bucketing repeated %d row(s) to fill bucket-local global microbatches", padded) + return ( + torch.cat([rounds[index] for index in round_order]), + tuple(round_bucket_ids[index] for index in round_order), + ) + def __iter__(self): indices = self.sp_group_local_indices for i in range(0, len(indices), self.batch_size): @@ -88,6 +166,36 @@ def __len__(self): return len(self.sp_group_local_indices) // self.batch_size +def _shape_bucket_id_from_parquet_path(file_path: str) -> str: + """Return and validate the sole ``bucket=...`` ancestor of a parquet.""" + matches = [part for part in Path(file_path).parts if part.startswith("bucket=")] + if len(matches) != 1: + raise ValueError( + "Native-shape parquet paths must have exactly one ancestor named " + "'bucket=x-f', got " + f"{file_path!r} with bucket ancestors {matches}" + ) + bucket_id = matches[0] + parse_video_shape_bucket_id(bucket_id) + return bucket_id + + +def shape_bucket_ids_from_parquet_files( + parquet_files: Sequence[str], + lengths: Sequence[int], +) -> list[str]: + """Expand canonical path bucket IDs to one identifier per dataset row.""" + if len(parquet_files) != len(lengths): + raise ValueError("parquet_files and lengths must have matching lengths") + sample_bucket_ids: list[str] = [] + for file_path, length in zip(parquet_files, lengths, strict=True): + if int(length) < 0: + raise ValueError(f"Parquet row counts must be non-negative, got {length}") + bucket_id = _shape_bucket_id_from_parquet_path(str(file_path)) + sample_bucket_ids.extend([bucket_id] * int(length)) + return sample_bucket_ids + + def _parse_data_path_specs(path: str | Sequence[str] | dict[str, int]) -> list[tuple[str, int]]: """Parse one or more dataset roots with old-framework repeat counts.""" if isinstance(path, dict): @@ -226,7 +334,12 @@ def get_parquet_files_and_length(path: str | Sequence[str] | dict[str, int]): return file_names_sorted, lengths_sorted -def read_row_from_parquet_file(parquet_files: list[str], global_row_idx: int, lengths: list[int]) -> dict[str, Any]: +def read_row_from_parquet_file( + parquet_files: list[str], + global_row_idx: int, + lengths: list[int], + columns: Sequence[str] | None = None, +) -> dict[str, Any]: ''' Read a row from a parquet file. Args: @@ -268,7 +381,10 @@ def read_row_from_parquet_file(parquet_files: list[str], global_row_idx: int, le # If we reach here, local_row_idx is out of bounds for this parquet file raise IndexError(f"local_row_idx {local_row_idx} is out of bounds for parquet file {parquet_files[file_index]}") - row_group = parquet_file.read_row_group(row_group_index).to_pydict() + # Project at the Parquet reader boundary. This is especially important for + # data-free training over a T2VA superset: the text-only schema must not + # pull hundreds of MiB of unused video/audio latent bytes into host memory. + row_group = parquet_file.read_row_group(row_group_index, columns=columns).to_pydict() row_dict = {k: v[local_index] for k, v in row_group.items()} del row_group @@ -295,6 +411,7 @@ def __init__( drop_last: bool = True, drop_first_row: bool = False, text_padding_length: int = 512, + native_shape_bucketing: bool = False, ): super().__init__() self.path = path @@ -307,6 +424,9 @@ def __init__( self.parquet_files, self.lengths = get_parquet_files_and_length(path) self.batch = batch_size self.text_padding_length = text_padding_length + self.sample_bucket_ids = ( + shape_bucket_ids_from_parquet_files(self.parquet_files, self.lengths) if native_shape_bucketing else None + ) self.sampler = DP_SP_BatchSampler( batch_size=batch_size, dataset_size=sum(self.lengths), @@ -316,6 +436,7 @@ def __init__( drop_last=drop_last, drop_first_row=drop_first_row, seed=seed, + sample_bucket_ids=self.sample_bucket_ids, ) logger.info("Dataset initialized with %d parquet files and %d rows", len(self.parquet_files), sum(self.lengths)) @@ -330,7 +451,12 @@ def get_validation_negative_prompt(self) -> tuple[torch.Tensor, torch.Tensor, st file_path = self.parquet_files[0] row_idx = 0 # Read the negative prompt data - row_dict = read_row_from_parquet_file([file_path], row_idx, [self.lengths[0]]) + row_dict = read_row_from_parquet_file( + [file_path], + row_idx, + [self.lengths[0]], + columns=self.parquet_schema.names, + ) batch = collate_rows_from_parquet_schema([row_dict], self.parquet_schema, @@ -352,7 +478,14 @@ def __getitems__(self, indices: list[int]) -> dict[str, Any]: """ Batch fetch using read_row_from_parquet_file for each index. """ - rows = [read_row_from_parquet_file(self.parquet_files, idx, self.lengths) for idx in indices] + rows = [ + read_row_from_parquet_file( + self.parquet_files, + idx, + self.lengths, + columns=self.parquet_schema.names, + ) for idx in indices + ] # Inject sample indices for deterministic CFG dropout # that is reproducible across checkpoint resume. @@ -364,6 +497,11 @@ def __getitems__(self, indices: list[int]) -> dict[str, Any]: self.text_padding_length, cfg_rate=self.cfg_rate, seed=self.seed) + if self.sample_bucket_ids is not None: + bucket_ids = {self.sample_bucket_ids[index] for index in indices} + if len(bucket_ids) != 1: + raise RuntimeError(f"Exact-shape sampler emitted a mixed bucket batch: {sorted(bucket_ids)}") + batch["_shape_bucket_id"] = bucket_ids.pop() return batch def __len__(self): @@ -385,7 +523,9 @@ def build_parquet_map_style_dataloader(path, drop_last=True, drop_first_row=False, text_padding_length=512, - seed=42) -> tuple[LatentsParquetMapStyleDataset, StatefulDataLoader]: + seed=42, + native_shape_bucketing=False) -> tuple[LatentsParquetMapStyleDataset, + StatefulDataLoader]: dataset = LatentsParquetMapStyleDataset(path, batch_size, cfg_rate=cfg_rate, @@ -393,7 +533,8 @@ def build_parquet_map_style_dataloader(path, drop_first_row=drop_first_row, text_padding_length=text_padding_length, parquet_schema=parquet_schema, - seed=seed) + seed=seed, + native_shape_bucketing=native_shape_bucketing) loader = StatefulDataLoader( dataset, diff --git a/fastvideo/dataset/shape_bucket.py b/fastvideo/dataset/shape_bucket.py new file mode 100644 index 0000000000..35dea8d239 --- /dev/null +++ b/fastvideo/dataset/shape_bucket.py @@ -0,0 +1,49 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Portable exact-shape bucket identifiers for preprocessed video data.""" + +from __future__ import annotations + +import re +from dataclasses import dataclass + +_VIDEO_SHAPE_BUCKET_PATTERN = re.compile( + r"^bucket=(?P[1-9][0-9]*)x(?P[1-9][0-9]*)-(?P[1-9][0-9]*)f$" +) + + +@dataclass(frozen=True, slots=True) +class ExactVideoShapeBucket: + """Pixel geometry and frame count encoded by one bucket directory.""" + + width: int + height: int + num_frames: int + + @property + def bucket_id(self) -> str: + return f"bucket={self.width}x{self.height}-{self.num_frames}f" + + +def parse_video_shape_bucket_id(bucket_id: str) -> ExactVideoShapeBucket: + """Parse ``bucket=x-f`` exactly. + + Width is deliberately first so the identifier matches common media + geometry notation. The strict spelling makes independently generated data + roots portable and prevents ranks from silently assigning the same shape + two different names. + """ + match = _VIDEO_SHAPE_BUCKET_PATTERN.fullmatch(str(bucket_id)) + if match is None: + raise ValueError( + "An exact video-shape bucket must be named " + "'bucket=x-f' with positive decimal " + f"integers (for example 'bucket=1344x768-124f'), got {bucket_id!r}" + ) + return ExactVideoShapeBucket( + width=int(match.group("width")), + height=int(match.group("height")), + num_frames=int(match.group("num_frames")), + ) + + +__all__ = ["ExactVideoShapeBucket", "parse_video_shape_bucket_id"] diff --git a/fastvideo/dataset/validation_dataset.py b/fastvideo/dataset/validation_dataset.py index 755df34f65..d9026ac732 100644 --- a/fastvideo/dataset/validation_dataset.py +++ b/fastvideo/dataset/validation_dataset.py @@ -1,5 +1,6 @@ # SPDX-License-Identifier: Apache-2.0 # adapted from: https://github.com/a-r-r-o-w/finetrainers/blob/main/finetrainers/data/dataset.py +import json import os import pathlib @@ -29,7 +30,15 @@ def __init__(self, filename: str): if self.filename.suffix == ".csv": data = datasets.load_dataset("csv", data_files=self.filename.as_posix(), split="train") elif self.filename.suffix == ".json": - data = datasets.load_dataset("json", data_files=self.filename.as_posix(), split="train", field="data") + # Historically validation JSON was wrapped as {"data": [...]}; + # native-shape held-out manifests are plain JSON arrays. Parse the + # tiny validation document directly so both contracts remain + # supported without guessing a datasets ``field`` value. + document = json.loads(self.filename.read_text(encoding="utf-8")) + rows = document.get("data") if isinstance(document, dict) else document + if not isinstance(rows, list): + raise ValueError("Validation JSON must be a row array or an object containing a 'data' row array") + data = datasets.Dataset.from_list(rows) elif self.filename.suffix == ".parquet": data = datasets.load_dataset("parquet", data_files=self.filename.as_posix(), split="train") elif self.filename.suffix == ".arrow": diff --git a/fastvideo/fastvideo_args.py b/fastvideo/fastvideo_args.py index a648a02538..aabb148b76 100644 --- a/fastvideo/fastvideo_args.py +++ b/fastvideo/fastvideo_args.py @@ -203,6 +203,10 @@ class FastVideoArgs: # Per-component flags below let callers compile additional submodules # independently; ``False`` leaves the component eager. enable_torch_compile: bool = False + # Regional fullgraph compile of repeated blocks (modular fastvideo/train + # stack). False preserves the legacy whole-model torch.compile semantics + # for fastvideo/training recipes; the modular moduleloader sets it True. + regional_compile: bool = False enable_torch_compile_text_encoder: bool = False enable_torch_compile_vae: bool = False enable_torch_compile_audio_vae: bool = False @@ -688,7 +692,9 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: type=str, default=None, help= - "JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'", + "JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'. " + "Note: the modular fastvideo/train stack uses regional fullgraph compile, which rejects 'mode' " + "(it injects inductor options); express mode effects via 'options' there.", ) parser.add_argument( "--inference-torch-compile", diff --git a/fastvideo/layers/quantization/__init__.py b/fastvideo/layers/quantization/__init__.py index 83f5076f1c..b484886eb8 100644 --- a/fastvideo/layers/quantization/__init__.py +++ b/fastvideo/layers/quantization/__init__.py @@ -11,6 +11,8 @@ "nvfp4_qat", "nvfp4_qat_train", "fp8_qat_train", + "INT8Affine", + "W4A16", ] QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods)) @@ -61,20 +63,22 @@ def get_quantization_config(quantization: str) -> type[QuantizationConfig]: # lazy import to avoid triggering `torch.compile` too early from .absmax_fp8 import AbsMaxFP8Config from .fp8_config import FP8Config - from .mxfp8_config import MXFP8Config from .nvfp4_config import NVFP4Config from .nvfp4_qat_config import NVFP4QATConfig from .nvfp4_qat_train_config import NVFP4QATTrainConfig from .fp8_qat_train_config import FP8QATTrainConfig + from .int8_affine_config import INT8AffineConfig + from .w4a16_config import W4A16Config method_to_config: dict[str, type[QuantizationConfig]] = { "AbsMaxFP8": AbsMaxFP8Config, "FP8": FP8Config, - "MXFP8": MXFP8Config, "NVFP4": NVFP4Config, "nvfp4_qat": NVFP4QATConfig, "nvfp4_qat_train": NVFP4QATTrainConfig, "fp8_qat_train": FP8QATTrainConfig, + "INT8Affine": INT8AffineConfig, + "W4A16": W4A16Config, } # Update the `method_to_config` with customized quantization methods. method_to_config.update(_CUSTOMIZED_METHOD_TO_QUANT_CONFIG) diff --git a/fastvideo/layers/quantization/int8_affine_config.py b/fastvideo/layers/quantization/int8_affine_config.py new file mode 100644 index 0000000000..4a5d1bcec1 --- /dev/null +++ b/fastvideo/layers/quantization/int8_affine_config.py @@ -0,0 +1,988 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Weight-only affine INT8 quantization (group size 64) for CUDA inference. + +This is the CUDA-side counterpart of the affine INT8 scheme the Apple +Silicon (MLX) deployment path already validates: per-group min/max affine +quantization with ``group_size=64`` and ``bits=8``, applied to the linear +*weights* only. Activations stay in bf16 — there is no activation +quantizer here and there should not be one until a fused INT8 GEMM lands. + +The quantizer math is transcribed from +``fastvideo/layers/quantization/mlx_affine_qat.py`` (itself a transcription +of MLX's CPU ``quantized.cpp::quantize`` at v0.31.2), so the integer codes, +per-group scales, and per-group biases this config produces are the same +decisions the MLX runtime's quantizer makes. ``int8_affine_quantize`` and +``int8_affine_dequantize`` were verified bit-identical to +``mlx_affine_quantize_reference`` / ``mlx_affine_dequantize_reference`` on +fp32, fp16, and bf16 inputs (codes, scales, biases, and the dequantized +tensor all exactly equal) — the only representation change is that codes are +stored as ``uint8`` rather than ``int32``. + +Design notes, deliberately different from ``nvfp4_config.py``: + +- **No hardcoded single-model layer list.** ``NVFP4Config`` hardcodes the + LTX-2 prefix set and its own docstring flags that as a wart. Here the + selection rule is a constructor field (``target_layers`` / + ``layer_suffixes``) with a model-agnostic default, and the MiniMax-H3 set + is built from H3's real module names by + :func:`minimax_h3_int8_affine_prefixes` / ``INT8AffineConfig.for_minimax_h3``. +- **A fail-closed deny list.** ``attn.to_gate_compress`` is H3's + sparse-attention (VSA) gate. Quantizing it perturbs a *discrete* routing + decision, so an error there is not a small output perturbation — it flips + which tiles the sparse attention attends to. H3's own deploy path leaves + it alone. The name matches none of the usual "norm"/"scale_shift_table" + exclusion heuristics, so it is excluded by an explicit deny list that the + constructor cannot widen away. + +Like ``nvfp4_config.py``, this module allocates a *dense bf16* weight and +converts to the low-precision form at load time (``convert_model_to_int8_affine``), +so a plain BF16 checkpoint quantizes with no pre-quantized weights. +""" + +from __future__ import annotations + +import json +import logging +import os +from collections.abc import Iterable +from typing import Any + +import torch +import torch.nn.functional as F +from torch.nn.parameter import Parameter + +from fastvideo.layers.quantization.base_config import ( + QuantizationConfig, + QuantizeMethodBase, +) +from fastvideo.models.utils import set_weight_attrs + +logger = logging.getLogger(__name__) + +DEFAULT_GROUP_SIZE = 64 +DEFAULT_BITS = 8 +_EPS = 1e-7 +# Affine codes span [0, 2**bits - 1]; for bits=8 that is [0, 255], which does +# NOT fit torch.int8. Codes are therefore stored as torch.uint8 (see +# ``int8_affine_quantize``) — the *scheme* is int8 affine, the container is +# unsigned because the zero-point convention is free-floating. +_MAX_UINT8_CODE = 255 + + +# --------------------------------------------------------------------------- +# Affine quantizer — MLX-parity math +# --------------------------------------------------------------------------- + + +def _group(w: torch.Tensor, group_size: int) -> torch.Tensor: + """View the last dim as ``(num_groups, group_size)``. + + Identical to ``mlx_affine_qat._group``: MLX groups along the last + (input) axis, which for a linear weight is the contraction dimension. + """ + if w.shape[-1] % group_size != 0: + raise ValueError(f"Last dim {w.shape[-1]} is not divisible by group_size {group_size}; " + "MLX affine quantization groups along the last axis.") + return w.reshape(*w.shape[:-1], w.shape[-1] // group_size, group_size) + + +def int8_affine_quantize( + w: torch.Tensor, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Quantize like ``mx.quantize(..., mode="affine")``. + + Transcribed from ``mlx_affine_qat.mlx_affine_quantize_reference``. The + decisions reproduced exactly are: per-group min/max in the input's + arithmetic, a *negative* scale when ``|w_max| >= |w_min|`` (the + quantizer anchors at the endpoint with the larger magnitude), the anchor + re-expressed as an exact integer multiple of the scale so the extreme + value round-trips exactly, ``rint`` (round-half-to-even) rounding, and + codes clamped to ``[0, 2**bits - 1]``. + + Bit-identical to ``mlx_affine_quantize_reference`` for a given input + dtype; the only representation change is that codes are ``torch.uint8`` + rather than ``torch.int32``, because ``2**bits - 1 == 255`` does not fit + a signed byte. Callers must not cast these to ``int8`` — 255 would wrap + to -1 and the dequantized weight would be wrong. + + ``scales``/``biases`` come back in ``w.dtype``, as in the reference. + Since the codes are decided *before* that cast, calling this on an fp32 + ``w`` costs nothing in code fidelity and avoids rounding the stored + scales to bf16 — which is what ``_quantize_layer_weight`` does (a bf16 + checkpoint value converts to fp32 exactly, so this is lossless input + with a higher-precision scale store). + + Returns ``(codes, scales, biases)``. Mirroring the reference, ``codes`` + comes back in the *grouped* shape ``w.shape[:-1] + (K // group_size, group_size)`` + (not ``w.shape``) and ``scales``/``biases`` in + ``w.shape[:-1] + (K // group_size,)``; pass ``out_shape=w.shape`` to + ``int8_affine_dequantize`` to flatten the grouping back out. + """ + n_bins = float((1 << bits) - 1) + grouped = _group(w, group_size).float() + + w_min = grouped.min(dim=-1).values + w_max = grouped.max(dim=-1).values + mask = w_min.abs() > w_max.abs() + scale = ((w_max - w_min) / n_bins).clamp_min(_EPS) + scale = torch.where(mask, scale, -scale) + edge = torch.where(mask, w_min, w_max) + q0 = torch.round(edge / scale) + nonzero_q0 = q0 != 0 + scale = torch.where(nonzero_q0, edge / torch.where(nonzero_q0, q0, torch.ones_like(q0)), scale) + bias = torch.where(nonzero_q0, edge, torch.zeros_like(edge)) + + codes = torch.round((grouped - bias.unsqueeze(-1)) / scale.unsqueeze(-1)) + codes = codes.clamp(min=0.0, max=n_bins) + if codes.max().item() > _MAX_UINT8_CODE: + raise ValueError(f"bits={bits} produced a code above {_MAX_UINT8_CODE}; only bits<=8 fits uint8 storage.") + return codes.to(torch.uint8), scale.to(w.dtype), bias.to(w.dtype) + + +def int8_affine_dequantize( + codes: torch.Tensor, + scales: torch.Tensor, + biases: torch.Tensor, + *, + out_shape: torch.Size | None = None, +) -> torch.Tensor: + """``code * scale + bias`` per group — the inverse of ``int8_affine_quantize``. + + Mirrors ``mlx_affine_qat.mlx_affine_dequantize_reference`` (which itself + matches MLX's *CPU* kernel). ``codes`` may be ``uint8``; the multiply is + done in the scales' dtype, so pass fp32 scales to get the fp32 stream. + + ``codes`` is accepted in either shape the quantizer's callers use: the + grouped ``(*, K // group_size, group_size)`` the reference returns, or the + flattened ``(*, K)`` a stored weight buffer naturally has. The two are + distinguished by rank (grouped codes are one rank above ``scales``). + """ + dtype = scales.dtype + if codes.dim() == scales.dim(): + codes = codes.reshape(*scales.shape, codes.shape[-1] // scales.shape[-1]) + deq = codes.to(dtype) * scales.unsqueeze(-1) + biases.unsqueeze(-1) + if out_shape is not None: + deq = deq.reshape(out_shape) + return deq + + +# --------------------------------------------------------------------------- +# Layer selection +# --------------------------------------------------------------------------- + +# Model-agnostic default: the transformer-block GEMMs every DiT in this +# repo names this way (H3, LTX-2's `attn1/attn2`, ...). Suffix matching is +# used rather than a literal prefix set so depth/prefix variations cannot +# silently drop layers. +_GENERIC_LINEAR_SUFFIXES: tuple[str, ...] = ( + "attn.to_q", + "attn.to_k", + "attn.to_v", + "attn.to_out", + "ff.fc_in", + "ff.fc_out", +) + +# Names that are NEVER quantized by this config, whatever else is configured. +# Checked before the allowlist, and unioned with (never replaced by) any +# caller-supplied list, so this is fail-closed: there is no constructor +# argument that re-enables them. +_NEVER_QUANTIZE_SUBSTRINGS: tuple[str, ...] = ( + # H3's VSA sparse-attention compression gate. Quantizing it perturbs a + # discrete routing decision and the checkpoint's gate is zero-initialized + # (the branch is exactly disabled until finetuned). H3's deploy path + # ignores it, so we must too. + "to_gate_compress", + # Global timestep-basis projector; feeds every block's modulation. + "adaln_basis", +) + +# Modules H3 already pins to fp32 (MiniMaxH3Transformer3DModel._keep_in_fp32_modules). +# They must not be targeted even if a suffix rule would otherwise reach them. +_H3_FP32_KEPT_SUBSTRINGS: tuple[str, ...] = ( + "proj_in", + "audio_proj_in", + "proj_out", + "audio_proj_out", + "time_embedder", +) + +# `context_embedder` is H3's text input projection. It is NOT in the model's +# fp32 keep set, but it is the same kind of module as `proj_in` / +# `audio_proj_in` (an input projection), and H3's deploy keeps input +# projections in fp32. Quantizing the text conditioning stream while leaving +# the video/audio input streams in fp32 is an unvalidated asymmetry, so it is +# excluded by default. ``include_context_embedder=True`` opts in. +_H3_INPUT_PROJECTION_SUBSTRINGS: tuple[str, ...] = ("context_embedder", ) + +MINIMAX_H3_PREFIX = "minimax_h3" +MINIMAX_H3_NUM_LAYERS = 50 +MINIMAX_H3_NUM_REFINER_LAYERS = 2 +# Both block stacks hold the same `MiniMaxH3Attention` / `MiniMaxH3FeedForward` +# modules; the refiner stack simply has no `adaln_proj`. +MINIMAX_H3_BLOCK_SCOPES: tuple[str, ...] = ( + "transformer_blocks", + "token_refiner.refiner_blocks", +) +# The H3 profile equals the generic set plus the per-block AdaLN modulation +# projection (`minimax_h3.transformer_blocks.{i}.adaln_proj.linear`), which is +# a real per-block GEMM in H3 and is listed in H3's linear inventory. +MINIMAX_H3_INT8_AFFINE_SUFFIXES: tuple[str, ...] = _GENERIC_LINEAR_SUFFIXES + ("adaln_proj.linear", ) +MINIMAX_H3_INT8_AFFINE_EXCLUSIONS: tuple[str, ...] = ( + _NEVER_QUANTIZE_SUBSTRINGS + _H3_FP32_KEPT_SUBSTRINGS + _H3_INPUT_PROJECTION_SUBSTRINGS) + + +def minimax_h3_int8_affine_prefixes( + *, + prefix: str = MINIMAX_H3_PREFIX, + num_layers: int = MINIMAX_H3_NUM_LAYERS, + num_refiner_layers: int = MINIMAX_H3_NUM_REFINER_LAYERS, + suffixes: Iterable[str] = MINIMAX_H3_INT8_AFFINE_SUFFIXES, +) -> frozenset[str]: + """Enumerate the exact H3 linear prefixes this config targets. + + Built from H3's real module names as constructed in + ``fastvideo/models/dits/minimax_h3.py``: + ``MiniMaxH3TransformerBlock`` builds ``{prefix}.transformer_blocks.{i}.attn``, + ``.ff``, ``.adaln_proj``; ``MiniMaxH3TokenRefiner`` builds + ``{prefix}.token_refiner.refiner_blocks.{i}.attn`` / ``.ff``. Defaults + match ``MiniMaxH3ArchConfig`` (``prefix="minimax_h3"``, ``num_layers=50``, + ``num_refiner_layers=2``). + + The enumerated set is *not* how selection runs at runtime (suffix + + deny-list matching is, so depth changes cannot silently drop layers); + it exists to be asserted against in tests and to give callers who want + a literal set one place to get it. + """ + suffixes = tuple(suffixes) + prefixes: set[str] = set() + for index in range(num_layers): + for suffix in suffixes: + prefixes.add(f"{prefix}.transformer_blocks.{index}.{suffix}") + for index in range(num_refiner_layers): + for suffix in suffixes: + # The refiner stack has no `adaln_proj`. + if suffix.startswith("adaln_proj"): + continue + prefixes.add(f"{prefix}.token_refiner.refiner_blocks.{index}.{suffix}") + return frozenset(prefixes) + + +class INT8AffineConfig(QuantizationConfig): + """Weight-only affine INT8 (group-64) quantization for CUDA DiT inference. + + Layer selection is a constructor field, not a hardcoded model list: + ``target_layers`` (explicit full prefixes) takes precedence when given, + otherwise ``layer_suffixes`` is matched with ``str.endswith``. Both are + subject to a fail-closed deny list — see ``exclude_substrings``. + + Weight-only: there is no activation quantizer, and ``INT8AffineQuantizeMethod.apply`` + dequantizes the stored codes back to the activation dtype and runs a + normal bf16/fp32 GEMM. That is the correctness-first path; a fused INT8 + GEMM is a follow-up, not a prerequisite. + + Only the INT8 arithmetic is scheme-specific — nothing here is H3-only. + Use :meth:`for_minimax_h3` for the H3 profile. + """ + + def __init__( + self, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + target_layers: Iterable[str] | None = None, + layer_suffixes: Iterable[str] | None = None, + exclude_substrings: Iterable[str] | None = None, + include_context_embedder: bool = False, + retain_original_weight: bool = True, + ) -> None: + super().__init__() + if bits < 2 or bits > 8: + raise ValueError(f"bits must be in [2, 8] (codes are stored as uint8), got {bits}") + if group_size <= 0: + raise ValueError(f"group_size must be positive, got {group_size}") + self.group_size = group_size + self.bits = bits + self.target_layers: frozenset[str] | None = (None if target_layers is None else frozenset(target_layers)) + self.layer_suffixes: tuple[str, ...] = (tuple(_GENERIC_LINEAR_SUFFIXES) + if layer_suffixes is None else tuple(layer_suffixes)) + # Deny list is always unioned with the hard exclusions: passing a + # custom list can only ever *add* exclusions, never remove one. This + # is what keeps `to_gate_compress` unquantizable. + self.exclude_substrings: tuple[str, ...] = tuple(_NEVER_QUANTIZE_SUBSTRINGS) + tuple( + exclude_substrings or ()) + self._include_context_embedder = include_context_embedder + if not include_context_embedder: + self.exclude_substrings = self.exclude_substrings + _H3_INPUT_PROJECTION_SUBSTRINGS + # Keep the dense bf16 `layer.weight` Parameter after conversion. + # Default True: several H3 forwards read `linear.weight.dtype` to + # cast their input (e.g. `MiniMaxH3AdaLayerNormModulation.forward`), + # so purging it breaks the model. Setting False frees the bf16 copy + # at the cost of requiring every caller to stop touching `.weight`. + self.retain_original_weight = retain_original_weight + + def get_name(self) -> str: + return "INT8Affine" + + def get_supported_act_dtypes(self) -> list[torch.dtype]: + return [torch.bfloat16, torch.float16, torch.float32] + + @classmethod + def get_min_capability(cls) -> int: + """Turing (75). + + The compute path is a plain bf16/fp32 GEMM over a dequantized weight, + so no INT8 tensor-core class is required; 75 matches ``AbsMaxFP8Config`` + and keeps the config loadable on the same hosts. + """ + return 75 + + @staticmethod + def get_config_filenames() -> list[str]: + return [] + + @classmethod + def from_config(cls, config: dict[str, Any]) -> INT8AffineConfig: + return cls( + group_size=config.get("group_size", DEFAULT_GROUP_SIZE), + bits=config.get("bits", DEFAULT_BITS), + target_layers=config.get("target_layers"), + layer_suffixes=config.get("layer_suffixes"), + exclude_substrings=config.get("exclude_substrings"), + include_context_embedder=config.get("include_context_embedder", False), + retain_original_weight=config.get("retain_original_weight", True), + ) + + @classmethod + def for_minimax_h3( + cls, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + include_context_embedder: bool = False, + retain_original_weight: bool = True, + ) -> INT8AffineConfig: + """The verified MiniMax-H3 profile: attention + FFN + AdaLN GEMMs. + + Excludes H3's fp32-pinned modules, the VSA gate, and (by default) + the text input projection. + """ + return cls( + group_size=group_size, + bits=bits, + layer_suffixes=MINIMAX_H3_INT8_AFFINE_SUFFIXES, + exclude_substrings=_H3_FP32_KEPT_SUBSTRINGS, + include_context_embedder=include_context_embedder, + retain_original_weight=retain_original_weight, + ) + + def is_target_layer(self, prefix: str) -> bool: + """Whether ``prefix`` is quantized under this config. + + Deny list first (fail-closed), then ``target_layers`` if supplied, + else suffix matching. Non-``LinearBase`` layers are filtered by + :meth:`get_quant_method`, not here, so this is safe to call on any + module name. + """ + for banned in self.exclude_substrings: + if banned in prefix: + return False + if self.target_layers is not None: + return prefix in self.target_layers + return any(prefix.endswith(suffix) for suffix in self.layer_suffixes) + + def get_quant_method(self, layer: torch.nn.Module, prefix: str): + from fastvideo.layers.linear import LinearBase + + if isinstance(layer, LinearBase) and self.is_target_layer(prefix): + return INT8AffineQuantizeMethod( + layer_prefix=prefix, + group_size=self.group_size, + bits=self.bits, + retain_original_weight=self.retain_original_weight, + ) + return None + + +class INT8AffineQuantizeMethod(QuantizeMethodBase): + """Linear method for weight-only affine INT8. + + ``create_weights`` allocates the same dense bf16 Parameter an + unquantized linear would (so the BF16 checkpoint loads unchanged), and + the INT8 codes/scales/biases arrive later as non-persistent buffers from + :func:`convert_model_to_int8_affine` — i.e. conversion happens at *load* + time, not construction time, exactly mirroring ``NVFP4QuantizeMethod``. + """ + + def __init__( + self, + layer_prefix: str = "", + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + retain_original_weight: bool = True, + ) -> None: + super().__init__() + self.layer_prefix = layer_prefix + self.group_size = group_size + self.bits = bits + self.retain_original_weight = retain_original_weight + + def create_weights( + self, + layer: torch.nn.Module, + input_size_per_partition: int, + output_partition_sizes: list[int], + input_size: int, + output_size: int, + params_dtype: torch.dtype, + **extra_weight_attrs, + ): + weight = Parameter( + torch.empty( + sum(output_partition_sizes), + input_size_per_partition, + dtype=params_dtype, + ), + requires_grad=False, + ) + set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0}) + layer.register_parameter("weight", weight) + set_weight_attrs(weight, extra_weight_attrs) + + def _ensure_quantized(self, layer: torch.nn.Module) -> bool: + """Convert on first use if the loader hook never ran. + + Returns False when the layer is intentionally left dense (grad-enabled + forward: a training step must see the master weight, not a frozen + dequantized copy). The loader path is ``_maybe_quantize_model`` -> + :func:`convert_model_to_int8_affine`; this fallback exists so the + config is still *correct* if that dispatch is missing, but it warns + because reaching it means the loader hook did not fire. + """ + if getattr(layer, "_int8_affine_codes", None) is not None: + return True + weight = getattr(layer, "weight", None) + if weight is None: + raise RuntimeError(f"INT8Affine layer {self.layer_prefix!r} has no weight and no quantized buffers.") + if torch.is_grad_enabled(): + return False + logger.warning( + "INT8Affine: layer %r reached apply() unquantized; converting lazily. The loader hook " + "(_maybe_quantize_model) did not dispatch to convert_model_to_int8_affine — check its " + "isinstance chain in fastvideo/models/loader/fsdp_load.py.", + self.layer_prefix, + ) + _quantize_layer_weight(layer, weight, group_size=self.group_size, bits=self.bits) + return True + + def apply( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + if not self._ensure_quantized(layer): + weight = layer.weight + return F.linear(x, weight.to(x.dtype) if weight.dtype != x.dtype else weight, bias) + + codes = layer._int8_affine_codes + # Dequantize in fp32 (the scales' dtype), then match the activation. + # `code * scale + bias` in fp32 is the more accurate side of the + # CPU/Metal split documented in mlx_affine_qat.py. + weight = int8_affine_dequantize( + codes, + layer._int8_affine_scales, + layer._int8_affine_biases, + out_shape=codes.shape, + ).to(x.dtype) + return F.linear(x, weight, bias) + + +# --------------------------------------------------------------------------- +# Load-time conversion +# --------------------------------------------------------------------------- + + +def _quantize_layer_weight( + mod: torch.nn.Module, + weight: torch.Tensor, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, +) -> None: + """Quantize one linear's weight in place into non-persistent buffers.""" + from torch.distributed.tensor import DTensor # type: ignore + + weight_local = weight.to_local() if isinstance(weight, DTensor) else weight # type: ignore[arg-type] + # fp32 source keeps the scale/bias solve out of bf16 (see int8_affine_quantize). + # nan_to_num matches convert_model_to_nvfp4: one NaN would otherwise poison + # every group it touches. + w32 = weight_local.detach().float().nan_to_num() + if w32.shape[-1] % group_size != 0: + raise ValueError(f"INT8Affine layer {mod!r}: input dim {w32.shape[-1]} is not divisible by " + f"group_size {group_size}.") + codes, scales, biases = int8_affine_quantize(w32, group_size=group_size, bits=bits) + # Store codes flattened back to the weight shape (the quantizer returns the + # grouped view) so `apply` can dequantize with out_shape=codes.shape. + mod.register_buffer("_int8_affine_codes", codes.reshape(w32.shape).contiguous(), persistent=False) + mod.register_buffer("_int8_affine_scales", scales.to(torch.float32).contiguous(), persistent=False) + mod.register_buffer("_int8_affine_biases", biases.to(torch.float32).contiguous(), persistent=False) + + +def convert_model_to_int8_affine(model: torch.nn.Module, ) -> None: + """Quantize every INT8Affine-tagged linear in-place after weights load. + + Mirrors ``convert_model_to_nvfp4`` / ``convert_model_to_fp8``: walk the + module tree once, convert each layer whose ``quant_method`` is an + :class:`INT8AffineQuantizeMethod`, and register the int8 codes plus + per-group scales/biases as non-persistent buffers (so they are not + written back into ``state_dict``/checkpoints). + + Callers: the loader hook ``_maybe_quantize_model`` in + ``fastvideo/models/loader/fsdp_load.py``. *That hook is not edited by + this module* — it dispatches on an explicit ``isinstance`` chain, so it + needs a matching branch (see the module report). Without it, + ``INT8AffineQuantizeMethod.apply`` converts lazily on first forward and + logs a warning, so inference is still correct, just later and noisier. + """ + converted = 0 + purged = 0 + shapes: set[tuple[int, int]] = set() + for mod in model.modules(): + qm = getattr(mod, "quant_method", None) + if not isinstance(qm, INT8AffineQuantizeMethod): + continue + weight = getattr(mod, "weight", None) + if weight is None: + continue + _quantize_layer_weight(mod, weight, group_size=qm.group_size, bits=qm.bits) + converted += 1 + shapes.add((qm.group_size, qm.bits)) + if not qm.retain_original_weight: + # register_parameter(None) (as convert_model_to_nvfp4 does) rather + # than popping the key: `layer.weight` then reads as None instead + # of raising AttributeError. + original = mod._parameters.get("weight") + if original is not None: + original.grad = None + mod.register_parameter("weight", None) + purged += 1 + + if converted: + logger.info("INT8Affine conversion receipt: quantized %d linear layers (%s); purged %d original bf16 " + "weight tensors.", converted, + ", ".join(f"group_size={g}, bits={b}" for g, b in sorted(shapes)), purged) + + +# --------------------------------------------------------------------------- +# Compact checkpoint sidecar — save/load of the quantized payload +# --------------------------------------------------------------------------- +# +# The INT8 buffers are registered with ``persistent=False`` (see +# ``_quantize_layer_weight``), so ``state_dict()`` does NOT carry them: a saved +# checkpoint holds dense bf16 weights only and the INT8 payload is rebuilt by +# re-running the conversion at every load. That is impossible on a host that +# cannot hold the dense weights at all — an RTX 5090 has 32 GB and H3's bf16 +# DiT is ~40 GB, so ``convert_model_to_int8_affine`` (which starts from the +# dense weight) can never run there. The sidecar is what makes a pre-quantized +# checkpoint servable on that hardware. +# +# Format (identical in shape to the NVFP4 sidecar, ``nvfp4_config.py``): one +# safetensors file keyed ``"::"`` plus a JSON manifest +# under the ``fastvideo_int8_affine`` metadata key describing the scheme. +# +# Buffer inventory restored by a load — these are the same three tensors +# ``_quantize_layer_weight`` registers, byte for byte: +# +# ``_int8_affine_codes`` uint8 ``(out_dim, in_dim)`` — one code per +# weight, flattened back out of the quantizer's +# grouped view so ``apply`` can dequantize with +# ``out_shape=codes.shape``. **uint8, not int8**: +# a bits=8 code spans [0, 255] and 255 does not fit +# a signed byte (it would wrap to -1 and silently +# invert that weight). The dtype is asserted at load +# rather than cast for exactly that reason. +# ``_int8_affine_scales`` float32 ``(out_dim, in_dim // group_size)`` +# ``_int8_affine_biases`` float32 ``(out_dim, in_dim // group_size)`` +# +# Both constants are stored fp32 by the converter regardless of the weight +# dtype, so fp32 is what a sidecar must carry; a bf16 copy would not be a +# bit-exact restore. + +INT8_AFFINE_SIDECAR_SUFFIX = ".int8affine.safetensors" +# Filename used when the sidecar sits inside a checkpoint *directory*; it does +# not carry the suffix above, which is what a sibling file is named with. +INT8_AFFINE_DIR_SIDECAR_NAME = "int8_affine.safetensors" +_INT8_AFFINE_SIDECAR_FORMAT = "fastvideo.int8_affine" +_INT8_AFFINE_SIDECAR_VERSION = 1 +_INT8_AFFINE_SIDECAR_METADATA_KEY = "fastvideo_int8_affine" +_INT8_AFFINE_SIDECAR_KEY_SEP = "::" +# Order matters only for the manifest; every buffer is optional on load so a +# future format can add tensors without breaking older readers. +_INT8_AFFINE_SIDECAR_BUFFERS = ( + "_int8_affine_codes", + "_int8_affine_scales", + "_int8_affine_biases", +) +# The exact container each buffer must have. A sidecar whose codes came back +# as int8 (or float32) would dequantize to garbage with no error anywhere, so +# these are validated, never coerced. +_INT8_AFFINE_SIDECAR_DTYPES = { + "_int8_affine_codes": torch.uint8, + "_int8_affine_scales": torch.float32, + "_int8_affine_biases": torch.float32, +} + + +def _sidecar_key(module_fqn: str, buffer_name: str) -> str: + return f"{module_fqn}{_INT8_AFFINE_SIDECAR_KEY_SEP}{buffer_name}" + + +def _int8_affine_tagged_modules(model: torch.nn.Module) -> list[tuple[str, torch.nn.Module, INT8AffineQuantizeMethod]]: + tagged = [] + for fqn, mod in model.named_modules(): + qm = getattr(mod, "quant_method", None) + if isinstance(qm, INT8AffineQuantizeMethod): + tagged.append((fqn, mod, qm)) + return tagged + + +def _is_dtensor(tensor: torch.Tensor) -> bool: + # Imported lazily: torch.distributed.tensor is not cheap to import and is + # absent on some builds. + try: + from torch.distributed.tensor import DTensor # type: ignore + except ImportError: # pragma: no cover - depends on the torch build + return False + return isinstance(tensor, DTensor) + + +def int8_affine_sidecar_state_dict(model: torch.nn.Module) -> dict[str, torch.Tensor]: + """Collect the quantized tensors of every INT8 affine linear in *model*. + + Keys are ``"::"`` and values are detached CPU + copies. Modules whose buffers are missing (never converted) are skipped; + the returned mapping is what :func:`save_int8_affine_checkpoint` writes. + + FSDP note: a DTensor buffer is saved as this rank's local shard, so a + sharded save is only reloadable into an identically sharded model. + """ + state: dict[str, torch.Tensor] = {} + for fqn, mod, _ in _int8_affine_tagged_modules(model): + for name in _INT8_AFFINE_SIDECAR_BUFFERS: + tensor = getattr(mod, name, None) + if tensor is None: + continue + if _is_dtensor(tensor): + tensor = tensor.to_local() # type: ignore[attr-defined] + state[_sidecar_key(fqn, name)] = tensor.detach().to("cpu", copy=True).contiguous() + return state + + +def save_int8_affine_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + extra_metadata: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Write the model's INT8 affine tensors to a compact sidecar safetensors file. + + The file is roughly the INT8 size (1 byte/code plus two fp32 per group — + 8.5 bits/weight at ``group_size=64``) instead of the dense bf16 size. + Returns a receipt dict (also logged) with the module count and both sizes. + Raises ``RuntimeError`` when the model has no INT8 affine linears, which + usually means the config's layer selection did not cover the model's layer + paths, and when the tagged layers carry no buffers (never converted). + + The manifest under the ``fastvideo_int8_affine`` metadata key carries the + format name/version, the scheme (``group_size``/``bits``), the per-layer + weight and buffer shapes, and the quantized module fqns (the keys of + ``layers``), so a loader can validate a sidecar against a model without + materializing the tensors. + """ + from safetensors.torch import save_file + + state = int8_affine_sidecar_state_dict(model) + tagged = _int8_affine_tagged_modules(model) + if not tagged: + raise RuntimeError("No INT8 affine linear layers found in this model; nothing to serialize. " + "Check that the model was built with an INT8AffineConfig whose layer selection " + "covers its layer paths (e.g. INT8AffineConfig.for_minimax_h3()).") + if not state: + raise RuntimeError(f"Found {len(tagged)} INT8Affine-tagged linear layers but none carry quantized " + "buffers. Call convert_model_to_int8_affine(model) before saving a sidecar.") + + layers: dict[str, dict[str, Any]] = {} + quant_prefixes: dict[str, str] = {} + dense_bytes = 0 + group_sizes: set[int] = set() + bit_widths: set[int] = set() + for fqn, mod, qm in tagged: + codes = getattr(mod, "_int8_affine_codes", None) + weight = getattr(mod, "weight", None) + if codes is None and weight is None: + continue + # When the dense weight was purged the codes still give the logical + # shape: they hold one code per weight element. + weight_shape = [int(dim) for dim in (weight if weight is not None else codes).shape] + tensors = { + name: [int(dim) for dim in getattr(mod, name).shape] + for name in _INT8_AFFINE_SIDECAR_BUFFERS if getattr(mod, name, None) is not None + } + layers[fqn] = { + "weight_shape": weight_shape, + "group_size": int(qm.group_size), + "bits": int(qm.bits), + "tensors": tensors, + } + quant_prefixes[fqn] = getattr(qm, "layer_prefix", "") or "" + dense_bytes += weight_shape[0] * weight_shape[1] * 2 + group_sizes.add(int(qm.group_size)) + bit_widths.add(int(qm.bits)) + + metadata: dict[str, Any] = { + "format": _INT8_AFFINE_SIDECAR_FORMAT, + "version": _INT8_AFFINE_SIDECAR_VERSION, + # Uniform scheme when every layer agrees (the normal case); None when a + # model mixes schemes, in which case the per-layer entries are + # authoritative. A loader validates the per-layer values. + "group_size": group_sizes.pop() if len(group_sizes) == 1 else None, + "bits": bit_widths.pop() if len(bit_widths) == 1 else None, + "num_layers": len(layers), + "layers": layers, + "quant_prefixes": quant_prefixes, + "model_class": type(model).__name__, + } + if extra_metadata: + metadata.update(extra_metadata) + + payload = dict(state) + serialized_bytes = sum(t.numel() * t.element_size() for t in payload.values()) + save_file(payload, os.fspath(path), metadata={_INT8_AFFINE_SIDECAR_METADATA_KEY: json.dumps(metadata)}) + + receipt = { + "path": os.fspath(path), + "num_layers": len(layers), + "num_tensors": len(payload), + "quantized_bytes": serialized_bytes, + "dense_bfloat16_bytes": dense_bytes, + "compression_ratio": (dense_bytes / serialized_bytes) if serialized_bytes else 0.0, + } + logger.info( + "INT8Affine sidecar: wrote %d quantized modules / %d tensors (%d bytes) to %s " + "(%.2f GiB quantized vs %.2f GiB dense bf16, %.2fx smaller).", + receipt["num_layers"], + receipt["num_tensors"], + serialized_bytes, + receipt["path"], + serialized_bytes / (1 << 30), + dense_bytes / (1 << 30), + receipt["compression_ratio"], + ) + return receipt + + +def read_int8_affine_sidecar_metadata(path: str | os.PathLike[str]) -> dict[str, Any]: + """Return the manifest of a sidecar file without materializing its tensors.""" + from safetensors import safe_open + + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + raw = handle.metadata() or {} + if _INT8_AFFINE_SIDECAR_METADATA_KEY not in raw: + raise ValueError(f"{os.fspath(path)} is not a FastVideo INT8 affine sidecar " + f"(no {_INT8_AFFINE_SIDECAR_METADATA_KEY!r} metadata).") + return json.loads(raw[_INT8_AFFINE_SIDECAR_METADATA_KEY]) + + +def int8_affine_sidecar_path_for(checkpoint_path: str | os.PathLike[str]) -> str: + """Conventional sidecar path for a transformer checkpoint or directory. + + ``.../transformer.safetensors`` -> ``.../transformer.int8affine.safetensors``; + a directory -> ``/int8_affine.safetensors``. + """ + raw = os.fspath(checkpoint_path) + if os.path.isdir(raw): + return os.path.join(raw, INT8_AFFINE_DIR_SIDECAR_NAME) + if raw.endswith(".safetensors"): + return raw[:-len(".safetensors")] + INT8_AFFINE_SIDECAR_SUFFIX + return raw + INT8_AFFINE_SIDECAR_SUFFIX + + +def _sidecar_target_device(mod: torch.nn.Module, name: str) -> torch.device | None: + """Device the restored buffer should live on. + + Mirrors ``_quantize_layer_weight``, which registers the buffers on the + (local) weight's device; falls back to an existing buffer, then to the + module's parameter device so a purge-then-restore still lands on GPU. + """ + weight = getattr(mod, "weight", None) + if weight is not None and weight.device.type != "meta": + return weight.device + existing = getattr(mod, name, None) + if existing is not None and existing.device.type != "meta": + return existing.device + for param in mod.parameters(recurse=False): + if param.device.type != "meta": + return param.device + return None + + +def _expected_sidecar_shapes(name: str, weight_shape: tuple[int, int], group_size: int) -> tuple[tuple[int, ...], ...]: + """The one shape a fresh conversion would produce for *name*. + + Unlike the NVFP4 sidecar there is no padded variant to accept: the + quantizer groups along the last axis and ``_group`` refuses a K that is + not divisible by ``group_size``, so an exact divisor is the only + legitimate layout. A padded scales tensor would not merely be unusual — + ``int8_affine_dequantize`` recovers the group width as + ``codes.shape[-1] // scales.shape[-1]``, so extra groups silently regroup + every code in the row. + """ + out_dim, in_dim = weight_shape + if name == "_int8_affine_codes": + return ((out_dim, in_dim), ) + if in_dim % group_size: + raise ValueError(f"Sidecar declares a ({out_dim}, {in_dim}) weight with group_size {group_size}, " + "which does not divide the input dim; this layout cannot be dequantized.") + return ((out_dim, in_dim // group_size), ) + + +def load_int8_affine_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + strict: bool = True, +) -> int: + """Restore INT8 affine tensors from a sidecar, skipping ``convert_model_to_int8_affine``. + + Registers ``_int8_affine_codes`` / ``_int8_affine_scales`` / + ``_int8_affine_biases`` on every INT8Affine-tagged linear from the + sidecar, byte-for-byte as the conversion would have produced them. The + dense bf16 weights are never touched (they may be absent entirely), and + nothing here needs a GPU or a fused INT8 kernel — the dequantize-then-GEMM + reference path in :meth:`INT8AffineQuantizeMethod.apply` is pure PyTorch, + so a pre-quantized checkpoint loads on any host. + + ``strict`` raises on any layer-set mismatch (a sidecar that does not + describe this model); with ``strict=False`` those are logged and skipped, + leaving those layers unconverted. Scheme (``group_size``/``bits``) and + per-tensor shape/dtype mismatches are **never** downgraded: a mis-read + code buffer produces garbage output with no error, so those always raise. + + Returns the number of layers restored. + """ + from safetensors import safe_open + + tagged = _int8_affine_tagged_modules(model) + if not tagged: + raise RuntimeError("No INT8 affine linear layers are attached to this model, so a sidecar cannot be " + "restored. This is the silent-dense failure mode: the model's INT8AffineConfig " + "layer selection does not cover its layer paths (for MiniMax-H3 use " + "INT8AffineConfig.for_minimax_h3()).") + + manifest = read_int8_affine_sidecar_metadata(path) + if manifest.get("format") != _INT8_AFFINE_SIDECAR_FORMAT: + raise ValueError(f"Unsupported INT8 affine sidecar format {manifest.get('format')!r} in " + f"{os.fspath(path)}.") + if int(manifest.get("version", -1)) != _INT8_AFFINE_SIDECAR_VERSION: + raise ValueError(f"Unsupported INT8 affine sidecar version {manifest.get('version')!r} in " + f"{os.fspath(path)} (this build reads version {_INT8_AFFINE_SIDECAR_VERSION}).") + + saved_layers: dict[str, dict[str, Any]] = manifest.get("layers", {}) + model_fqns = {fqn for fqn, _, _ in tagged} + missing = sorted(model_fqns - set(saved_layers)) + extra = sorted(set(saved_layers) - model_fqns) + if missing or extra: + message = (f"INT8 affine sidecar {os.fspath(path)} does not match this model: " + f"{len(missing)} layers missing from the sidecar, {len(extra)} layers not in the model. " + f"First missing={missing[:3]}, first extra={extra[:3]}.") + if strict: + raise ValueError(message) + logger.warning("%s Restoring the intersection only.", message) + + restored = 0 + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + available = set(handle.keys()) + for fqn, mod, qm in tagged: + if fqn not in saved_layers: + continue + entry = saved_layers[fqn] + try: + weight_shape = tuple(int(dim) for dim in entry["weight_shape"]) + group_size = int(entry["group_size"]) + bits = int(entry["bits"]) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError(f"INT8 affine sidecar {os.fspath(path)} entry for {fqn!r} is malformed: " + f"{entry!r} does not carry an integer weight_shape/group_size/bits.") from exc + if len(weight_shape) != 2: + raise ValueError(f"INT8 affine sidecar entry {fqn!r} declares weight shape {list(weight_shape)}; " + "a linear weight is 2-D.") + # A different scheme means different dequantize arithmetic over the + # same bytes: never recoverable, so never relaxed by strict=False. + if group_size != qm.group_size or bits != qm.bits: + raise ValueError(f"INT8 affine sidecar {os.fspath(path)} was written for {fqn!r} with " + f"group_size={group_size}, bits={bits}, but this model quantizes it with " + f"group_size={qm.group_size}, bits={qm.bits}.") + weight = getattr(mod, "weight", None) + if weight is not None and tuple(int(dim) for dim in weight.shape) != weight_shape: + raise ValueError(f"INT8 affine sidecar entry {fqn!r} describes a {list(weight_shape)} weight, but " + f"this model's layer has shape {list(weight.shape)}.") + tensors: dict[str, torch.Tensor] = {} + for name in _INT8_AFFINE_SIDECAR_BUFFERS: + key = _sidecar_key(fqn, name) + if key not in available: + continue + tensor = handle.get_tensor(key) + expected_dtype = _INT8_AFFINE_SIDECAR_DTYPES[name] + if tensor.dtype != expected_dtype: + raise ValueError(f"INT8 affine sidecar tensor {key} has dtype {tensor.dtype}, expected " + f"{expected_dtype}. Codes are uint8 because a bits=8 code spans [0, 255] " + "and does not fit int8; casting would silently corrupt them.") + expected = _expected_sidecar_shapes(name, weight_shape, group_size) + if tuple(tensor.shape) not in expected: + raise ValueError(f"INT8 affine sidecar tensor {key} has shape {tuple(tensor.shape)}, expected " + f"one of {list(expected)} for a {list(weight_shape)} linear with " + f"group_size={group_size}.") + device = _sidecar_target_device(mod, name) + if device is not None: + tensor = tensor.to(device=device, non_blocking=True) + tensors[name] = tensor + if set(tensors) != set(_INT8_AFFINE_SIDECAR_BUFFERS): + message = (f"INT8 affine sidecar entry for {fqn!r} is incomplete (has {sorted(tensors)}); " + f"all of {list(_INT8_AFFINE_SIDECAR_BUFFERS)} are required.") + if strict: + raise ValueError(message) + logger.warning("%s Skipping this layer.", message) + continue + for name, tensor in tensors.items(): + mod.register_buffer(name, tensor, persistent=False) + restored += 1 + + logger.info("INT8Affine sidecar: restored %d quantized modules from %s (dense weights untouched).", restored, + os.fspath(path)) + return restored + + +__all__ = [ + "DEFAULT_BITS", + "DEFAULT_GROUP_SIZE", + "INT8AffineConfig", + "INT8AffineQuantizeMethod", + "INT8_AFFINE_DIR_SIDECAR_NAME", + "INT8_AFFINE_SIDECAR_SUFFIX", + "MINIMAX_H3_BLOCK_SCOPES", + "MINIMAX_H3_INT8_AFFINE_EXCLUSIONS", + "MINIMAX_H3_INT8_AFFINE_SUFFIXES", + "MINIMAX_H3_PREFIX", + "convert_model_to_int8_affine", + "int8_affine_dequantize", + "int8_affine_quantize", + "int8_affine_sidecar_path_for", + "int8_affine_sidecar_state_dict", + "load_int8_affine_checkpoint", + "minimax_h3_int8_affine_prefixes", + "read_int8_affine_sidecar_metadata", + "save_int8_affine_checkpoint", +] diff --git a/fastvideo/layers/quantization/nvfp4_config.py b/fastvideo/layers/quantization/nvfp4_config.py index 284622d595..48fdf45b21 100644 --- a/fastvideo/layers/quantization/nvfp4_config.py +++ b/fastvideo/layers/quantization/nvfp4_config.py @@ -1,5 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -"""NVFP4 linear quantization for supported transformer layer sets. +"""NVFP4 quantization (FlashInfer-backed) for LTX-2 and MiniMax-H3. NVFP4 is NVIDIA's block-scaled FP4 format (e2m1 mantissa, fp32 alpha, ``layout_128x4`` scale layout, group size 16) — distinct from @@ -8,8 +8,27 @@ variants that may land later (e.g. AMD's MX-FP4 or vendor-neutral e3m0). -The registered config targets the curated LTX-2 deployment set and the -main MiniMax-H3 transformer-block FFN linears. +Upstreamed from ``FastVideo-internal`` so consumers that load LTX-2 +weights with NVFP4 quantization can drive the public package +end-to-end. + +The set of quantized linears is **per-model configuration**, not a +hardcoded constant: ``NVFP4Config(layer_prefixes=...)`` selects it, and +``layer_prefixes=None`` keeps the historical LTX-2 set the default so +nothing regresses. A model whose prefix set is not supplied therefore +attaches **no** quant methods and runs dense in silence — see +``is_nvfp4_linear_prefix`` and the module docs in +``docs/quantization/h3_nvfp4.md``. + +``NVFP4Config.for_minimax_h3()`` returns the MiniMax-H3 set (300 block +linears), with H3's VSA compression gate ``attn.to_gate_compress`` +excluded. + +Quantized weights can be written to / restored from a compact sidecar +safetensors file (packed FP4 codes + block scales + global scale, ~4x +smaller than the dense bf16 weights) instead of being re-derived from +dense weights at load time — see ``save_nvfp4_checkpoint`` / +``load_nvfp4_checkpoint``. `flashinfer` is imported lazily inside the call paths that need it. This keeps ``import fastvideo`` cheap on hosts where flashinfer is @@ -18,8 +37,9 @@ """ from __future__ import annotations +import json import logging -import re +import os from typing import Any import torch @@ -75,17 +95,64 @@ def _require_flashinfer() -> tuple[Any, Any, Any]: _LTX2_NVFP4_LINEAR_PREFIXES = frozenset(f"ltx2.blocks.{block_idx}.{suffix}" for block_idx in range(48) for suffix in _LTX2_NVFP4_BLOCK_LINEAR_SUFFIXES) | frozenset( ("ltx2.adaln_single.linear", )) -_MINIMAX_H3_NVFP4_FF_PREFIX = re.compile(r"(?:^|\.)transformer_blocks\.\d+\.ff\.(?:fc_in|fc_out)$") + +# --- MiniMax-H3 layer set ------------------------------------------------- +# +# H3's DiT (``fastvideo/models/dits/minimax_h3.py``) names its block linears +# ``{prefix}.transformer_blocks.{i}.{suffix}`` with the default +# ``prefix="minimax_h3"`` and ``num_layers=50`` +# (``fastvideo/configs/models/dits/minimax_h3.py``). The six suffixes below are +# the always-dense block linears — attention QKV/out and the SwiGLU FFN — which +# carry most of H3's parameters. The token refiner, the per-block AdaLN +# modulation (``adaln_proj.linear``), the patch/audio/context projections and +# ``norm_out.linear`` are deliberately left dense: they are small, run once per +# block or once per forward, and quantizing them buys little while adding +# activation-quantize overhead on every call. +MINIMAX_H3_DIT_PREFIX = "minimax_h3" +MINIMAX_H3_NUM_LAYERS = 50 +MINIMAX_H3_BLOCK_LINEAR_SUFFIXES = ( + "attn.to_q", + "attn.to_k", + "attn.to_v", + "attn.to_out", + "ff.fc_in", + "ff.fc_out", +) +# 50 blocks x 6 linears = 300 quantized linears. +MINIMAX_H3_NVFP4_LINEAR_PREFIXES = frozenset(f"{MINIMAX_H3_DIT_PREFIX}.transformer_blocks.{block_idx}.{suffix}" + for block_idx in range(MINIMAX_H3_NUM_LAYERS) + for suffix in MINIMAX_H3_BLOCK_LINEAR_SUFFIXES) + +# Linears that must NEVER be quantized, whatever a caller passes in +# ``layer_prefixes``. ``attn.to_gate_compress`` is H3's VSA sparse-attention +# compression gate: H3's own deployment path loads it dense, and the +# zero-initialized gate is probed structurally in the forward +# (``MiniMaxH3Attention._gate_active``) to skip a guaranteed-zero branch. +# Quantizing it would perturb the gate's numerics and defeat the exact-zero +# skip that keeps the VSA branch free while it is untrained. +_ALWAYS_EXCLUDED_LINEAR_SUFFIXES = ("attn.to_gate_compress", ) +MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES = _ALWAYS_EXCLUDED_LINEAR_SUFFIXES + + +def _matches_linear_suffix(prefix: str, suffixes: frozenset[str] | tuple[str, ...]) -> bool: + """True when *prefix* is one of *suffixes*, or ends at a dot boundary. + + Entries may be a full module path or a trailing suffix + (``"attn.to_gate_compress"``), so one pattern covers every block of a + stack. The dot boundary keeps ``"ff.fc_in"`` from matching a hypothetical + ``"cross_ff.fc_in"``. + """ + return any(prefix == suffix or prefix.endswith("." + suffix) for suffix in suffixes) def is_ltx2_nvfp4_linear_prefix(prefix: str) -> bool: - """Return whether *prefix* belongs to the LTX-2 NVFP4 deployment set.""" - return prefix in _LTX2_NVFP4_LINEAR_PREFIXES - + """Return whether *prefix* belongs to the LTX-2 NVFP4 deployment set. -def is_minimax_h3_nvfp4_linear_prefix(prefix: str) -> bool: - """Return whether *prefix* is a main MiniMax-H3 transformer-block FFN linear.""" - return _MINIMAX_H3_NVFP4_FF_PREFIX.search(prefix) is not None + Kept as a module-level predicate for callers that predate the + ``NVFP4Config.layer_prefixes`` field; per-instance configs should use + :meth:`NVFP4Config.is_nvfp4_linear_prefix` instead. + """ + return prefix in _LTX2_NVFP4_LINEAR_PREFIXES def _is_ltx2_refine_only_prefix(prefix: str) -> bool: @@ -427,18 +494,31 @@ def apply( class NVFP4Config(QuantizationConfig): - """Select NVFP4 for the supported LTX-2 and MiniMax-H3 linear sets. + """NVFP4 quantization configuration, parameterized by layer paths. NVFP4 is NVIDIA's block-scaled FP4 (e2m1 mantissa, fp32 alpha, - ``layout_128x4`` scale layout, group size 16). LTX-2 uses its curated - attention and FFN deployment set. MiniMax-H3 uses only ``fc_in`` and - ``fc_out`` in each main transformer-block FFN. + ``layout_128x4`` scale layout, group size 16). + + Which linears get quantized is set by ``layer_prefixes``. The default is + the historical LTX-2 set, so ``NVFP4Config()`` behaves exactly as it did + before the field existed. Other models must pass their own set (see + :meth:`for_minimax_h3`) — a model whose layer paths are not covered + attaches no quant methods at all and silently runs dense, which is the + failure mode this field exists to remove. """ - def __init__(self, layer_profile: str = "refine", retain_original_weights: bool | None = None): + def __init__( + self, + layer_profile: str = "refine", + retain_original_weights: bool | None = None, + layer_prefixes: frozenset[str] | set[str] | list[str] | None = None, + exclude_prefixes: frozenset[str] | set[str] | list[str] | None = None, + ): super().__init__() # ``base``: stage-1 set (no attn2.to_out, no cross-modal AV - # projections). ``refine``: full stage-2 set. + # projections). ``refine``: full stage-2 set. LTX-2 streaming only: + # other models have no stage split (their ``_is_refine_only_layer`` is + # always False, so every quantized layer stays on the FP4 path). self.layer_profile = layer_profile # Original bf16 ``layer.weight`` retention after FP4 conversion. # Default (None/False): purge the purgeable originals -- every @@ -448,6 +528,29 @@ def __init__(self, layer_profile: str = "refine", retain_original_weights: bool # single-stage deploy. True: retain everything (debugging / # pre-purge behavior). self.retain_original_weights = retain_original_weights + # Full module paths of the linears to quantize. None -> the LTX-2 set + # (unchanged default). Frozen so a config instance is hashable and + # cannot be mutated after it has been handed to a model. + self.layer_prefixes: frozenset[str] = (frozenset(_LTX2_NVFP4_LINEAR_PREFIXES) + if layer_prefixes is None else frozenset(layer_prefixes)) + # Additional never-quantize patterns (full path or trailing suffix). + # Applied on top of ``_ALWAYS_EXCLUDED_LINEAR_SUFFIXES``, which no + # caller can override. + self.exclude_prefixes: frozenset[str] = frozenset(exclude_prefixes) if exclude_prefixes else frozenset() + + def is_nvfp4_linear_prefix(self, prefix: str) -> bool: + """Whether *prefix* is quantized under this config. + + Exclusions are checked first, and ``_ALWAYS_EXCLUDED_LINEAR_SUFFIXES`` + is unconditional: a prefix listed in ``layer_prefixes`` by mistake + (e.g. a glob that swept up ``attn.to_gate_compress``) still comes back + False here, so the gate cannot be quantized by any caller. + """ + if _matches_linear_suffix(prefix, _ALWAYS_EXCLUDED_LINEAR_SUFFIXES): + return False + if _matches_linear_suffix(prefix, self.exclude_prefixes): + return False + return prefix in self.layer_prefixes def get_name(self): return "nvfp4" @@ -463,20 +566,33 @@ def get_min_capability(cls): def get_config_filenames(): return [] + @classmethod + def for_minimax_h3(cls, **kwargs: Any) -> NVFP4Config: + """The MiniMax-H3 NVFP4 layer set (300 block linears). + + ``attn.to_gate_compress`` (the VSA sparse-attention gate) is excluded; + see ``MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES``. Keyword arguments + (e.g. ``retain_original_weights``) are forwarded to the constructor. + """ + kwargs.setdefault("layer_prefixes", MINIMAX_H3_NVFP4_LINEAR_PREFIXES) + kwargs.setdefault("exclude_prefixes", MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES) + return cls(**kwargs) + @classmethod def from_config(cls, config: dict[str, Any]) -> NVFP4Config: return cls( layer_profile=config.get("layer_profile", "refine"), retain_original_weights=config.get("retain_original_weights"), + layer_prefixes=config.get("layer_prefixes"), + exclude_prefixes=config.get("exclude_prefixes"), ) def get_quant_method(self, layer: torch.nn.Module, prefix: str): from fastvideo.layers.linear import LinearBase - # LTX-2 switches its active subset by stage at runtime. MiniMax-H3 - # uses the fixed main-transformer FFN set selected by its prefix. - if isinstance(layer, LinearBase) and (is_ltx2_nvfp4_linear_prefix(prefix) - or is_minimax_h3_nvfp4_linear_prefix(prefix)): + # Use the superset at build/load time, then switch active subset + # dynamically in NVFP4QuantizeMethod.apply based on stage profile. + if isinstance(layer, LinearBase) and self.is_nvfp4_linear_prefix(prefix): method = NVFP4QuantizeMethod(layer_prefix=prefix) method._retain_original_weights = self.retain_original_weights return method @@ -487,9 +603,6 @@ def convert_model_to_nvfp4(model: torch.nn.Module) -> None: SfLayout, _, _ = _require_flashinfer() from torch.distributed.tensor import DTensor # type: ignore - purged = 0 - retained = 0 - purged_bytes = 0 for mod in model.modules(): qm = getattr(mod, "quant_method", None) if isinstance(qm, NVFP4QuantizeMethod): @@ -522,25 +635,47 @@ def convert_model_to_nvfp4(model: torch.nn.Module) -> None: persistent=False, ) - retain_flag = getattr(qm, "_retain_original_weights", None) - # Refine-only layers are NEVER purgeable: the "base" stage profile - # runs them dense by deployment contract (the distilled - # single-stage deploy included — its forward context is the base - # profile, so e.g. audio_to_video_attn routes dense every step). - # retain_original_weights therefore only widens retention - # (True = keep everything); it cannot narrow it below the - # dense-capable set. - retain = qm._is_refine_only_layer or retain_flag is True - if retain: - retained += 1 - elif isinstance(weight, DTensor): - # ponytail: purging FSDP-sharded originals needs per-shard - # resharding bookkeeping; skip until a sharded deploy needs it. - retained += 1 - else: - purged_bytes += weight.numel() * weight.element_size() - purged += 1 - mod.register_parameter("weight", None) + _apply_dense_weight_policy(model) + + +def _apply_dense_weight_policy(model: torch.nn.Module) -> None: + """Drop the bf16 ``weight`` of every NVFP4 linear that no longer needs it. + + Refine-only layers are NEVER purgeable: the "base" stage profile runs them + dense by deployment contract (the distilled single-stage deploy included — + its forward context is the base profile, so e.g. ``audio_to_video_attn`` + routes dense every step). ``retain_original_weights`` therefore only widens + retention (True = keep everything); it cannot narrow it below the + dense-capable set. + + Shared by :func:`convert_model_to_nvfp4` and + :func:`load_nvfp4_checkpoint` so a restored sidecar frees the same memory a + fresh conversion would. + """ + from torch.distributed.tensor import DTensor # type: ignore + + purged = 0 + retained = 0 + purged_bytes = 0 + for mod in model.modules(): + qm = getattr(mod, "quant_method", None) + if not isinstance(qm, NVFP4QuantizeMethod): + continue + weight = getattr(mod, "weight", None) + if weight is None: + continue + retain_flag = getattr(qm, "_retain_original_weights", None) + retain = getattr(qm, "_is_refine_only_layer", False) or retain_flag is True + if retain: + retained += 1 + elif isinstance(weight, DTensor): + # ponytail: purging FSDP-sharded originals needs per-shard + # resharding bookkeeping; skip until a sharded deploy needs it. + retained += 1 + else: + purged_bytes += weight.numel() * weight.element_size() + purged += 1 + mod.register_parameter("weight", None) if purged or retained: logger.info( @@ -553,10 +688,368 @@ def convert_model_to_nvfp4(model: torch.nn.Module) -> None: ) +# --- compact NVFP4 checkpoint sidecar ------------------------------------- +# +# ``convert_model_to_nvfp4`` registers the packed FP4 tensors with +# ``persistent=False``, so they never enter a ``state_dict`` and a saved +# checkpoint carries dense bf16 weights plus a load-time quantization pass. +# A *sidecar* writes those tensors to their own safetensors file, keyed by +# module path, and restores them without re-running the conversion. +# +# Layout contract (must match ``apply`` / ``quantize_input`` / +# ``convert_model_to_nvfp4`` exactly, or the restored weights are garbage): +# +# ``_nvfp4_weight`` uint8 ``(out_dim, ceil(in_dim / 2))`` — two e2m1 +# codes per byte, K packed, N unpacked, exactly as +# ``nvfp4_quantize`` returns it (``apply`` passes +# ``.T`` to ``mm_fp4``). +# ``_nvfp4_weight_scale`` uint8 ``(out_dim, ceil(in_dim / 16))`` — e4m3 +# block-scale bit patterns, ``SfLayout.layout_128x4`` +# with ``do_shuffle=False`` (row-padded to the +# 128-row tile by ``_nvfp4_quantize`` and narrowed +# back; the padding is not stored). +# ``_weight_global_sf`` bfloat16 scalar — ``(448 * 6) / max|W|``. +# ``_nvfp4_alpha`` float32 scalar — ``1 / weight_global_sf`` at fp32 +# precision, so it is stored separately rather than +# recomputed from the bf16-rounded global sf. +# +# Block size is 16 throughout, and the per-row activation scale +# ``NVFP4QuantizeMethod.x_global_sf`` is not persisted: it is a constant +# ``1.0`` on the method (never data-derived), so it must be identical on both +# sides of a save/load. + +NVFP4_SIDECAR_SUFFIX = ".nvfp4.safetensors" +# Filename used when the sidecar sits inside a checkpoint *directory*; it does +# not carry the suffix above, which is what a sibling file is named with. +NVFP4_DIR_SIDECAR_NAME = "nvfp4.safetensors" +_NVFP4_SIDECAR_FORMAT = "fastvideo.nvfp4" +_NVFP4_SIDECAR_VERSION = 1 +_NVFP4_SIDECAR_METADATA_KEY = "fastvideo_nvfp4" +_NVFP4_SIDECAR_KEY_SEP = "::" +_NVFP4_SIDECAR_SF_LAYOUT = "layout_128x4" +_NVFP4_SIDECAR_DO_SHUFFLE = False +_NVFP4_SIDECAR_BLOCK_SIZE = 16 +# Order matters only for the manifest; every buffer is optional on load so a +# future format can add tensors without breaking older readers. +_NVFP4_SIDECAR_BUFFERS = ( + "_nvfp4_weight", + "_nvfp4_weight_scale", + "_weight_global_sf", + "_nvfp4_alpha", +) + + +def _sidecar_key(module_fqn: str, buffer_name: str) -> str: + return f"{module_fqn}{_NVFP4_SIDECAR_KEY_SEP}{buffer_name}" + + +def _split_sidecar_key(key: str) -> tuple[str, str]: + module_fqn, _, buffer_name = key.rpartition(_NVFP4_SIDECAR_KEY_SEP) + return module_fqn, buffer_name + + +def _nvfp4_tagged_modules(model: torch.nn.Module) -> list[tuple[str, torch.nn.Module, NVFP4QuantizeMethod]]: + tagged = [] + for fqn, mod in model.named_modules(): + qm = getattr(mod, "quant_method", None) + if isinstance(qm, NVFP4QuantizeMethod): + tagged.append((fqn, mod, qm)) + return tagged + + +def _is_dtensor(tensor: torch.Tensor) -> bool: + # Imported lazily: torch.distributed.tensor is not cheap to import and is + # absent on some builds. + try: + from torch.distributed.tensor import DTensor # type: ignore + except ImportError: # pragma: no cover - depends on the torch build + return False + return isinstance(tensor, DTensor) + + +def nvfp4_sidecar_state_dict(model: torch.nn.Module) -> dict[str, torch.Tensor]: + """Collect the quantized tensors of every NVFP4 linear in *model*. + + Keys are ``"::"`` and values are detached CPU + copies. Modules whose buffers are missing (never converted) are skipped; + the returned mapping is what :func:`save_nvfp4_checkpoint` writes. + + FSDP note: a DTensor buffer is saved as this rank's local shard, so a + sharded save is only reloadable into an identically sharded model. + """ + state: dict[str, torch.Tensor] = {} + for fqn, mod, _ in _nvfp4_tagged_modules(model): + for name in _NVFP4_SIDECAR_BUFFERS: + tensor = getattr(mod, name, None) + if tensor is None: + continue + if _is_dtensor(tensor): + tensor = tensor.to_local() # type: ignore[attr-defined] + state[_sidecar_key(fqn, name)] = tensor.detach().to("cpu", copy=True).contiguous() + return state + + +def save_nvfp4_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + extra_metadata: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Write the model's NVFP4 tensors to a compact sidecar safetensors file. + + The file is roughly the packed-FP4 size (0.5 byte/code plus 1 byte per 16 + codes) instead of the dense bf16 size — about 3.6x smaller for a + power-of-two K. Returns a receipt dict (also logged) describing the layer + count and both sizes. Raises ``RuntimeError`` when the model has no NVFP4 + linears, which usually means the config's ``layer_prefixes`` did not cover + the model's layer paths. + """ + from safetensors.torch import save_file + + state = nvfp4_sidecar_state_dict(model) + tagged = _nvfp4_tagged_modules(model) + if not tagged: + raise RuntimeError("No NVFP4 linear layers found in this model; nothing to serialize. " + "Check that the model was built with an NVFP4Config whose " + "layer_prefixes cover its layer paths (e.g. NVFP4Config.for_minimax_h3()).") + if not state: + raise RuntimeError(f"Found {len(tagged)} NVFP4-tagged linear layers but none carry quantized " + "buffers. Call convert_model_to_nvfp4(model) before saving a sidecar.") + + layers: dict[str, list[int]] = {} + quant_prefixes: dict[str, str] = {} + dense_bytes = 0 + for fqn, mod, qm in tagged: + packed = getattr(mod, "_nvfp4_weight", None) + weight = getattr(mod, "weight", None) + if packed is None and weight is None: + continue + if weight is not None: + out_dim, in_dim = int(weight.shape[0]), int(weight.shape[1]) + else: + # Dense weight already purged: K = 2 codes/byte (H3's dims are + # even, so the ceil in the packed shape is exact). + out_dim, in_dim = int(packed.shape[0]), int(packed.shape[1]) * 2 + layers[fqn] = [out_dim, in_dim] + quant_prefixes[fqn] = getattr(qm, "layer_prefix", "") or "" + dense_bytes += out_dim * in_dim * 2 + + metadata = { + "format": _NVFP4_SIDECAR_FORMAT, + "version": _NVFP4_SIDECAR_VERSION, + "sf_layout": _NVFP4_SIDECAR_SF_LAYOUT, + "do_shuffle": _NVFP4_SIDECAR_DO_SHUFFLE, + "block_size": _NVFP4_SIDECAR_BLOCK_SIZE, + "num_layers": len(layers), + "layers": layers, + "quant_prefixes": quant_prefixes, + "model_class": type(model).__name__, + } + if extra_metadata: + metadata.update(extra_metadata) + + payload = dict(state) + serialized_bytes = sum(t.numel() * t.element_size() for t in payload.values()) + save_file(payload, os.fspath(path), metadata={_NVFP4_SIDECAR_METADATA_KEY: json.dumps(metadata)}) + + receipt = { + "path": os.fspath(path), + "num_layers": len(layers), + "num_tensors": len(payload), + "quantized_bytes": serialized_bytes, + "dense_bfloat16_bytes": dense_bytes, + "compression_ratio": (dense_bytes / serialized_bytes) if serialized_bytes else 0.0, + } + logger.info( + "NVFP4 sidecar: wrote %d layers / %d tensors to %s (%.2f GiB quantized vs " + "%.2f GiB dense bf16, %.2fx smaller).", + receipt["num_layers"], + receipt["num_tensors"], + receipt["path"], + serialized_bytes / (1 << 30), + dense_bytes / (1 << 30), + receipt["compression_ratio"], + ) + return receipt + + +def read_nvfp4_sidecar_metadata(path: str | os.PathLike[str]) -> dict[str, Any]: + """Return the manifest of a sidecar file without materializing its tensors.""" + from safetensors import safe_open + + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + raw = handle.metadata() or {} + if _NVFP4_SIDECAR_METADATA_KEY not in raw: + raise ValueError(f"{os.fspath(path)} is not a FastVideo NVFP4 sidecar " + f"(no {_NVFP4_SIDECAR_METADATA_KEY!r} metadata).") + return json.loads(raw[_NVFP4_SIDECAR_METADATA_KEY]) + + +def nvfp4_sidecar_path_for(checkpoint_path: str | os.PathLike[str]) -> str: + """Conventional sidecar path for a transformer checkpoint or directory. + + ``.../transformer.safetensors`` -> ``.../transformer.nvfp4.safetensors``; + a directory -> ``/nvfp4.safetensors``. + """ + raw = os.fspath(checkpoint_path) + if os.path.isdir(raw): + return os.path.join(raw, NVFP4_DIR_SIDECAR_NAME) + if raw.endswith(".safetensors"): + return raw[:-len(".safetensors")] + NVFP4_SIDECAR_SUFFIX + return raw + NVFP4_SIDECAR_SUFFIX + + +def _sidecar_target_device(mod: torch.nn.Module, name: str) -> torch.device | None: + """Device the restored buffer should live on. + + Mirrors ``convert_model_to_nvfp4``, which registers the buffers on the + (local) weight's device; falls back to an existing buffer, then to the + module's parameter device so a purge-then-restore still lands on GPU. + """ + weight = getattr(mod, "weight", None) + if weight is not None and weight.device.type != "meta": + return weight.device + existing = getattr(mod, name, None) + if existing is not None and existing.device.type != "meta": + return existing.device + for param in mod.parameters(recurse=False): + if param.device.type != "meta": + return param.device + return None + + +def load_nvfp4_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + strict: bool = True, + purge_dense_weights: bool = True, +) -> int: + """Restore NVFP4 tensors from a sidecar, skipping ``convert_model_to_nvfp4``. + + Registers ``_nvfp4_weight`` / ``_nvfp4_weight_scale`` / + ``_weight_global_sf`` / ``_nvfp4_alpha`` on every NVFP4-tagged linear from + the sidecar, byte-for-byte as the conversion would have produced them. The + dense bf16 weights are never touched (they may be absent entirely), and no + flashinfer call is made — this works on a host that only needs to *serve* a + pre-quantized checkpoint. + + ``strict`` raises on any layer-set or shape mismatch (a sidecar that does + not describe this model); with ``strict=False`` the mismatches are logged + and skipped, leaving those layers unconverted. ``purge_dense_weights`` + applies the same retention policy as ``convert_model_to_nvfp4``. + + Returns the number of layers restored. + """ + from safetensors import safe_open + + tagged = _nvfp4_tagged_modules(model) + if not tagged: + raise RuntimeError("No NVFP4 linear layers are attached to this model, so a sidecar cannot be " + "restored. This is the silent-dense failure mode: the model's NVFP4Config " + "layer_prefixes do not cover its layer paths (for MiniMax-H3 use " + "NVFP4Config.for_minimax_h3()).") + + manifest = read_nvfp4_sidecar_metadata(path) + if manifest.get("format") != _NVFP4_SIDECAR_FORMAT: + raise ValueError(f"Unsupported NVFP4 sidecar format {manifest.get('format')!r} in {os.fspath(path)}.") + if int(manifest.get("version", -1)) != _NVFP4_SIDECAR_VERSION: + raise ValueError(f"Unsupported NVFP4 sidecar version {manifest.get('version')!r} in " + f"{os.fspath(path)} (this build reads version {_NVFP4_SIDECAR_VERSION}).") + # A layout mismatch is unrecoverable (the packed nibbles would be read with + # the wrong swizzle), so it is never downgraded by strict=False. + expected_layout = { + "sf_layout": _NVFP4_SIDECAR_SF_LAYOUT, + "do_shuffle": _NVFP4_SIDECAR_DO_SHUFFLE, + "block_size": _NVFP4_SIDECAR_BLOCK_SIZE, + } + for key, expected in expected_layout.items(): + if manifest.get(key) != expected: + raise ValueError(f"NVFP4 sidecar {os.fspath(path)} was written with {key}=" + f"{manifest.get(key)!r}, but this build quantizes with {key}={expected!r}.") + + saved_layers: dict[str, list[int]] = manifest.get("layers", {}) + model_fqns = {fqn for fqn, _, _ in tagged} + missing = sorted(model_fqns - set(saved_layers)) + extra = sorted(set(saved_layers) - model_fqns) + if missing or extra: + message = (f"NVFP4 sidecar {os.fspath(path)} does not match this model: " + f"{len(missing)} layers missing from the sidecar, {len(extra)} layers not in the model. " + f"First missing={missing[:3]}, first extra={extra[:3]}.") + if strict: + raise ValueError(message) + logger.warning("%s Restoring the intersection only.", message) + + restored = 0 + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + available = set(handle.keys()) + for fqn, mod, _ in tagged: + if fqn not in saved_layers: + continue + out_dim, in_dim = (int(value) for value in saved_layers[fqn]) + tensors: dict[str, torch.Tensor] = {} + for name in _NVFP4_SIDECAR_BUFFERS: + key = _sidecar_key(fqn, name) + if key not in available: + continue + tensor = handle.get_tensor(key) + expected = _expected_sidecar_shapes(name, out_dim, in_dim) + if tuple(tensor.shape) not in expected: + raise ValueError(f"NVFP4 sidecar tensor {key} has shape {tuple(tensor.shape)}, expected one of " + f"{list(expected)} for a ({out_dim}, {in_dim}) linear.") + device = _sidecar_target_device(mod, name) + if device is not None: + tensor = tensor.to(device=device, non_blocking=True) + tensors[name] = tensor + if "_nvfp4_weight" not in tensors or "_nvfp4_weight_scale" not in tensors: + message = (f"NVFP4 sidecar entry for {fqn!r} is incomplete (has " + f"{sorted(tensors)}); the packed weight and its block scales are both required.") + if strict: + raise ValueError(message) + logger.warning("%s Skipping this layer.", message) + continue + for name, tensor in tensors.items(): + mod.register_buffer(name, tensor, persistent=False) + restored += 1 + + if purge_dense_weights: + _apply_dense_weight_policy(model) + + logger.info("NVFP4 sidecar: restored %d quantized layers from %s (dense weights %s).", restored, os.fspath(path), + "purged per policy" if purge_dense_weights else "left in place") + return restored + + +def _expected_sidecar_shapes(name: str, out_dim: int, in_dim: int) -> tuple[tuple[int, ...], ...]: + """The shapes a fresh conversion could produce for *name*. + + ``_nvfp4_quantize`` narrows the packed weight back to the logical row count + but returns the block scales as the kernel emitted them, i.e. still padded + to the 128-row tile: a layer whose output dim is not a multiple of 128 gets + a scale tensor with more rows than the weight. Both are accepted so a + sidecar written by a real flashinfer conversion validates. + """ + if name == "_nvfp4_weight": + return ((out_dim, (in_dim + 1) // 2), ) + if name == "_nvfp4_weight_scale": + padded_rows = ((out_dim + 127) // 128) * 128 + shapes = {(out_dim, (in_dim + 15) // 16), (padded_rows, (in_dim + 15) // 16)} + return tuple(sorted(shapes)) + # The two global scales are scalars produced by conversion. + return ((), ) + + __all__ = [ + "MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES", + "MINIMAX_H3_NVFP4_LINEAR_PREFIXES", "NVFP4Config", "NVFP4QuantizeMethod", + "NVFP4_SIDECAR_SUFFIX", "convert_model_to_nvfp4", "is_ltx2_nvfp4_linear_prefix", - "is_minimax_h3_nvfp4_linear_prefix", + "load_nvfp4_checkpoint", + "nvfp4_sidecar_path_for", + "nvfp4_sidecar_state_dict", + "read_nvfp4_sidecar_metadata", + "save_nvfp4_checkpoint", ] diff --git a/fastvideo/layers/quantization/nvfp4_qat_config.py b/fastvideo/layers/quantization/nvfp4_qat_config.py index bcd9fd04c7..63c10b563a 100644 --- a/fastvideo/layers/quantization/nvfp4_qat_config.py +++ b/fastvideo/layers/quantization/nvfp4_qat_config.py @@ -52,6 +52,10 @@ # share a substring with Wan's "to_out"/"ffn.fc_in"/"ffn.fc_out") need to be # listed explicitly below. DEFAULT_FP4_LAYERS = ( + # MiniMax-H3 names its FFN "ff.", not "ffn." -- without these the + # FFN is silently left dense and QAD only covers attention. + "ff.fc_in", + "ff.fc_out", "ffn.fc_in", "ffn.fc_out", "to_q", diff --git a/fastvideo/layers/quantization/w4a16_config.py b/fastvideo/layers/quantization/w4a16_config.py new file mode 100644 index 0000000000..271aa75ca7 --- /dev/null +++ b/fastvideo/layers/quantization/w4a16_config.py @@ -0,0 +1,1099 @@ +# SPDX-License-Identifier: Apache-2.0 +"""W4A16 — 4-bit weight, 16-bit activation quantization for CUDA DiT inference. + +W4A16 means exactly what it says: the *weights* are stored as 4-bit integers +with a per-group scale (and zero point), and the *activations* stay in a +16-bit float (bf16/fp16). Nothing here quantizes activations, and there is no +fused W4A16 GEMM in this environment either (see "Kernel status" below). The +compute path is: dequantize the stored 4-bit codes back to the activation +dtype, then run an ordinary dense GEMM. + +Why this lane exists +-------------------- +W4A16 is a **primary deployment path for Ada-generation consumer GPUs** +(RTX 4090 24 GB, RTX 6000 Ada 48 GB). H3's 20B DiT does not fit those cards +in BF16, and Ada has no native NVFP4 tensor-core path — so the low-bit +options that actually apply there are INT8, FP8 and W4A16. W4A16 is the one +that shrinks the weight *storage* the most (5.0 bits/weight with +``group_size=64`` — 4 for the codes plus two fp32 per-group constants — versus +8 for INT8/FP8 and 16 for BF16), which is what +matters when the bottleneck is "does the checkpoint fit", not "how fast is +one GEMM". + +Kernel status — read this before quoting a number +------------------------------------------------- +**There is no W4A16 GEMM kernel in this repository or in its installed +dependencies.** ``fastvideo-kernel`` ships an INT8 GEMM +(``csrc/turbodiffusion/gemm/gemm.cu`` is ``int8_gemm`` over +``cutlass::NumericConverter``), FP4 attention for sm_100/sm_120, +and block-sparse attention — no 4-bit weight GEMM on any architecture. The +``int4`` matches in ``csrc/turbodiffusion`` are the CUDA 16-byte vector type, +not 4-bit quantization. ``autoawq`` / ``auto_gptq`` / ``gptqmodel`` / +``marlin`` / ``bitsandbytes`` are not installed. + +So what this module provides is a **correctness reference**: a faithful +quantize/dequantize pair and a layer method that runs it end-to-end. Its +steady-state weight storage really is 4-bit, but every forward pays a full +dequantize plus a dense 16-bit GEMM, so it is *slower* than BF16 and its +per-forward peak memory transiently includes one dense weight. It is the +schema and the wiring a real kernel would consume — not a speed path. Do not +benchmark it against BF16 and call the result "W4A16 on Ada". + +Relationship to the other precision lanes +----------------------------------------- +- ``NVFP4`` (``nvfp4_config.py``) — Blackwell (sm_100+) block-scaled FP4 with + a real FlashInfer GEMM. Not an Ada path. +- ``INT8Affine`` (``int8_affine_config.py``) — group-64 affine INT8, Ada-viable, + same "dequantize then dense GEMM" reference shape. +- ``FP8`` / ``AbsMaxFP8`` — also Ada-viable (sm_89 supports FP8 tensor cores). +- **This module** — half the weight storage of INT8/FP8, unconditionally + weight-only, no fused kernel. + +Design notes carried over from the neighbouring configs +------------------------------------------------------- +- **Load-time conversion, dense allocation.** ``W4A16QuantizeMethod.create_weights`` + allocates the same dense Parameter an unquantized linear would, so a plain + BF16 checkpoint loads unchanged; the 4-bit codes arrive afterwards from + :func:`convert_model_to_w4a16` (or lazily on first forward). Same shape as + ``NVFP4QuantizeMethod`` / ``INT8AffineQuantizeMethod``. +- **Layer selection is a constructor field, not a hardcoded model list.** + ``target_layers`` is an explicit allowlist of full module paths; + ``layer_suffixes`` is a generic suffix rule for models without an enumerable + list. ``NVFP4Config``'s hardcoded LTX-2 frozenset is deliberately not + repeated here. +- **A fail-closed deny list.** ``attn.to_gate_compress`` is H3's VSA + sparse-attention compression gate. Quantizing it perturbs a *discrete* + routing decision, so an error there is not a small output perturbation — it + changes which tiles sparse attention attends to. The constructor can only + ever *add* exclusions, never remove one, so no caller can quantize it. +""" + +from __future__ import annotations + +import json +import logging +import os +from collections.abc import Iterable +from typing import Any + +import torch +import torch.nn.functional as F +from torch.nn.parameter import Parameter + +from fastvideo.layers.quantization.base_config import ( + QuantizationConfig, + QuantizeMethodBase, +) +from fastvideo.models.utils import set_weight_attrs + +logger = logging.getLogger(__name__) + +DEFAULT_GROUP_SIZE = 64 +DEFAULT_BITS = 4 +# Affine codes span [0, 2**bits - 1]; for bits=4 that is [0, 15], stored as +# torch.uint8 (two codes per byte, see ``w4a16_quantize``). +# Guard against a degenerate group (all-equal values) producing a zero scale +# and a division by zero in the zero-point solve. +_EPS = 1e-8 +# A model with 362 targeted linears would otherwise emit 362 copies of the +# "loader hook did not fire" warning on its first forward. +_LAZY_CONVERSION_WARNED = False + +# --------------------------------------------------------------------------- +# Group-wise affine 4-bit quantizer +# --------------------------------------------------------------------------- + + +def _group(w: torch.Tensor, group_size: int) -> torch.Tensor: + """View the last dim as ``(num_groups, group_size)``. + + Grouping is along the last (input/contraction) axis, matching + ``int8_affine_config._group`` and MLX's affine quantizer. + """ + if w.shape[-1] % group_size != 0: + raise ValueError(f"Last dim {w.shape[-1]} is not divisible by group_size {group_size}.") + return w.reshape(*w.shape[:-1], w.shape[-1] // group_size, group_size) + + +def w4a16_quantize( + w: torch.Tensor, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Group-wise affine quantization of ``w`` to ``bits``-bit unsigned codes. + + Per group of ``group_size`` contiguous values along the last axis, this + solves ``w ~= (code - zero) * scale`` with a **min/max affine** fit: + ``scale = (max - min) / (2**bits - 1)`` and ``zero`` the code that + reproduces ``min`` exactly, so both endpoints of the group round-trip + exactly. Codes are ``rint``-rounded (round-half-to-even) and clamped to + ``[0, 2**bits - 1]``. + + The fit is done in fp32 regardless of ``w.dtype``, so a bf16 checkpoint + value (which converts to fp32 exactly) yields a higher-precision scale + store than solving in bf16 would. + + Returns ``(codes, scales, zeros)``: + + - ``codes`` — ``torch.uint8``, shape ``w.shape[:-1] + (K // 2,)``, **two + 4-bit codes packed per byte**. For ``bits=4`` only; ``bits=8`` returns + one code per byte at shape ``w.shape``. + - ``scales`` — ``torch.float32``, shape ``w.shape[:-1] + (K // group_size,)``. + - ``zeros`` — ``torch.float32``, same shape as ``scales``. + + Packing convention (ours — no kernel consumes it yet): the **low nibble is + the lower K index**, i.e. ``packed[..., j] = codes[..., 2j] | codes[..., 2j+1] << 4``. + A future kernel has to be written against this layout; it does not match + AWQ's interleaved layout, GPTQ's ``g_idx`` layout, or bitsandbytes' order. + """ + if bits != 4 and bits != 8: + raise ValueError(f"W4A16 stores codes as uint8; bits must be 4 or 8, got {bits}") + if group_size <= 0: + raise ValueError(f"group_size must be positive, got {group_size}") + + w32 = w.detach().float().nan_to_num() + max_code = float((1 << bits) - 1) + grouped = _group(w32, group_size) + + w_min = grouped.amin(dim=-1) + w_max = grouped.amax(dim=-1) + scales = (w_max - w_min).clamp_min(_EPS) / max_code + zeros = torch.round(-w_min / scales).clamp_(0.0, max_code) + + codes = torch.round(grouped / scales.unsqueeze(-1) + zeros.unsqueeze(-1)) + codes = codes.clamp_(0.0, max_code).to(torch.uint8).reshape(w32.shape) + + if bits == 8: + return codes, scales, zeros + return _pack_4bit(codes), scales, zeros + + +def _pack_4bit(codes: torch.Tensor) -> torch.Tensor: + """Pack ``[0, 15]`` codes two-per-byte along the last axis. + + Low nibble = lower K index. The last dim must be even, which every H3 + linear input dim is (5376, 7168, 14336, 2688 are all even). + """ + if codes.shape[-1] % 2: + raise ValueError(f"4-bit packing needs an even last dim, got {codes.shape[-1]}") + codes = codes.to(torch.uint8).reshape(*codes.shape[:-1], codes.shape[-1] // 2, 2) + low = codes[..., 0] + high = codes[..., 1] & 0x0F + return (low | (high << 4)).contiguous() + + +def _unpack_4bit(packed: torch.Tensor) -> torch.Tensor: + """Inverse of :func:`_pack_4bit`; returns uint8 codes at ``2 * packed.shape[-1]``.""" + low = packed & 0x0F + high = (packed >> 4) & 0x0F + return torch.stack((low, high), dim=-1).reshape(*packed.shape[:-1], packed.shape[-1] * 2) + + +def w4a16_dequantize( + codes: torch.Tensor, + scales: torch.Tensor, + zeros: torch.Tensor, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + out_shape: tuple[int, ...] | torch.Size | None = None, + out_dtype: torch.dtype | None = None, +) -> torch.Tensor: + """Reconstruct the dense weight from 4-bit codes, group scales and zeros. + + ``codes`` is the packed tensor :func:`w4a16_quantize` returned (or, for + ``bits=8``, the unpacked one). Pass ``out_shape`` to name the logical + weight shape; otherwise the reconstruction keeps the code layout's own + last dim (``2 * packed.shape[-1]`` for 4-bit). + + The arithmetic runs in fp32 and is cast to ``out_dtype`` at the end — + ``(code - zero) * scale`` in fp32 is the more accurate side of the split, + same choice ``int8_affine_config`` makes. + """ + if bits == 4: + codes = _unpack_4bit(codes) + if out_shape is not None and tuple(codes.shape) != tuple(out_shape): + codes = codes.reshape(out_shape) + grouped = _group(codes.float(), group_size) + dense = (grouped - zeros.unsqueeze(-1)) * scales.unsqueeze(-1) + dense = dense.reshape(codes.shape) + return dense if out_dtype is None else dense.to(out_dtype) + + +# --- MiniMax-H3 layer set ------------------------------------------------- +# +# H3's DiT (``fastvideo/models/dits/minimax_h3.py``) names its linears +# ``{prefix}.{scope}.{i}.{suffix}`` with ``prefix="minimax_h3"``, +# ``num_layers=50``, ``num_refiner_layers=2`` +# (``fastvideo/configs/models/dits/minimax_h3.py``). Both block stacks hold the +# same ``MiniMaxH3Attention`` / ``MiniMaxH3FeedForward``; only the main stack +# has ``adaln_proj``. +MINIMAX_H3_PREFIX = "minimax_h3" +MINIMAX_H3_NUM_LAYERS = 50 +MINIMAX_H3_NUM_REFINER_LAYERS = 2 +MINIMAX_H3_BLOCK_SCOPES: tuple[str, ...] = ( + "transformer_blocks", + "token_refiner.refiner_blocks", +) +MINIMAX_H3_BLOCK_LINEAR_SUFFIXES: tuple[str, ...] = ( + "attn.to_q", + "attn.to_k", + "attn.to_v", + "attn.to_out", + "ff.fc_in", + "ff.fc_out", +) +# The main stack additionally carries the per-block AdaLN modulation GEMM. +MINIMAX_H3_MAIN_STACK_LINEAR_SUFFIXES: tuple[str, ...] = MINIMAX_H3_BLOCK_LINEAR_SUFFIXES + ("adaln_proj.linear", ) + +# Suffix rules for a generic (non-H3) transformer, so the config is usable +# before a model has an enumerable prefix list. +_GENERIC_LINEAR_SUFFIXES: tuple[str, ...] = ( + "attn.to_q", + "attn.to_k", + "attn.to_v", + "attn.to_out", + "ff.fc_in", + "ff.fc_out", +) + +# Linears that must NEVER be quantized, whatever a caller passes in +# ``target_layers``. ``attn.to_gate_compress`` is H3's VSA sparse-attention +# compression gate: its output steers a *discrete* tile-selection decision, so +# quantizing it does not merely perturb the output, it can change which tiles +# the sparse attention reads. H3 also probes the loaded gate structurally +# (``MiniMaxH3Attention._gate_active`` tests ``weight != 0`` once to skip a +# guaranteed-zero branch), which a dequantized weight would break. It matches +# no generic "norm"/"embedder" exclusion heuristic, so it is named explicitly. +_NEVER_QUANTIZE_SUBSTRINGS: tuple[str, ...] = ("attn.to_gate_compress", ) + +# Modules H3 already pins to fp32 (``MiniMaxH3Transformer3DModel._keep_in_fp32_modules``). +_H3_FP32_KEPT_SUBSTRINGS: tuple[str, ...] = ( + "proj_in", + "audio_proj_in", + "proj_out", + "audio_proj_out", + "time_embedder", +) + + +def _matches_linear_suffix(prefix: str, suffixes: Iterable[str]) -> bool: + """True when *prefix* is one of *suffixes* or ends at a dot boundary. + + The dot boundary keeps ``"ff.fc_in"`` from matching a hypothetical + ``"cross_ff.fc_in"``. + """ + return any(prefix == suffix or prefix.endswith("." + suffix) for suffix in suffixes) + + +def minimax_h3_w4a16_prefixes( + *, + prefix: str = MINIMAX_H3_PREFIX, + num_layers: int = MINIMAX_H3_NUM_LAYERS, + num_refiner_layers: int = MINIMAX_H3_NUM_REFINER_LAYERS, +) -> frozenset[str]: + """Enumerate the exact H3 linear prefixes the H3 W4A16 profile targets. + + Built from H3's real module names as constructed in + ``fastvideo/models/dits/minimax_h3.py``: ``MiniMaxH3TransformerBlock`` + builds ``{prefix}.transformer_blocks.{i}.attn`` / ``.ff`` / ``.adaln_proj``; + ``MiniMaxH3TokenRefiner`` builds ``{prefix}.token_refiner.refiner_blocks.{i}.attn`` / ``.ff``. + + 50 main blocks x 7 linears + 2 refiner blocks x 6 linears = **362 linears**. + + This is the allowlist :meth:`W4A16Config.for_minimax_h3` hands to + ``target_layers``. It is a plain function so a caller (or a test) can + regenerate it from the architecture constants rather than trusting a + literal — the H3 profile is *derived*, not hardcoded. + """ + prefixes: set[str] = set() + for index in range(num_layers): + for suffix in MINIMAX_H3_MAIN_STACK_LINEAR_SUFFIXES: + prefixes.add(f"{prefix}.transformer_blocks.{index}.{suffix}") + for index in range(num_refiner_layers): + for suffix in MINIMAX_H3_BLOCK_LINEAR_SUFFIXES: + prefixes.add(f"{prefix}.token_refiner.refiner_blocks.{index}.{suffix}") + return frozenset(prefixes) + + +class W4A16Config(QuantizationConfig): + """Weight-only 4-bit (group-wise affine) quantization with 16-bit activations. + + Layer selection is a constructor field, not a hardcoded model list: + ``target_layers`` is an explicit allowlist of full module paths and takes + precedence when given; otherwise ``layer_suffixes`` is matched with a + dot-boundary suffix rule. Both are subject to the fail-closed deny list + (``_NEVER_QUANTIZE_SUBSTRINGS``), which the constructor can only widen. + + Weight-only by construction: there is no activation quantizer and + :class:`W4A16QuantizeMethod` runs a dense 16-bit GEMM over a dequantized + weight. Use :meth:`for_minimax_h3` for the H3 profile. + """ + + def __init__( + self, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + target_layers: Iterable[str] | None = None, + layer_suffixes: Iterable[str] | None = None, + exclude_substrings: Iterable[str] | None = None, + retain_original_weight: bool = True, + ) -> None: + super().__init__() + if bits not in (4, 8): + raise ValueError(f"W4A16 stores codes as uint8; bits must be 4 or 8, got {bits}") + if group_size <= 0: + raise ValueError(f"group_size must be positive, got {group_size}") + self.group_size = group_size + self.bits = bits + self.target_layers: frozenset[str] | None = (None if target_layers is None else frozenset(target_layers)) + self.layer_suffixes: tuple[str, ...] = (tuple(_GENERIC_LINEAR_SUFFIXES) + if layer_suffixes is None else tuple(layer_suffixes)) + # Fail-closed deny list. The hard exclusions are always present; + # ``exclude_substrings`` can only add, never remove -- this is what + # keeps ``attn.to_gate_compress`` unquantizable no matter what a + # caller passes as ``target_layers``. + self.exclude_substrings: tuple[str, ...] = tuple(_NEVER_QUANTIZE_SUBSTRINGS) + tuple(exclude_substrings or ()) + # Keep the dense bf16 ``layer.weight`` Parameter after conversion. + # Default True and it matters here: ``MiniMaxH3AdaLayerNormModulation.forward`` + # reads ``self.linear.weight.dtype`` to cast its input, so purging the + # weight raises AttributeError. Setting False frees the bf16 copy at + # the cost of requiring every caller to stop touching ``.weight``. + self.retain_original_weight = retain_original_weight + + def get_name(self) -> str: + return "W4A16" + + def get_supported_act_dtypes(self) -> list[torch.dtype]: + return [torch.bfloat16, torch.float16, torch.float32] + + @classmethod + def get_min_capability(cls) -> int: + """Turing (75). + + The compute path is a plain 16-bit GEMM over a dequantized weight, so + no 4-bit tensor-core class is required and the config stays loadable + wherever the other reference paths are. **This is not a claim that + 4-bit runs fast there.** The deployment target for this lane is Ada + (sm_89, RTX 4090 / RTX 6000 Ada); making it a *fast* path needs a + fused W4A16 GEMM that does not exist in this repository yet. + """ + return 75 + + @staticmethod + def get_config_filenames() -> list[str]: + return [] + + @classmethod + def from_config(cls, config: dict[str, Any]) -> W4A16Config: + return cls( + group_size=config.get("group_size", DEFAULT_GROUP_SIZE), + bits=config.get("bits", DEFAULT_BITS), + target_layers=config.get("target_layers"), + layer_suffixes=config.get("layer_suffixes"), + exclude_substrings=config.get("exclude_substrings"), + retain_original_weight=config.get("retain_original_weight", True), + ) + + @classmethod + def for_minimax_h3( + cls, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + retain_original_weight: bool = True, + ) -> W4A16Config: + """The MiniMax-H3 profile: 362 attention / FFN / AdaLN linears. + + The allowlist comes from :func:`minimax_h3_w4a16_prefixes`, and the + deny list additionally carries H3's fp32-pinned modules (``proj_in``, + ``audio_proj_in``, ``time_embedder``, ``proj_out``, ``audio_proj_out``) + — those are excluded both by construction (they are not in the + allowlist) and by name, so a later widening of the allowlist cannot + silently reach them. + + ``group_size`` must divide the input dim of every targeted linear: + with the released H3 config (hidden 5376, inner 7168, ffn 14336, + adaln 2688) that holds for 32, 64 and 128. + """ + return cls( + group_size=group_size, + bits=bits, + target_layers=minimax_h3_w4a16_prefixes(), + exclude_substrings=_H3_FP32_KEPT_SUBSTRINGS, + retain_original_weight=retain_original_weight, + ) + + def is_target_layer(self, prefix: str) -> bool: + """Whether ``prefix`` is quantized under this config. + + Deny list first (fail-closed), then ``target_layers`` if supplied, + else suffix matching. Non-``LinearBase`` layers are filtered by + :meth:`get_quant_method`, not here, so this is safe to call on any + module name. + """ + for banned in self.exclude_substrings: + if banned in prefix: + return False + if self.target_layers is not None: + return prefix in self.target_layers + return _matches_linear_suffix(prefix, self.layer_suffixes) + + def get_quant_method(self, layer: torch.nn.Module, prefix: str): + from fastvideo.layers.linear import LinearBase + + if not isinstance(layer, LinearBase) or not self.is_target_layer(prefix): + return None + # A group must sit entirely inside one weight row. Skipping (rather + # than raising) keeps a config change from hard-failing model + # construction; the warning is what makes the skip visible. + input_size = getattr(layer, "input_size", None) + if input_size is not None and input_size % self.group_size: + logger.warning( + "W4A16: skipping layer %r — input dim %d is not divisible by group_size %d. " + "The layer runs dense.", prefix, input_size, self.group_size) + return None + if self.bits == 4 and input_size is not None and input_size % 2: + logger.warning( + "W4A16: skipping layer %r — input dim %d is odd, so 4-bit codes cannot be " + "packed two-per-byte. The layer runs dense.", prefix, input_size) + return None + return W4A16QuantizeMethod( + layer_prefix=prefix, + group_size=self.group_size, + bits=self.bits, + retain_original_weight=self.retain_original_weight, + ) + + +class W4A16QuantizeMethod(QuantizeMethodBase): + """Linear method for weight-only 4-bit affine quantization. + + ``create_weights`` allocates the same dense Parameter an unquantized linear + would (so a BF16 checkpoint loads unchanged), and the 4-bit codes, group + scales and zeros arrive later as non-persistent buffers from + :func:`convert_model_to_w4a16` — conversion happens at *load* time, not + construction time, exactly mirroring ``NVFP4QuantizeMethod`` and + ``INT8AffineQuantizeMethod``. + + ``apply`` is the **reference path**: dequantize the whole weight to the + activation dtype, then ``F.linear``. It is bit-for-bit a dense 16-bit GEMM + over an approximation of the original weight, which is what makes it + useful as a correctness oracle — and it is why it is not a performance + path (see the module docstring's "Kernel status"). + """ + + def __init__( + self, + layer_prefix: str = "", + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, + retain_original_weight: bool = True, + ) -> None: + super().__init__() + self.layer_prefix = layer_prefix + self.group_size = group_size + self.bits = bits + self.retain_original_weight = retain_original_weight + + def create_weights( + self, + layer: torch.nn.Module, + input_size_per_partition: int, + output_partition_sizes: list[int], + input_size: int, + output_size: int, + params_dtype: torch.dtype, + **extra_weight_attrs, + ): + weight = Parameter( + torch.empty( + sum(output_partition_sizes), + input_size_per_partition, + dtype=params_dtype, + ), + requires_grad=False, + ) + set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0}) + layer.register_parameter("weight", weight) + set_weight_attrs(weight, extra_weight_attrs) + + def _ensure_quantized(self, layer: torch.nn.Module) -> bool: + """Convert on first use if the loader hook never ran. + + Returns False when the layer is intentionally left dense (grad-enabled + forward: a training step must see the master weight, not a frozen + dequantized copy). The loader path is ``_maybe_quantize_model`` -> + :func:`convert_model_to_w4a16`; this fallback exists so the config is + still *correct* if that dispatch is missing, but it warns because + reaching it means the loader hook did not fire. + """ + if getattr(layer, "_w4a16_codes", None) is not None: + return True + weight = getattr(layer, "weight", None) + if weight is None: + raise RuntimeError(f"W4A16 layer {self.layer_prefix!r} has no weight and no quantized buffers.") + if torch.is_grad_enabled(): + return False + global _LAZY_CONVERSION_WARNED + if not _LAZY_CONVERSION_WARNED: + _LAZY_CONVERSION_WARNED = True + logger.warning( + "W4A16: layer %r reached apply() unconverted; converting lazily (this message is logged once " + "per process, not once per layer). The loader hook (_maybe_quantize_model) did not dispatch to " + "convert_model_to_w4a16 — check its isinstance chain in " + "fastvideo/models/loader/fsdp_load.py.", self.layer_prefix) + _quantize_layer_weight(layer, weight, group_size=self.group_size, bits=self.bits) + return True + + def apply( + self, + layer: torch.nn.Module, + x: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + if not self._ensure_quantized(layer): + weight = layer.weight + if weight is None: + raise RuntimeError( + f"W4A16 layer {self.layer_prefix!r} is in the dense (grad-enabled) branch, but its " + "original bf16 weight was purged (W4A16Config(retain_original_weight=False)). Training " + "needs the master weight; load with retain_original_weight left at its default for any " + "run that takes gradients.") + return F.linear(x, weight.to(x.dtype) if weight.dtype != x.dtype else weight, bias) + + # Reference path: one dense dequantize per forward, then a normal + # 16-bit GEMM. No cached dense copy -- caching it would hold both the + # 4-bit codes and a full bf16 weight on the device, which is exactly + # the memory this lane exists to avoid. + weight = w4a16_dequantize( + layer._w4a16_codes, + layer._w4a16_scales, + layer._w4a16_zeros, + group_size=self.group_size, + bits=self.bits, + out_shape=layer._w4a16_weight_shape, + out_dtype=x.dtype, + ) + return F.linear(x, weight, bias) + + +# --------------------------------------------------------------------------- +# Load-time conversion +# --------------------------------------------------------------------------- + + +def _quantize_layer_weight( + mod: torch.nn.Module, + weight: torch.Tensor, + *, + group_size: int = DEFAULT_GROUP_SIZE, + bits: int = DEFAULT_BITS, +) -> None: + """Quantize one linear's weight in place into non-persistent buffers.""" + from torch.distributed.tensor import DTensor # type: ignore + + weight_local = weight.to_local() if isinstance(weight, DTensor) else weight # type: ignore[arg-type] + if weight_local.shape[-1] % group_size: + raise ValueError(f"W4A16 layer {mod!r}: input dim {weight_local.shape[-1]} is not divisible by " + f"group_size {group_size}.") + codes, scales, zeros = w4a16_quantize(weight_local, group_size=group_size, bits=bits) + mod.register_buffer("_w4a16_codes", codes.contiguous(), persistent=False) + mod.register_buffer("_w4a16_scales", scales.to(torch.float32).contiguous(), persistent=False) + mod.register_buffer("_w4a16_zeros", zeros.to(torch.float32).contiguous(), persistent=False) + # The logical weight shape is recorded rather than inferred at dequantize + # time: the packed code layout is (..., K // 2) and the orthogonal shape + # is not recoverable from it alone. + mod._w4a16_weight_shape = tuple(weight_local.shape) + + +def convert_model_to_w4a16(model: torch.nn.Module) -> None: + """Quantize every W4A16-tagged linear in-place after weights load. + + Mirrors ``convert_model_to_nvfp4`` / ``convert_model_to_int8_affine``: walk + the module tree once, convert each layer whose ``quant_method`` is a + :class:`W4A16QuantizeMethod`, and register the 4-bit codes plus per-group + scales/zeros as non-persistent buffers (so they are not written back into + ``state_dict``/checkpoints). + + Callers: the loader hook ``_maybe_quantize_model`` in + ``fastvideo/models/loader/fsdp_load.py``. *That hook is not edited by this + module* — it dispatches on an explicit ``isinstance`` chain, so it needs a + matching branch (see the module report). Without it, + :meth:`W4A16QuantizeMethod.apply` converts lazily on first forward and logs + a warning, so inference is still correct, just later and noisier. + """ + converted = 0 + purged = 0 + schemes: set[tuple[int, int]] = set() + for mod in model.modules(): + qm = getattr(mod, "quant_method", None) + if not isinstance(qm, W4A16QuantizeMethod): + continue + weight = getattr(mod, "weight", None) + if weight is None: + continue + _quantize_layer_weight(mod, weight, group_size=qm.group_size, bits=qm.bits) + converted += 1 + schemes.add((qm.group_size, qm.bits)) + if not qm.retain_original_weight: + # register_parameter(None) (as convert_model_to_nvfp4 does) rather + # than popping the key: `layer.weight` then reads as None instead + # of raising AttributeError. + original = mod._parameters.get("weight") + if original is not None: + original.grad = None + mod.register_parameter("weight", None) + purged += 1 + + if converted: + # Say plainly what was produced. This is the 4-bit *storage* receipt, + # not a throughput claim: apply() still runs a dense 16-bit GEMM. + logger.info( + "W4A16 conversion receipt: quantized %d linear layers (%s, reference dequantize-then-GEMM " + "path); purged %d original bf16 weight tensors.", converted, + ", ".join(f"group_size={g}, bits={b}" for g, b in sorted(schemes)), purged) + logger.info("W4A16: no fused 4-bit GEMM is available in this build, so forward compute is a dense " + "16-bit GEMM over a dequantized weight. Expect BF16-comparable VRAM for the transient " + "dense weight and slower-than-BF16 step times.") + + +# --------------------------------------------------------------------------- +# Compact checkpoint sidecar — save/load of the quantized payload +# --------------------------------------------------------------------------- +# +# The 4-bit buffers are registered with ``persistent=False`` (see +# ``_quantize_layer_weight``), so ``state_dict()`` does NOT carry them: a saved +# checkpoint holds dense bf16 weights only and the 4-bit payload is rebuilt by +# re-running the conversion at every load. That is impossible on a host that +# cannot hold the dense weights at all — the target cards for this lane are +# 24-48 GB Ada parts and a 32 GB RTX 5090, and ``convert_model_to_w4a16`` +# starts from the dense weight. The sidecar is what makes a pre-quantized +# checkpoint servable on that hardware. +# +# Format (identical in shape to the NVFP4 sidecar, ``nvfp4_config.py``): one +# safetensors file keyed ``"::"`` plus a JSON manifest +# under the ``fastvideo_w4a16`` metadata key describing the scheme. +# +# Buffer inventory restored by a load — these are the same tensors +# ``_quantize_layer_weight`` registers, byte for byte: +# +# ``_w4a16_codes`` uint8 two 4-bit codes packed per byte, shape +# ``(out_dim, in_dim // 2)`` for ``bits=4`` and the full +# ``(out_dim, in_dim)`` for ``bits=8``. The packing is +# the module's own convention (``_pack_4bit``): the **low +# nibble is the lower K index**, i.e. +# ``packed[..., j] = codes[..., 2j] | codes[..., 2j+1] << 4``. +# It matches neither AWQ, GPTQ nor bitsandbytes, so the +# bytes are only meaningful to a reader that unpacks them +# the same way. +# ``_w4a16_scales`` float32 ``(out_dim, in_dim // group_size)`` +# ``_w4a16_zeros`` float32 ``(out_dim, in_dim // group_size)`` +# +# A load also restores ``mod._w4a16_weight_shape``. That one is *not* a buffer +# — it is a plain tuple attribute set by ``_quantize_layer_weight``, and +# ``W4A16QuantizeMethod.apply`` passes it to ``w4a16_dequantize(out_shape=...)`` +# because the logical weight shape is not recoverable from the packed code +# shape alone. A sidecar load that skipped it would leave every layer raising +# ``AttributeError`` on first forward, so it is restored from the manifest's +# per-layer ``weight_shape``. + +W4A16_SIDECAR_SUFFIX = ".w4a16.safetensors" +# Filename used when the sidecar sits inside a checkpoint *directory*; it does +# not carry the suffix above, which is what a sibling file is named with. +W4A16_DIR_SIDECAR_NAME = "w4a16.safetensors" +_W4A16_SIDECAR_FORMAT = "fastvideo.w4a16" +_W4A16_SIDECAR_VERSION = 1 +_W4A16_SIDECAR_METADATA_KEY = "fastvideo_w4a16" +_W4A16_SIDECAR_KEY_SEP = "::" +# Order matters only for the manifest; every buffer is optional on load so a +# future format can add tensors without breaking older readers. +_W4A16_SIDECAR_BUFFERS = ( + "_w4a16_codes", + "_w4a16_scales", + "_w4a16_zeros", +) +# The exact container each buffer must have. A sidecar whose codes came back +# as int8 (or float32) would unpack to garbage with no error anywhere, so +# these are validated, never coerced. +_W4A16_SIDECAR_DTYPES = { + "_w4a16_codes": torch.uint8, + "_w4a16_scales": torch.float32, + "_w4a16_zeros": torch.float32, +} + + +def _sidecar_key(module_fqn: str, buffer_name: str) -> str: + return f"{module_fqn}{_W4A16_SIDECAR_KEY_SEP}{buffer_name}" + + +def _w4a16_tagged_modules(model: torch.nn.Module) -> list[tuple[str, torch.nn.Module, W4A16QuantizeMethod]]: + tagged = [] + for fqn, mod in model.named_modules(): + qm = getattr(mod, "quant_method", None) + if isinstance(qm, W4A16QuantizeMethod): + tagged.append((fqn, mod, qm)) + return tagged + + +def _is_dtensor(tensor: torch.Tensor) -> bool: + # Imported lazily: torch.distributed.tensor is not cheap to import and is + # absent on some builds. + try: + from torch.distributed.tensor import DTensor # type: ignore + except ImportError: # pragma: no cover - depends on the torch build + return False + return isinstance(tensor, DTensor) + + +def w4a16_sidecar_state_dict(model: torch.nn.Module) -> dict[str, torch.Tensor]: + """Collect the quantized tensors of every W4A16 linear in *model*. + + Keys are ``"::"`` and values are detached CPU + copies. Modules whose buffers are missing (never converted) are skipped; + the returned mapping is what :func:`save_w4a16_checkpoint` writes. + + ``_w4a16_weight_shape`` is not a tensor and so is not collected here; it + travels in the manifest instead (see :func:`save_w4a16_checkpoint`). + + FSDP note: a DTensor buffer is saved as this rank's local shard, so a + sharded save is only reloadable into an identically sharded model. + """ + state: dict[str, torch.Tensor] = {} + for fqn, mod, _ in _w4a16_tagged_modules(model): + for name in _W4A16_SIDECAR_BUFFERS: + tensor = getattr(mod, name, None) + if tensor is None: + continue + if _is_dtensor(tensor): + tensor = tensor.to_local() # type: ignore[attr-defined] + state[_sidecar_key(fqn, name)] = tensor.detach().to("cpu", copy=True).contiguous() + return state + + +def save_w4a16_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + extra_metadata: dict[str, Any] | None = None, +) -> dict[str, Any]: + """Write the model's W4A16 tensors to a compact sidecar safetensors file. + + The file is roughly the 4-bit size (5.0 bits/weight at ``group_size=64``: + 4 for the packed codes plus two fp32 per group) instead of the dense bf16 + size. Returns a receipt dict (also logged) with the module count and both + sizes. Raises ``RuntimeError`` when the model has no W4A16 linears, which + usually means the config's layer selection did not cover the model's layer + paths, and when the tagged layers carry no buffers (never converted). + + The manifest under the ``fastvideo_w4a16`` metadata key carries the format + name/version, the scheme (``group_size``/``bits``), the per-layer weight and + buffer shapes, and the quantized module fqns (the keys of ``layers``), so a + loader can validate a sidecar against a model without materializing the + tensors. The per-layer ``weight_shape`` is what a load uses to restore + ``_w4a16_weight_shape``, which the packed codes cannot express. + """ + from safetensors.torch import save_file + + state = w4a16_sidecar_state_dict(model) + tagged = _w4a16_tagged_modules(model) + if not tagged: + raise RuntimeError("No W4A16 linear layers found in this model; nothing to serialize. Check that the " + "model was built with a W4A16Config whose layer selection covers its layer paths " + "(e.g. W4A16Config.for_minimax_h3()).") + if not state: + raise RuntimeError(f"Found {len(tagged)} W4A16-tagged linear layers but none carry quantized buffers. " + "Call convert_model_to_w4a16(model) before saving a sidecar.") + + layers: dict[str, dict[str, Any]] = {} + quant_prefixes: dict[str, str] = {} + dense_bytes = 0 + group_sizes: set[int] = set() + bit_widths: set[int] = set() + for fqn, mod, qm in tagged: + codes = getattr(mod, "_w4a16_codes", None) + weight = getattr(mod, "weight", None) + if codes is None and weight is None: + continue + if weight is not None: + weight_shape = [int(dim) for dim in weight.shape] + elif getattr(mod, "_w4a16_weight_shape", None) is not None: + # Dense weight already purged: the converter recorded the logical + # shape separately, since 2 packed codes/byte lose it. + weight_shape = [int(dim) for dim in mod._w4a16_weight_shape] + else: + raise RuntimeError(f"W4A16 layer {fqn!r} has no dense weight and no recorded " + "_w4a16_weight_shape; its logical weight shape cannot be recovered from the " + "packed codes. Re-run convert_model_to_w4a16(model) before saving.") + tensors = { + name: [int(dim) for dim in getattr(mod, name).shape] + for name in _W4A16_SIDECAR_BUFFERS if getattr(mod, name, None) is not None + } + layers[fqn] = { + "weight_shape": weight_shape, + "group_size": int(qm.group_size), + "bits": int(qm.bits), + "tensors": tensors, + } + quant_prefixes[fqn] = getattr(qm, "layer_prefix", "") or "" + dense_bytes += weight_shape[0] * weight_shape[1] * 2 + group_sizes.add(int(qm.group_size)) + bit_widths.add(int(qm.bits)) + + metadata: dict[str, Any] = { + "format": _W4A16_SIDECAR_FORMAT, + "version": _W4A16_SIDECAR_VERSION, + # Uniform scheme when every layer agrees (the normal case); None when a + # model mixes schemes, in which case the per-layer entries are + # authoritative. A loader validates the per-layer values. + "group_size": group_sizes.pop() if len(group_sizes) == 1 else None, + "bits": bit_widths.pop() if len(bit_widths) == 1 else None, + "num_layers": len(layers), + "layers": layers, + "quant_prefixes": quant_prefixes, + "model_class": type(model).__name__, + } + if extra_metadata: + metadata.update(extra_metadata) + + payload = dict(state) + serialized_bytes = sum(t.numel() * t.element_size() for t in payload.values()) + save_file(payload, os.fspath(path), metadata={_W4A16_SIDECAR_METADATA_KEY: json.dumps(metadata)}) + + receipt = { + "path": os.fspath(path), + "num_layers": len(layers), + "num_tensors": len(payload), + "quantized_bytes": serialized_bytes, + "dense_bfloat16_bytes": dense_bytes, + "compression_ratio": (dense_bytes / serialized_bytes) if serialized_bytes else 0.0, + } + logger.info( + "W4A16 sidecar: wrote %d quantized modules / %d tensors (%d bytes) to %s " + "(%.2f GiB quantized vs %.2f GiB dense bf16, %.2fx smaller).", + receipt["num_layers"], + receipt["num_tensors"], + serialized_bytes, + receipt["path"], + serialized_bytes / (1 << 30), + dense_bytes / (1 << 30), + receipt["compression_ratio"], + ) + return receipt + + +def read_w4a16_sidecar_metadata(path: str | os.PathLike[str]) -> dict[str, Any]: + """Return the manifest of a sidecar file without materializing its tensors.""" + from safetensors import safe_open + + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + raw = handle.metadata() or {} + if _W4A16_SIDECAR_METADATA_KEY not in raw: + raise ValueError(f"{os.fspath(path)} is not a FastVideo W4A16 sidecar " + f"(no {_W4A16_SIDECAR_METADATA_KEY!r} metadata).") + return json.loads(raw[_W4A16_SIDECAR_METADATA_KEY]) + + +def w4a16_sidecar_path_for(checkpoint_path: str | os.PathLike[str]) -> str: + """Conventional sidecar path for a transformer checkpoint or directory. + + ``.../transformer.safetensors`` -> ``.../transformer.w4a16.safetensors``; + a directory -> ``/w4a16.safetensors``. + """ + raw = os.fspath(checkpoint_path) + if os.path.isdir(raw): + return os.path.join(raw, W4A16_DIR_SIDECAR_NAME) + if raw.endswith(".safetensors"): + return raw[:-len(".safetensors")] + W4A16_SIDECAR_SUFFIX + return raw + W4A16_SIDECAR_SUFFIX + + +def _sidecar_target_device(mod: torch.nn.Module, name: str) -> torch.device | None: + """Device the restored buffer should live on. + + Mirrors ``_quantize_layer_weight``, which registers the buffers on the + (local) weight's device; falls back to an existing buffer, then to the + module's parameter device so a purge-then-restore still lands on GPU. + """ + weight = getattr(mod, "weight", None) + if weight is not None and weight.device.type != "meta": + return weight.device + existing = getattr(mod, name, None) + if existing is not None and existing.device.type != "meta": + return existing.device + for param in mod.parameters(recurse=False): + if param.device.type != "meta": + return param.device + return None + + +def _expected_sidecar_shapes(name: str, weight_shape: tuple[int, int], group_size: int, + bits: int) -> tuple[tuple[int, ...], ...]: + """The one shape a fresh conversion would produce for *name*. + + Unlike the NVFP4 sidecar there is no padded variant to accept: the + quantizer groups along the last axis and ``_group`` refuses a K that is + not divisible by ``group_size``, and the packed code interpretation is + fixed by ``bits``. A sidecar that disagrees would not merely be unusual — + reading it back would unpack the wrong nibble order or the wrong number of + groups, which is silent corruption rather than an error. + """ + out_dim, in_dim = weight_shape + if name == "_w4a16_codes": + if bits == 8: + return ((out_dim, in_dim), ) + if in_dim % 2: + raise ValueError(f"Sidecar declares a ({out_dim}, {in_dim}) weight at bits={bits}, but 4-bit codes " + "pack two per byte and need an even input dim.") + return ((out_dim, in_dim // 2), ) + if in_dim % group_size: + raise ValueError(f"Sidecar declares a ({out_dim}, {in_dim}) weight with group_size {group_size}, " + "which does not divide the input dim; this layout cannot be dequantized.") + return ((out_dim, in_dim // group_size), ) + + +def load_w4a16_checkpoint( + model: torch.nn.Module, + path: str | os.PathLike[str], + *, + strict: bool = True, +) -> int: + """Restore W4A16 tensors from a sidecar, skipping ``convert_model_to_w4a16``. + + Registers ``_w4a16_codes`` / ``_w4a16_scales`` / ``_w4a16_zeros`` on every + W4A16-tagged linear from the sidecar, byte-for-byte as the conversion would + have produced them, and restores the ``_w4a16_weight_shape`` attribute that + :meth:`W4A16QuantizeMethod.apply` needs (it is not a buffer, so nothing else + would). The dense bf16 weights are never touched (they may be absent + entirely), and nothing here needs a GPU or a 4-bit kernel — the + dequantize-then-GEMM reference path in ``apply`` is pure PyTorch, so a + pre-quantized checkpoint loads on any host. + + ``strict`` raises on any layer-set mismatch (a sidecar that does not + describe this model); with ``strict=False`` those are logged and skipped, + leaving those layers unconverted. Scheme (``group_size``/``bits``) and + per-tensor shape/dtype mismatches are **never** downgraded: mis-read or + mis-unpacked codes produce garbage output with no error, so those always + raise. + + Returns the number of layers restored. + """ + from safetensors import safe_open + + tagged = _w4a16_tagged_modules(model) + if not tagged: + raise RuntimeError("No W4A16 linear layers are attached to this model, so a sidecar cannot be restored. " + "This is the silent-dense failure mode: the model's W4A16Config layer selection " + "does not cover its layer paths (for MiniMax-H3 use " + "W4A16Config.for_minimax_h3()).") + + manifest = read_w4a16_sidecar_metadata(path) + if manifest.get("format") != _W4A16_SIDECAR_FORMAT: + raise ValueError(f"Unsupported W4A16 sidecar format {manifest.get('format')!r} in {os.fspath(path)}.") + if int(manifest.get("version", -1)) != _W4A16_SIDECAR_VERSION: + raise ValueError(f"Unsupported W4A16 sidecar version {manifest.get('version')!r} in {os.fspath(path)} " + f"(this build reads version {_W4A16_SIDECAR_VERSION}).") + + saved_layers: dict[str, dict[str, Any]] = manifest.get("layers", {}) + model_fqns = {fqn for fqn, _, _ in tagged} + missing = sorted(model_fqns - set(saved_layers)) + extra = sorted(set(saved_layers) - model_fqns) + if missing or extra: + message = (f"W4A16 sidecar {os.fspath(path)} does not match this model: {len(missing)} layers missing " + f"from the sidecar, {len(extra)} layers not in the model. First missing={missing[:3]}, " + f"first extra={extra[:3]}.") + if strict: + raise ValueError(message) + logger.warning("%s Restoring the intersection only.", message) + + restored = 0 + with safe_open(os.fspath(path), framework="pt", device="cpu") as handle: + available = set(handle.keys()) + for fqn, mod, qm in tagged: + if fqn not in saved_layers: + continue + entry = saved_layers[fqn] + try: + weight_shape = tuple(int(dim) for dim in entry["weight_shape"]) + group_size = int(entry["group_size"]) + bits = int(entry["bits"]) + except (KeyError, TypeError, ValueError) as exc: + raise ValueError(f"W4A16 sidecar {os.fspath(path)} entry for {fqn!r} is malformed: {entry!r} " + "does not carry an integer weight_shape/group_size/bits.") from exc + if len(weight_shape) != 2: + raise ValueError(f"W4A16 sidecar entry {fqn!r} declares weight shape {list(weight_shape)}; " + "a linear weight is 2-D.") + # A different scheme means a different nibble layout and different + # dequantize arithmetic over the same bytes: never recoverable, so + # never relaxed by strict=False. + if group_size != qm.group_size or bits != qm.bits: + raise ValueError(f"W4A16 sidecar {os.fspath(path)} was written for {fqn!r} with " + f"group_size={group_size}, bits={bits}, but this model quantizes it with " + f"group_size={qm.group_size}, bits={qm.bits}.") + weight = getattr(mod, "weight", None) + if weight is not None and tuple(int(dim) for dim in weight.shape) != weight_shape: + raise ValueError(f"W4A16 sidecar entry {fqn!r} describes a {list(weight_shape)} weight, but " + f"this model's layer has shape {list(weight.shape)}.") + tensors: dict[str, torch.Tensor] = {} + for name in _W4A16_SIDECAR_BUFFERS: + key = _sidecar_key(fqn, name) + if key not in available: + continue + tensor = handle.get_tensor(key) + expected_dtype = _W4A16_SIDECAR_DTYPES[name] + if tensor.dtype != expected_dtype: + raise ValueError(f"W4A16 sidecar tensor {key} has dtype {tensor.dtype}, expected " + f"{expected_dtype}. Codes are uint8 bit patterns; a cast would silently " + "unpack to different 4-bit values.") + expected = _expected_sidecar_shapes(name, weight_shape, group_size, bits) + if tuple(tensor.shape) not in expected: + raise ValueError(f"W4A16 sidecar tensor {key} has shape {tuple(tensor.shape)}, expected " + f"one of {list(expected)} for a {list(weight_shape)} linear with " + f"group_size={group_size}, bits={bits}.") + device = _sidecar_target_device(mod, name) + if device is not None: + tensor = tensor.to(device=device, non_blocking=True) + tensors[name] = tensor + if set(tensors) != set(_W4A16_SIDECAR_BUFFERS): + message = (f"W4A16 sidecar entry for {fqn!r} is incomplete (has {sorted(tensors)}); " + f"all of {list(_W4A16_SIDECAR_BUFFERS)} are required.") + if strict: + raise ValueError(message) + logger.warning("%s Skipping this layer.", message) + continue + for name, tensor in tensors.items(): + mod.register_buffer(name, tensor, persistent=False) + # Not a buffer: without it ``apply`` raises AttributeError on the + # first forward, because 2 codes/byte do not encode the logical + # weight shape. + mod._w4a16_weight_shape = weight_shape + restored += 1 + + logger.info("W4A16 sidecar: restored %d quantized modules from %s (dense weights untouched).", restored, + os.fspath(path)) + return restored + + +__all__ = [ + "DEFAULT_BITS", + "DEFAULT_GROUP_SIZE", + "MINIMAX_H3_BLOCK_LINEAR_SUFFIXES", + "MINIMAX_H3_BLOCK_SCOPES", + "MINIMAX_H3_NUM_LAYERS", + "MINIMAX_H3_NUM_REFINER_LAYERS", + "MINIMAX_H3_PREFIX", + "W4A16Config", + "W4A16QuantizeMethod", + "W4A16_DIR_SIDECAR_NAME", + "W4A16_SIDECAR_SUFFIX", + "convert_model_to_w4a16", + "load_w4a16_checkpoint", + "minimax_h3_w4a16_prefixes", + "read_w4a16_sidecar_metadata", + "save_w4a16_checkpoint", + "w4a16_dequantize", + "w4a16_quantize", + "w4a16_sidecar_path_for", + "w4a16_sidecar_state_dict", +] diff --git a/fastvideo/models/dits/minimax_h3.py b/fastvideo/models/dits/minimax_h3.py index 370d11d47b..0815b24466 100644 --- a/fastvideo/models/dits/minimax_h3.py +++ b/fastvideo/models/dits/minimax_h3.py @@ -229,6 +229,11 @@ def _gate_active(self) -> bool: reach a zero gate for it to ever train). """ if torch.is_grad_enabled(): + # Training can turn the zero-init gate nonzero; drop the cached + # answer so the next no-grad forward (validation sampling) + # re-tests the weight instead of skipping a branch that has + # started contributing. + self._gate_compress_active = None return True if self._gate_compress_active is None: if torch.compiler.is_compiling(): @@ -481,7 +486,6 @@ def __init__( fuse_modulate: bool = False, fuse_qknorm_rope: bool = False, fuse_swiglu: bool = False, - fa4_packed_varlen: bool = False, ) -> None: super().__init__() self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps) @@ -494,7 +498,10 @@ def __init__( quant_config, prefix=f"{prefix}.attn", fuse_qknorm_rope=fuse_qknorm_rope, - fa4_packed_varlen=fa4_packed_varlen, + # The packed multimodal document is a single long self-attention + # sequence. FA4's varlen API is substantially faster for this + # shape; the backend keeps grad/training and non-FA4 calls fixed. + fa4_packed_varlen=True, ) self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps) self.ff = MiniMaxH3FeedForward( @@ -592,16 +599,12 @@ class MiniMaxH3Transformer3DModel(BaseDiT): def _get_parameter_dtype(self, name: str, default_dtype: torch.dtype) -> torch.dtype: """Keep the released input, timestep, and output projections in FP32. - Factorized AdaLN uses FP16; BF16 is ~1.7x worse there. + Folded AdaLN parameters follow the enclosing FSDP policy. In + particular, an FP32 training load must keep them as FP32 optimizer + masters; pinning the folded weights to FP16 here quantizes every + small Adam update before the next forward. Release checkpoints are + still exported in BF16, matching the validated 42-block parent. """ - # Precedence: the factorized-AdaLN FP16 pin wins over - # uniform_parameter_dtype on purpose. Under FSDP's one-dtype rule the - # resulting mix hard-fails at load time, which beats silently training - # AdaLN in BF16. Rank-reduced checkpoints are inference artifacts -- - # train from the full-rank release. - if getattr(self, "adaln_rank", None) is not None and ( - ".adaln_proj." in name or name.startswith(("norm_out.linear.", "adaln_basis."))): - return torch.float16 if self.config.uniform_parameter_dtype: return default_dtype return torch.float32 if name.split(".", 1)[0] in self._keep_in_fp32_modules else default_dtype @@ -666,14 +669,6 @@ def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None: prefix=f"{config.prefix}.time_embedder", ) self.adaln_rank: int | None = arch.adaln_rank - if self.adaln_rank is not None and config.uniform_parameter_dtype: - raise ValueError( - "Rank-reduced AdaLN checkpoints (adaln_rank set) cannot be trained: " - "uniform_parameter_dtype needs one dtype for every trainable " - "parameter, but factorized AdaLN weights are pinned to FP16 " - "(BF16 reconstructs them ~1.7x worse). Fine-tune the full-rank " - "checkpoint instead, then re-fit the basis with " - "scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py.") adaln_dim = self.adaln_rank or arch.time_embed_dim self.adaln_basis = ReplicatedLinear( arch.time_embed_dim, @@ -720,7 +715,6 @@ def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None: fuse_modulate="modulate" in self.enabled_fusions, fuse_qknorm_rope="qknorm_rope" in self.enabled_fusions, fuse_swiglu="swiglu" in self.enabled_fusions, - fa4_packed_varlen=envs.FASTVIDEO_MINIMAX_H3_FA4_PACKED_VARLEN, ) for index in range(arch.num_layers) ]) self.norm_out = MiniMaxH3AdaLayerNormOut( diff --git a/fastvideo/models/loader/component_loader.py b/fastvideo/models/loader/component_loader.py index 148a6d1bb0..c70ae77436 100644 --- a/fastvideo/models/loader/component_loader.py +++ b/fastvideo/models/loader/component_loader.py @@ -45,12 +45,34 @@ safetensors_weights_iterator, ) from fastvideo.models.registry import ModelRegistry +from fastvideo.platforms import AttentionBackendEnum from fastvideo.utils import PRECISION_TO_TYPE, is_pin_memory_available from fastvideo.hooks.layerwise_offload import enable_layerwise_offload logger = init_logger(__name__) +def _teacher_critic_attention_context(fastvideo_args: FastVideoArgs): + """Mask only a generator-only QAT request for teacher/critic loads. + + ``_loading_teacher_critic_model`` predates role-local attention backends and + also controls custom-weight/quant-config handling. Treating the flag as a + blanket request for automatic attention used to erase an explicit + ``FLASH_ATTN`` request from DMD teacher and critic roles. Preserve every + explicit dense request; only suppress ``ATTN_QAT_TRAIN``, which is the + generator-only policy the flag was introduced to isolate. + """ + if not hasattr(fastvideo_args, "_loading_teacher_critic_model"): + return nullcontext() + + active_scope = _active_component_attention_backend_scope() + requested = (active_scope.backend if active_scope is not None else + coerce_attn_backend(getattr(fastvideo_args, "attention_backend", None))) + if requested is AttentionBackendEnum.ATTN_QAT_TRAIN: + return _component_attention_backend_scope(None, component="transformer") + return nullcontext() + + class ComponentLoader(ABC): """Base class for loading a specific type of model component.""" @@ -1052,14 +1074,11 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): # Generator-only QAT for DMD distillation: the teacher (real_score) and # critic (fake_score) transformers load with this flag set and must stay - # full precision. Drop the nvfp4_qat quant from their copied config, and - # build their attention under a scope that ignores any process-wide - # ATTN_QAT_TRAIN request so it falls back to dense. The generator loads - # without the flag and keeps both. The scope is exception-safe and - # needs no env mutation or selector cache flush (the request is part - # of the resolution cache key). - _qat_generator_only = hasattr(fastvideo_args, "_loading_teacher_critic_model") - if _qat_generator_only: + # full precision. Drop nvfp4_qat from their copied config. Attention is + # narrowed separately below only when the inherited request is + # ATTN_QAT_TRAIN; an explicit role-local FLASH_ATTN/SDPA request wins. + _teacher_or_critic = hasattr(fastvideo_args, "_loading_teacher_critic_model") + if _teacher_or_critic: dit_config.quant_config = None model_cls, _ = ModelRegistry.resolve_model_cls(cls_name) @@ -1107,8 +1126,7 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): # non-strictly for Cosmos2.5 only; keep upstream strict behavior for others. strict_load = not (cls_name.startswith("Cosmos25") or cls_name == "Cosmos25Transformer3DModel" or getattr(fastvideo_args.pipeline_config, "prefix", "") == "Cosmos25") - attention_context = (_component_attention_backend_scope(None, component="transformer") - if _qat_generator_only else nullcontext()) + attention_context = _teacher_critic_attention_context(fastvideo_args) with attention_context: # dit_config is what the model is handed and keeps as `self.config`, # so recording here makes the decision readable from the loaded @@ -1142,6 +1160,8 @@ def load(self, model_path: str, fastvideo_args: FastVideoArgs): training_mode=fastvideo_args.training_mode, enable_torch_compile=fastvideo_args.enable_torch_compile, torch_compile_kwargs=fastvideo_args.torch_compile_kwargs, + regional_compile=getattr(fastvideo_args, "regional_compile", False), + pre_fsdp_transform=getattr(fastvideo_args, "_pre_fsdp_transform", None), inference_regional_compile=fastvideo_args.inference_torch_compile, inference_vsa_tile_size=fastvideo_args.VSA_tile_size, # Only the whole-parameter half of the adapter is applied here, while diff --git a/fastvideo/models/loader/fsdp_load.py b/fastvideo/models/loader/fsdp_load.py index e076359ee6..260841afe1 100644 --- a/fastvideo/models/loader/fsdp_load.py +++ b/fastvideo/models/loader/fsdp_load.py @@ -16,11 +16,17 @@ from torch import nn from torch.distributed import DeviceMesh, init_device_mesh from torch.distributed._tensor import distribute_tensor +from torch.distributed.tensor import DTensor, Replicate, Shard from torch.distributed.fsdp import (CPUOffloadPolicy, FSDPModule, MixedPrecisionPolicy, fully_shard) from torch.nn.modules.module import _IncompatibleKeys from fastvideo.logger import init_logger from fastvideo.models.loader.lora_patch import DenseLoRAPatch +from fastvideo.models.loader.shard_cache import ( + shard_cache_context, + try_load_from_shard_cache, + write_shard_cache, +) from fastvideo.models.loader.utils import (get_param_names_mapping, hf_to_custom_state_dict) from fastvideo.models.loader.weight_utils import safetensors_weights_iterator from fastvideo.utils import set_mixed_precision_policy, is_pin_memory_available @@ -28,6 +34,86 @@ logger = init_logger(__name__) +def _dtensor_from_cpu_full_tensor( + full_tensor: torch.Tensor, + meta_sharded_param: DTensor, + device: torch.device, + target_dtype: torch.dtype, +) -> DTensor | None: + """Materialize only this rank's DTensor shard on CUDA. + + The checkpoint iterator already yields CPU tensors. Moving the complete + tensor to every GPU before ``distribute_tensor`` creates an avoidable + full-tensor H2D peak (about 1 GiB for H3's largest matrices), which can + OOM a multi-role DMD2 process even though the final local shards fit. + Replicate+Shard meshes can be sliced exactly on CPU and reconstructed as + a DTensor from the local piece. Unknown placements retain the established + full-tensor fallback. + """ + coordinate = meta_sharded_param.device_mesh.get_coordinate() + if coordinate is None: + return None + local = full_tensor + for mesh_dim, placement in enumerate(meta_sharded_param.placements): + if isinstance(placement, Replicate): + continue + if not isinstance(placement, Shard): + return None + chunks = torch.tensor_split( + local, + int(meta_sharded_param.device_mesh.size(mesh_dim)), + dim=int(placement.dim), + ) + local = chunks[int(coordinate[mesh_dim])] + local = local.contiguous().to(device=device, dtype=target_dtype) + return DTensor.from_local( + local, + meta_sharded_param.device_mesh, + meta_sharded_param.placements, + run_check=False, + shape=meta_sharded_param.shape, + stride=meta_sharded_param.stride(), + ) + + +def _mixed_precision_module_groups( + model: nn.Module, + default_param_dtype: torch.dtype | None, +) -> tuple[list[tuple[str, nn.Module]], set[nn.Parameter]]: + """Resolve declared FP32 FSDP groups and uncovered mixed parameters.""" + dtype_selector = getattr(model, "_get_parameter_dtype", None) + if not callable(dtype_selector) or default_param_dtype is None: + return [], set() + + # All selector lookups and declared-group matches use canonical + # (checkpoint-key) names: activation-checkpoint wrapping inserts + # ``_checkpoint_wrapped_module`` segments into named_parameters()/ + # named_modules() FQNs while dtype selectors are written against the + # clean architecture names. + mixed_params = { + parameter + for name, parameter in model.named_parameters() + if dtype_selector(_strip_checkpoint_wrapper_prefix(name), default_param_dtype) != default_param_dtype + } + declared = set(getattr(model, "_keep_in_fp32_modules", ())) + groups: list[tuple[str, nn.Module]] = [] + covered: set[nn.Parameter] = set() + for name, module in model.named_modules(): + clean_name = _strip_checkpoint_wrapper_prefix(name) + if clean_name not in declared or any(clean_name.startswith(f"{parent}.") for parent, _ in groups): + continue + parameters = list(module.named_parameters()) + if not parameters: + continue + if not all( + dtype_selector(_strip_checkpoint_wrapper_prefix(f"{clean_name}.{child_name}"), + default_param_dtype) == torch.float32 for child_name, _ in parameters): + continue + groups.append((clean_name, module)) + covered.update(parameter for _, parameter in parameters) + return groups, mixed_params - covered + + def _summarize_param_names(names: set[str]) -> str: """Collapse per-layer parameter names into one ``blocks.*.suffix xN`` entry each.""" families: dict[str, int] = {} @@ -70,6 +156,14 @@ def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lor FP8QuantizeMethod, convert_model_to_fp8, ) + from fastvideo.layers.quantization.int8_affine_config import ( + INT8AffineQuantizeMethod, + convert_model_to_int8_affine, + ) + from fastvideo.layers.quantization.w4a16_config import ( + W4A16QuantizeMethod, + convert_model_to_w4a16, + ) from fastvideo.layers.quantization.mxfp8_config import ( MXFP8QuantizeMethod, convert_model_to_mxfp8, @@ -94,6 +188,14 @@ def _maybe_quantize_model(model: nn.Module, *, defer_weight_conversion_until_lor logger.info("Converting loaded model weights for FP8 linear layers") convert_model_to_fp8(model) return + if isinstance(qm, INT8AffineQuantizeMethod): + logger.info("Converting loaded model weights for INT8 affine linear layers") + convert_model_to_int8_affine(model) + return + if isinstance(qm, W4A16QuantizeMethod): + logger.info("Converting loaded model weights for W4A16 linear layers") + convert_model_to_w4a16(model) + return if isinstance(qm, MXFP8QuantizeMethod): if defer_weight_conversion_until_lora_merge: logger.info("Deferring MXFP8 weight conversion until the inference LoRA merge completes") @@ -197,6 +299,8 @@ def maybe_load_fsdp_model( pin_cpu_memory: bool = True, enable_torch_compile: bool = False, torch_compile_kwargs: dict[str, Any] | None = None, + pre_fsdp_transform: Callable[[nn.Module], nn.Module] | None = None, + regional_compile: bool = False, inference_regional_compile: bool = False, inference_vsa_tile_size: int | None = None, lora_path: str | None = None, @@ -229,12 +333,24 @@ def maybe_load_fsdp_model( with set_default_dtype(default_dtype), torch.device("meta"): model = model_cls(**init_params) + if pre_fsdp_transform is not None: + model = pre_fsdp_transform(model) + dtype_selector = getattr(model, "_get_parameter_dtype", None) - has_mixed_parameter_dtypes = callable(dtype_selector) and any( - dtype_selector(name, param_dtype) != param_dtype for name, _ in model.named_parameters()) - if training_mode and has_mixed_parameter_dtypes: + parameter_dtype_overrides = [] + if callable(dtype_selector): + # Canonicalize activation-checkpoint-wrapped names so the selector + # (and the shard-cache manifest) always sees checkpoint-key names. + parameter_dtype_overrides = [ + (clean_name, str(selected_dtype)) for name, _ in model.named_parameters() + if (selected_dtype := dtype_selector(clean_name := _strip_checkpoint_wrapper_prefix(name), + default_dtype)) != default_dtype + ] + _, ungrouped_mixed_params = _mixed_precision_module_groups(model, param_dtype) + if training_mode and ungrouped_mixed_params: raise NotImplementedError("FSDP training with model-selected mixed parameter dtypes requires " - "separate gradient synchronization for replicated parameters.") + "separate gradient synchronization for replicated parameters or " + "declared FP32 module groups.") # Check if we should use FSDP use_fsdp = training_mode or fsdp_inference @@ -275,11 +391,6 @@ def maybe_load_fsdp_model( fsdp_shard_conditions=model._fsdp_shard_conditions, pin_cpu_memory=pin_cpu_memory) - # Host offload is already disabled on unified memory (GB10). Staging the - # 35B FastH3 DiT on CPU and then copying to CUDA doubled that working set - # and took minutes. Follow cpu_offload: read onto the accelerator. - weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=cpu_offload) - logger.info("Loading transformer weights with to_cpu=%s", cpu_offload) param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping) dense_lora_patch = DenseLoRAPatch.from_adapter( lora_path, @@ -287,9 +398,6 @@ def maybe_load_fsdp_model( 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 - # replacements and continue with a silently incomplete model. model_parameter_names = {name for name, _ in model.named_parameters()} missing_vsa_gates = sorted(name for name in dense_lora_patch.replacement_parameters if "gate_compress" in name and name not in model_parameter_names) @@ -298,16 +406,44 @@ def maybe_load_fsdp_model( "This LoRA adapter provides MiniMax H3 VSA compression gates, but the selected attention backend " "did not construct them. Use attention_backend='VIDEO_SPARSE_ATTN_H3'. Missing parameters: " + ", ".join(missing_vsa_gates[:3]) + (" ..." if len(missing_vsa_gates) > 3 else "")) - load_model_from_full_model_state_dict( - model, - weight_iterator, - device, - default_dtype, - strict=strict, - cpu_offload=cpu_offload, - param_names_mapping=param_names_mapping_fn, - dense_lora_patch=dense_lora_patch, - ) + + # Sharded base-weight cache (opt-in via FASTVIDEO_WEIGHT_SHARD_CACHE): + # rebuild local DTensor chunks from tmpfs instead of re-reading and + # re-scattering the full checkpoint on every relaunch. Any miss or + # validation failure falls through to the full load below. + shard_cache_ctx = None + # Adapter-applied weights must not read or populate a cache keyed only by + # the dense base checkpoint. + if use_fsdp and not cpu_offload and dense_lora_patch is None: + shard_cache_ctx = shard_cache_context( + weight_dir_list=weight_dir_list, + device_mesh=device_mesh, + hsdp_replicate_dim=hsdp_replicate_dim, + hsdp_shard_dim=hsdp_shard_dim, + default_dtype=default_dtype, + param_dtype=param_dtype, + param_names_mapping=model.param_names_mapping, + parameter_dtype_overrides=parameter_dtype_overrides, + ) + cache_hit = (shard_cache_ctx is not None + and try_load_from_shard_cache(model, shard_cache_ctx, device, strict=strict)) + if not cache_hit: + # Host offload is already disabled on unified-memory systems. Follow + # that policy instead of staging a second full copy on CPU. + weight_iterator = safetensors_weights_iterator(weight_dir_list, to_cpu=cpu_offload) + logger.info("Loading transformer weights with to_cpu=%s", cpu_offload) + load_model_from_full_model_state_dict( + model, + weight_iterator, + device, + default_dtype, + strict=strict, + cpu_offload=cpu_offload, + param_names_mapping=param_names_mapping_fn, + dense_lora_patch=dense_lora_patch, + ) + if shard_cache_ctx is not None: + write_shard_cache(model, shard_cache_ctx) if hasattr(model, "materialize_non_persistent_buffers"): model.materialize_non_persistent_buffers(device=device, dtype=default_dtype) for n, p in chain(model.named_parameters(), model.named_buffers()): @@ -325,16 +461,23 @@ def maybe_load_fsdp_model( # are present (lazy imports inside the helper). _maybe_quantize_model(model, defer_weight_conversion_until_lora_merge=lora_path is not None) - compile_in_loader = enable_torch_compile and training_mode - if compile_in_loader: - unsupported = _prepare_model_for_compile(model, regional=False) - if unsupported is not None: - logger.warning("Training torch.compile requested but disabled: %s. Model stays eager.", unsupported) + if enable_torch_compile and training_mode: + if not regional_compile: + unsupported = _prepare_model_for_compile(model, regional=False) + if unsupported is not None: + logger.warning("Training torch.compile requested but disabled: %s. Model stays eager.", unsupported) + else: + compile_kwargs = torch_compile_kwargs or {} + logger.info("Enabling whole-model torch.compile with kwargs=%s", compile_kwargs) + model = torch.compile(model, **compile_kwargs) else: - compile_kwargs = torch_compile_kwargs or {} - logger.info("Enabling torch.compile for FSDP training module with kwargs=%s", compile_kwargs) - model = torch.compile(model, **compile_kwargs) - logger.info("torch.compile enabled for %s", type(model).__name__) + unsupported = _regional_compile_unsupported_reason(init_params, training_mode=True) + if unsupported is not None: + logger.warning( + "enable_torch_compile requested but disabled: %s. " + "Training continues in eager mode.", unsupported) + else: + _compile_model_regions(model, torch_compile_kwargs or {}) elif inference_regional_compile and not training_mode: # Inference-side counterpart of the #1718 training regional compile: # per-block fullgraph compile right after the transformer loads, no @@ -342,6 +485,7 @@ def maybe_load_fsdp_model( unsupported = _regional_compile_unsupported_reason( init_params, vsa_tile_size=inference_vsa_tile_size, + training_mode=False, ) if unsupported is None: unsupported = _prepare_model_for_compile(model, regional=True) @@ -356,41 +500,44 @@ def maybe_load_fsdp_model( return model +def _strip_checkpoint_wrapper_prefix(name: str) -> str: + """Canonicalize an FQN produced after activation-checkpoint wrapping.""" + return name.replace("._checkpoint_wrapped_module.", ".").removeprefix("_checkpoint_wrapped_module.") + + def _regional_compile_unsupported_reason( init_params: dict[str, Any], *, vsa_tile_size: int | None = None, + training_mode: bool = False, ) -> str | None: - """Return why regional fullgraph compile cannot run, or None if it can. - - Dense FA2, FA3, and FA4 inference all route through compile-visible - custom-op boundaries. FA3's raw autograd.Function carve-out applies only - to grad-enabled calls, outside this inference-only loader path. - - The legacy VSA backend remains outside the fullgraph support envelope. - MiniMax H3's VSA backend is supported only through the inference-only - sm_100a tile-64 route; its regional hook resolves loaded compression - gates and probes the kernel before block capture. - """ + """Return why regional fullgraph compile cannot run, or ``None``.""" try: - from fastvideo.attention.layer import _attention_compile_explicitly_disabled + from fastvideo.attention.layer import (_attention_compile_disabled, + _attention_compile_explicitly_disabled) except Exception: # pragma: no cover - attention stack not importable pass else: - if _attention_compile_explicitly_disabled(): + compile_disabled = (_attention_compile_disabled() + if training_mode else _attention_compile_explicitly_disabled()) + if compile_disabled: # The escape hatch wraps attention forwards in # torch.compiler.disable, which is a hard dynamo error inside a # fullgraph region ("Skip inlining `torch.compiler.disable()`d - # function"). Degrade to eager instead, matching the hatch's + # function" at the first training step — h3-compile-ab job 2610). + # Degrade the role to eager instead, matching the hatch's # debugging intent. return ("FASTVIDEO_DISABLE_ATTENTION_COMPILE=1 keeps attention " "forwards out of compiled graphs via torch.compiler." "disable, which fullgraph regional compile cannot trace; " - "this model stays eager") + "this role stays eager") config = init_params.get("config") resolved = getattr(config, "_resolved_attention_backend", None) resolved_name = getattr(resolved, "name", "") if resolved_name == "VIDEO_SPARSE_ATTN_H3": + if training_mode: + return ("VIDEO_SPARSE_ATTN_H3 training uses sparse kernels and collectives " + "that are not fullgraph-traceable; this role stays eager") if os.environ.get("FASTVIDEO_H3_VSA_PROBE"): return ("FASTVIDEO_H3_VSA_PROBE records tensors and files from the VSA-H3 attention body, which " "regional fullgraph compile cannot capture; this model stays eager") @@ -404,7 +551,18 @@ def _regional_compile_unsupported_reason( return (f"attention backend resolved to {resolved_name}, whose Triton " "kernels, sequence-parallel collectives, and sync metadata " "guard graph-break (incompatible with fullgraph regional " - "compile); this model stays eager") + "compile); this role stays eager") + if not training_mode or resolved is None or resolved_name != "FLASH_ATTN": + return None + try: + from fastvideo.attention.utils.flash_attn_default import fa_version + except Exception: # pragma: no cover - flash-attn stack not importable + return None + if fa_version == "3": + return ("attention backend resolved to FLASH_ATTN with flash-attn 3, " + "whose grad-enabled path graph-breaks (incompatible with " + "fullgraph regional compile); use FA2, FA4 (FASTVIDEO_FA4=1), " + "or TORCH_SDPA for compiled training") return None @@ -421,25 +579,25 @@ def _enable_regional_attention_compile(model: nn.Module) -> int: def _compile_model_regions(model: nn.Module, compile_kwargs: dict[str, Any]) -> int: - """Compile repeated mathematical regions of a loaded model. + """Compile repeated mathematical regions after FSDP setup. Only the selected module ``forward`` is replaced. This keeps activation - checkpoint wrappers structurally transparent while any module-level hooks - (FSDP pre/post, layerwise offload) execute outside the compiled region. + checkpoint wrappers structurally transparent while FSDP pre/post hooks + execute outside the compiled region. """ compile_conditions = getattr(model, "_compile_conditions", None) if not compile_conditions: raise ValueError(f"{type(model).__name__} does not declare _compile_conditions") if compile_kwargs.get("fullgraph", True) is not True: - raise ValueError("Regional compile requires fullgraph=True") + raise ValueError("Regional training compile requires fullgraph=True") if "mode" in compile_kwargs: # torch.compile forbids passing both `mode` and `options`, and # regional compile always injects options (emulate_precision_casts) - # to match the training-side regional-compile configuration. Fail here - # with an actionable message instead of letting torch raise a - # mode/options conflict about an `options` key the user never wrote. - raise ValueError("Regional compile sets inductor options " + # for bf16 numerics parity. Fail here with an actionable message + # instead of letting torch raise a mode/options conflict about an + # `options` key the user never wrote. + raise ValueError("Regional training compile sets inductor options " "(emulate_precision_casts) and cannot be combined " "with torch_compile_kwargs['mode']. Remove 'mode' or " "express its effect via torch_compile_kwargs['options'].") @@ -462,7 +620,7 @@ def _compile_model_regions(model: nn.Module, compile_kwargs: dict[str, Any]) -> if compiled_count == 0: raise ValueError(f"No submodules in {type(model).__name__} matched _compile_conditions") logger.info( - "Enabled regional torch.compile for %d submodules in %s with kwargs=%s", + "Enabled regional torch.compile for %d submodules in %s after FSDP setup with kwargs=%s", compiled_count, type(model).__name__, kwargs, @@ -514,15 +672,12 @@ def shard_model( return default_param_dtype = getattr(mp_policy, "param_dtype", None) - dtype_selector = getattr(model, "_get_parameter_dtype", None) - ignored_params: set[nn.Parameter] = set() - if callable(dtype_selector) and default_param_dtype is not None: - ignored_params = { - parameter - for name, parameter in model.named_parameters() - if dtype_selector(name, default_param_dtype) != default_param_dtype - } + fp32_groups, ignored_params = _mixed_precision_module_groups(model, default_param_dtype) named_modules = list(model.named_modules()) + fp32_group_ids = {id(module) for _, module in fp32_groups} + fp32_group_params = { + parameter for _, module in fp32_groups for parameter in module.parameters() + } ignored_params_by_module = { id(module): ignored_params.intersection(set(module.parameters())) for _, module in named_modules @@ -547,6 +702,10 @@ def shard_model( for n, m in reversed(named_modules): if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]): + if id(m) in fp32_group_ids: + continue + if fp32_group_params.intersection(set(m.parameters())): + raise ValueError(f"FSDP shard condition for {n!r} contains a declared FP32 compute group") # Count all parameters param_count = sum(p.numel() for p in m.parameters(recurse=True)) @@ -568,6 +727,10 @@ def shard_model( # Shard all modules matching conditions for n, m in reversed(named_modules): if any([shard_condition(n, m) for shard_condition in fsdp_shard_conditions]): + if id(m) in fp32_group_ids: + continue + if fp32_group_params.intersection(set(m.parameters())): + raise ValueError(f"FSDP shard condition for {n!r} contains a declared FP32 compute group") module_kwargs = fsdp_kwargs local_ignored_params = ignored_params_by_module[id(m)] if local_ignored_params: @@ -578,6 +741,20 @@ def shard_model( if num_layers_sharded == 0: raise ValueError("No layer modules were sharded. Please check if shard conditions are working as expected.") + if fp32_groups: + fp32_kwargs = { + **fsdp_kwargs, + "mp_policy": MixedPrecisionPolicy( + param_dtype=torch.float32, + reduce_dtype=mp_policy.reduce_dtype, + output_dtype=mp_policy.output_dtype, + cast_forward_inputs=mp_policy.cast_forward_inputs, + ), + } + for _, module in fp32_groups: + fully_shard(module, **fp32_kwargs) + logger.info("Sharded FP32 compute modules: %s", [name for name, _ in fp32_groups]) + # Finally shard the entire model to account for any stragglers root_kwargs = fsdp_kwargs if ignored_params: @@ -585,6 +762,24 @@ def shard_model( fully_shard(model, **root_kwargs) +# Parameters the model registers at build time that a checkpoint never carries, +# so their absence from the incoming state dict is expected rather than a mapping +# bug. Quantization configs are the common source: a quant linear method creates +# its own scale tensors while the checkpoint only holds `weight`/`bias` +# (fastvideo/layers/quantization/absmax_fp8.py registers `scale_weight` and +# `scale_input`). Scales kept in `persistent=False` buffers — NVFP4, FP8, +# INT8Affine — never reach `state_dict()` and so need no entry here. +# `gate_compress` (VSA gate) and `proj_l` (SLA) are likewise built by the +# attention backend instead of loaded. Anything else in the model but missing +# from the checkpoint is a real mismatch and still raises below. +ALLOWED_NEW_PARAM_PATTERNS: tuple[str, ...] = ("gate_compress", "proj_l", "scale_weight", "scale_input") + + +def is_allowed_new_param(param_name: str) -> bool: + """Whether ``param_name`` may be zero-initialized instead of loaded.""" + return any(pattern in param_name for pattern in ALLOWED_NEW_PARAM_PATTERNS) + + # TODO(PY): device mesh for cfg parallel def load_model_from_full_model_state_dict( model: FSDPModule | torch.nn.Module, @@ -621,8 +816,14 @@ def load_model_from_full_model_state_dict( NotImplementedError: If got FSDP with more than 1D. """ meta_sd = model.state_dict() - named_parameters = dict(model.named_parameters()) - named_buffers = dict(model.named_buffers()) + # state_dict() keys are clean (checkpoint-wrapper hooks strip the AC + # prefix) but named_parameters()/named_buffers() are not; checkpoint keys + # are clean, so canonicalize before any name-keyed lookup. Without this, + # a loaded buffer inside an AC-wrapped block misses the named_buffers + # membership test below and is silently converted into a trainable + # nn.Parameter by load_state_dict(assign=True). + named_parameters = {_strip_checkpoint_wrapper_prefix(k): v for k, v in model.named_parameters()} + named_buffers = {_strip_checkpoint_wrapper_prefix(k): v for k, v in model.named_buffers()} sharded_sd = {} custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(full_sd_iterator, param_names_mapping) # type: ignore @@ -686,12 +887,25 @@ def load_model_from_full_model_state_dict( # In cases where parts of the model aren't sharded, some parameters will be plain tensors. sharded_tensor = full_tensor else: - full_tensor = full_tensor.to(device=device, dtype=target_dtype) - sharded_tensor = distribute_tensor( - full_tensor, - meta_sharded_param.device_mesh, - meta_sharded_param.placements, - ) + sharded_tensor = None + if full_tensor.device.type == "cpu" and isinstance(meta_sharded_param, DTensor): + sharded_tensor = _dtensor_from_cpu_full_tensor( + full_tensor, + meta_sharded_param, + device, + target_dtype, + ) + if sharded_tensor is None: + full_tensor = full_tensor.to(device=device, dtype=target_dtype) + # Every rank read the identical full tensor from the checkpoint, + # so each can slice its own shard locally; src_data_rank=None + # avoids a redundant rank-zero scatter. + sharded_tensor = distribute_tensor( + full_tensor, + meta_sharded_param.device_mesh, + meta_sharded_param.placements, + src_data_rank=None, + ) if cpu_offload: sharded_tensor = sharded_tensor.cpu() if target_param_name in named_buffers: @@ -716,18 +930,16 @@ def load_model_from_full_model_state_dict( logger.warning("Found unloaded parameters in meta state dict, zero-initializing: %d (%s)", len(zero_init), _summarize_param_names(zero_init)) - # List of allowed parameter name patterns - ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l"] # Can be extended as needed for new_param_name in unused_keys: # An adapter that ships the parameter outright both supplies the value and # authorizes it: the allowlist exists to catch a checkpoint silently missing a # weight, which is not the case when something deliberately provides one. adapter_value = (dense_lora_patch.replacement_for(new_param_name) if dense_lora_patch is not None else None) - if adapter_value is None and not any(pattern in new_param_name for pattern in ALLOWED_NEW_PARAM_PATTERNS): + if adapter_value is None and not is_allowed_new_param(new_param_name): logger.error("Unsupported new parameter: %s. Allowed patterns: %s", new_param_name, - ALLOWED_NEW_PARAM_PATTERNS) + list(ALLOWED_NEW_PARAM_PATTERNS)) raise ValueError(f"New parameter '{new_param_name}' is not supported. " - f"Currently only parameters containing {ALLOWED_NEW_PARAM_PATTERNS} are allowed.") + f"Currently only parameters containing {list(ALLOWED_NEW_PARAM_PATTERNS)} are allowed.") meta_sharded_param = meta_sd.get(new_param_name) target_dtype = param_dtype dtype_selector = getattr(model, "_get_parameter_dtype", None) diff --git a/fastvideo/models/loader/shard_cache.py b/fastvideo/models/loader/shard_cache.py new file mode 100644 index 0000000000..af820e16dd --- /dev/null +++ b/fastvideo/models/loader/shard_cache.py @@ -0,0 +1,373 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Per-rank sharded base-weight cache for fast training relaunches. + +After a full checkpoint load, each shard rank persists its local DTensor +chunks (post-rename, post-cast — exactly what ``assign=True`` installed) as +one safetensors file in a cache directory (typically tmpfs). Subsequent +launches with an identical (checkpoint, mesh layout, dtype, name-mapping) +tuple rebuild the model from those chunks via ``DTensor.from_local`` — +skipping the full-tensor reads, per-rank H2D of the whole checkpoint, and +the distribute/scatter step. + +Opt-in via ``FASTVIDEO_WEIGHT_SHARD_CACHE=`` (e.g. ``/dev/shm/fv-wcache``). +Any validation failure or exception degrades to the normal full load — the +cache can never fail a run. +""" + +from __future__ import annotations + +import hashlib +import json +import os +import shutil +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import torch +import torch.distributed as dist +import torch.nn as nn +from torch.distributed.tensor import DTensor + +from fastvideo.logger import init_logger + +logger = init_logger(__name__) + +_FORMAT_VERSION = 1 +_ENV_DIR = "FASTVIDEO_WEIGHT_SHARD_CACHE" +_ENV_MAX_GB = "FASTVIDEO_WEIGHT_SHARD_CACHE_MAX_GB" +# Set to "1" when the cache dir is node-local (tmpfs) on a multi-node run so +# every replicate rank writes its own node's copy. +_ENV_PER_NODE = "FASTVIDEO_WEIGHT_SHARD_CACHE_PER_NODE" +# Mirrors fsdp_load's zero-init allowance for params absent from checkpoints. +# Kept as a literal (not an import) because fsdp_load imports this module; a +# test asserts the two tuples stay identical so the mirror cannot drift. +_ALLOWED_NEW_PARAM_PATTERNS = ("gate_compress", "proj_l", "scale_weight", "scale_input") +_WRITE_MARGIN_BYTES = 5 << 30 + + +@dataclass +class ShardCacheContext: + entry_dir: Path + key: str + shard_index: int + num_shards: int + is_writer: bool # replicate-coordinate 0 writes; replicas are identical + + +def _expand_weight_files(weight_dir_list: list[str]) -> list[str]: + files: list[str] = [] + for entry in weight_dir_list: + if os.path.isdir(entry): + files.extend(str(p) for p in sorted(Path(entry).glob("*.safetensors"))) + else: + files.append(entry) + return files + + +def _shard_file(entry_dir: Path, shard_index: int, num_shards: int) -> Path: + return entry_dir / f"shard{shard_index}-of-{num_shards}.safetensors" + + +def shard_cache_context( + *, + weight_dir_list: list[str], + device_mesh: Any, + hsdp_replicate_dim: int, + hsdp_shard_dim: int, + default_dtype: torch.dtype, + param_dtype: torch.dtype, + param_names_mapping: dict[str, str] | None, + parameter_dtype_overrides: list[tuple[str, str]] | None = None, +) -> ShardCacheContext | None: + root = os.environ.get(_ENV_DIR) + if not root: + return None + try: + files = _expand_weight_files(weight_dir_list) + if not files: + return None + stats = sorted((os.path.basename(p), os.stat(p).st_size, os.stat(p).st_mtime_ns) for p in files) + mapping_items = sorted((param_names_mapping or {}).items()) + key_fields = [ + _FORMAT_VERSION, + stats, + int(hsdp_replicate_dim), + int(hsdp_shard_dim), + str(default_dtype), + str(param_dtype), + mapping_items, + ] + if parameter_dtype_overrides: + key_fields.append(sorted(parameter_dtype_overrides)) + key_material = json.dumps(key_fields, sort_keys=True) + key = hashlib.sha256(key_material.encode()).hexdigest()[:16] + coordinate = device_mesh.get_coordinate() + if coordinate is None: + return None + # mesh dims are ("replicate", "shard") + replicate_index, shard_index = int(coordinate[0]), int(coordinate[1]) + # With a shared cache dir, replicate-coordinate 0 writes and replicas + # (identical shards) skip to avoid same-file collisions. With a + # node-local dir (e.g. /dev/shm on a multi-node HSDP run), every + # replica must write its own node's copy — same-name collisions are + # impossible across nodes, and a single-writer rule would leave every + # non-zero replica's node permanently cold (all-reduce MIN then turns + # that into a global miss). + per_node_root = os.environ.get(_ENV_PER_NODE, "0") == "1" + return ShardCacheContext( + entry_dir=Path(root) / key, + key=key, + shard_index=shard_index, + num_shards=int(hsdp_shard_dim), + is_writer=per_node_root or replicate_index == 0, + ) + except Exception as exc: # noqa: BLE001 - cache must never fail a load + logger.warning("shard cache disabled for this load (context error): %s", exc) + return None + + +def _all_ranks_agree(local_ok: bool, device: torch.device) -> bool: + if not (dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1): + return local_ok + flag = torch.tensor([1 if local_ok else 0], device=device, dtype=torch.int32) + dist.all_reduce(flag, op=dist.ReduceOp.MIN) + return bool(flag.item()) + + +def _validate_entry( + entry: dict[str, Any], + meta_param: torch.Tensor, + expected_dtype: torch.dtype, +) -> bool: + if entry["dtype"] != str(expected_dtype): + return False + if list(entry["global_shape"]) != list(meta_param.shape): + return False + is_dtensor = isinstance(meta_param, DTensor) + if entry["kind"] != ("dtensor" if is_dtensor else "tensor"): + return False + if is_dtensor: + if entry["placements"] != [str(p) for p in meta_param.placements]: + return False + if list(entry["local_shape"]) != list(meta_param.to_local().shape): + return False + return True + + +def try_load_from_shard_cache( + model: nn.Module, + ctx: ShardCacheContext, + device: torch.device, + *, + strict: bool = True, +) -> bool: + """Assemble the model's state dict from cached local shards. Returns False + (without mutating the model) on any mismatch.""" + try: + manifest_path = ctx.entry_dir / "manifest.json" + shard_path = _shard_file(ctx.entry_dir, ctx.shard_index, ctx.num_shards) + local_ok = manifest_path.is_file() and shard_path.is_file() + manifest: dict[str, Any] = {} + if local_ok: + manifest = json.loads(manifest_path.read_text()) + local_ok = (manifest.get("format_version") == _FORMAT_VERSION and manifest.get("key") == ctx.key) + meta_sd = model.state_dict() + dtype_selector = getattr(model, "_get_parameter_dtype", None) + if local_ok: + params_table = manifest["params"] + for name, meta_param in meta_sd.items(): + entry = params_table.get(name) + if entry is None: + if not any(pattern in name for pattern in _ALLOWED_NEW_PARAM_PATTERNS): + local_ok = False + break + continue + expected_dtype = meta_param.dtype + if callable(dtype_selector): + expected_dtype = dtype_selector(name, expected_dtype) + if not _validate_entry(entry, meta_param, expected_dtype): + local_ok = False + break + if not _all_ranks_agree(local_ok, device): + if local_ok: + logger.info("shard cache: another rank missed entry %s; falling back to full load", ctx.key) + return False + + from safetensors import safe_open + + # meta_sd (and the manifest) use clean checkpoint keys, but a model + # activation-checkpoint-wrapped before load (pre-FSDP AC) yields + # `_checkpoint_wrapped_module.`-prefixed names from named_buffers(). + # Canonicalize like the full-load path, or the membership test below + # rebuilds a cached buffer as a trainable nn.Parameter on warm boots. + from fastvideo.models.loader.fsdp_load import ( + _strip_checkpoint_wrapper_prefix, ) + + named_buffers = {_strip_checkpoint_wrapper_prefix(k): v for k, v in model.named_buffers()} + sharded_sd: dict[str, Any] = {} + with safe_open(str(shard_path), framework="pt", device=str(device)) as f: + cached_keys = set(f.keys()) + for name, meta_param in meta_sd.items(): + if name in cached_keys: + local = f.get_tensor(name) + if isinstance(meta_param, DTensor): + tensor: torch.Tensor = DTensor.from_local( + local, + meta_param.device_mesh, + meta_param.placements, + run_check=False, + shape=meta_param.shape, + stride=meta_param.stride(), + ) + else: + tensor = local + else: + # Zero-init new params exactly like the full-load path. + target_dtype = meta_param.dtype + if callable(dtype_selector): + target_dtype = dtype_selector(name, target_dtype) + if isinstance(meta_param, DTensor): + local = torch.zeros(meta_param.to_local().shape, device=device, dtype=target_dtype) + tensor = DTensor.from_local( + local, + meta_param.device_mesh, + meta_param.placements, + run_check=False, + shape=meta_param.shape, + stride=meta_param.stride(), + ) + else: + tensor = torch.zeros(meta_param.shape, device=device, dtype=target_dtype) + sharded_sd[name] = tensor if name in named_buffers else nn.Parameter(tensor) + + reverse_map = manifest.get("reverse_param_names_mapping", {}) + model.reverse_param_names_mapping = {k: tuple(v) for k, v in reverse_map.items()} + model.load_state_dict(sharded_sd, strict=strict, assign=True) + # Freshen mtimes so mtime-based tmpfs cleaners (and our own LRU GC) + # treat actively used entries as recent. + for p in (shard_path, manifest_path): + try: + os.utime(p) + except OSError: + pass + logger.info( + "shard cache HIT %s: %d tensors from %s", + ctx.key, + len(sharded_sd), + shard_path, + ) + return True + except Exception as exc: # noqa: BLE001 - cache must never fail a load + logger.warning("shard cache load failed (%s); falling back to full load", exc) + return False + + +def write_shard_cache(model: nn.Module, ctx: ShardCacheContext) -> None: + """Persist this rank's local shards after a successful full load.""" + try: + from safetensors.torch import save_file + + tensors: dict[str, torch.Tensor] = {} + params_table: dict[str, Any] = {} + for name, value in model.state_dict().items(): + if isinstance(value, DTensor): + local = value.to_local().detach().to("cpu", copy=True).contiguous() + params_table[name] = { + "kind": "dtensor", + "dtype": str(value.dtype), + "global_shape": list(value.shape), + "placements": [str(p) for p in value.placements], + "local_shape": list(local.shape), + } + else: + local = value.detach().to("cpu", copy=True).contiguous() + params_table[name] = { + "kind": "tensor", + "dtype": str(value.dtype), + "global_shape": list(value.shape), + } + tensors[name] = local + + ctx.entry_dir.mkdir(parents=True, exist_ok=True) + needed = sum(t.numel() * t.element_size() for t in tensors.values()) + free = shutil.disk_usage(ctx.entry_dir).free + if free < needed + _WRITE_MARGIN_BYTES: + logger.warning( + "shard cache: skipping write (%.1f GiB needed, %.1f GiB free at %s)", + needed / 2**30, + free / 2**30, + ctx.entry_dir, + ) + _barrier_if_initialized() + return + + if ctx.is_writer: + shard_path = _shard_file(ctx.entry_dir, ctx.shard_index, ctx.num_shards) + tmp_path = shard_path.with_suffix(".safetensors.tmp") + save_file(tensors, str(tmp_path)) + os.replace(tmp_path, shard_path) + _barrier_if_initialized() + + rank = dist.get_rank() if (dist.is_available() and dist.is_initialized()) else 0 + if ctx.is_writer: + # Every shard writer emits the manifest, not just global rank 0: + # with a node-local cache root (PER_NODE=1) each node holds its own + # entry copy, and a manifest written only on rank 0's node leaves + # every other node's entry manifest-less — the all-rank agreement + # vote in try_load_from_shard_cache then fails on EVERY multi-node + # warm boot and silently degrades relaunches to full loads. The + # content is rank-invariant for uniformly divisible shards; the + # rank-suffixed tmp name keeps concurrent same-directory writers + # (shared root, or several local ranks per node) from clobbering + # each other's half-written file before the atomic replace. + reverse_map = { + k: list(v) + for k, v in getattr(model, "reverse_param_names_mapping", {}).items() + } + manifest = { + "format_version": _FORMAT_VERSION, + "key": ctx.key, + "num_shards": ctx.num_shards, + "params": params_table, + "reverse_param_names_mapping": reverse_map, + } + manifest_tmp = ctx.entry_dir / f"manifest.json.tmp.{ctx.shard_index}" + manifest_tmp.write_text(json.dumps(manifest)) + os.replace(manifest_tmp, ctx.entry_dir / "manifest.json") + if rank == 0: + logger.info("shard cache WRITE %s: %d tensors -> %s", ctx.key, len(tensors), ctx.entry_dir) + _gc_cache_root(ctx.entry_dir.parent, keep=ctx.entry_dir.name) + except Exception as exc: # noqa: BLE001 - cache must never fail a run + logger.warning("shard cache write failed (non-fatal): %s", exc) + _barrier_if_initialized() + + +def _barrier_if_initialized() -> None: + if dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1: + dist.barrier() + + +def _gc_cache_root(root: Path, keep: str) -> None: + """Drop least-recently-written entries beyond the size cap.""" + try: + max_bytes = float(os.environ.get(_ENV_MAX_GB, "300")) * 2**30 + entries = [] + for child in root.iterdir(): + if not child.is_dir(): + continue + manifest = child / "manifest.json" + size = sum(f.stat().st_size for f in child.glob("*") if f.is_file()) + mtime = manifest.stat().st_mtime if manifest.is_file() else 0.0 + entries.append((mtime, size, child)) + total = sum(size for _, size, _ in entries) + for mtime, size, child in sorted(entries): + if total <= max_bytes: + break + if child.name == keep: + continue + shutil.rmtree(child, ignore_errors=True) + total -= size + logger.info("shard cache GC: evicted %s (%.1f GiB)", child, size / 2**30) + except Exception as exc: # noqa: BLE001 + logger.warning("shard cache GC failed (non-fatal): %s", exc) diff --git a/fastvideo/models/schedulers/scheduling_minimax_h3.py b/fastvideo/models/schedulers/scheduling_minimax_h3.py index 647075caa2..f76c996cfc 100644 --- a/fastvideo/models/schedulers/scheduling_minimax_h3.py +++ b/fastvideo/models/schedulers/scheduling_minimax_h3.py @@ -51,6 +51,11 @@ def set_shift(self, shift: float) -> None: raise ValueError(f"`shift` must be positive, got {shift}.") self._shift = float(shift) + def shift_sigmas(self, base_sigmas: torch.Tensor) -> torch.Tensor: + """Warp unshifted base noise amounts onto this scheduler's shifted grid.""" + base = torch.as_tensor(base_sigmas, dtype=torch.float32) + return self._shift * base / (1 + (self._shift - 1) * base) + def set_timesteps( self, num_inference_steps: int | None = None, @@ -62,7 +67,7 @@ def set_timesteps( raise ValueError("`set_timesteps` requires explicit `sigmas` or " f"`num_inference_steps` >= 2, got {num_inference_steps}.") base = torch.linspace(1.0, 0.0, int(num_inference_steps), dtype=torch.float32) - sigma_tensor = self._shift * base / (1 + (self._shift - 1) * base) + sigma_tensor = self.shift_sigmas(base) sigma_tensor = torch.unique_consecutive(sigma_tensor) else: sigma_tensor = torch.as_tensor(sigmas, dtype=torch.float32).flatten().cpu() diff --git a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_decoding.py b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_decoding.py index 1eae7005e5..ce3bc2a0b8 100644 --- a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_decoding.py +++ b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_decoding.py @@ -7,7 +7,7 @@ import torch -from fastvideo.distributed import get_local_torch_device, get_sp_group, get_world_group, model_parallel_is_initialized +from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized from fastvideo.fastvideo_args import FastVideoArgs from fastvideo.logger import init_logger from fastvideo.models.vaes.minimax_h3_audio import MiniMaxH3AudioVAE @@ -40,9 +40,9 @@ def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout: def _decode_participation(fastvideo_args: FastVideoArgs, want_parallel: bool) -> tuple[Any, bool, bool]: """Resolve (sp_group, is_output_rank, parallel) for the VAE decode stages. - The existing serial path keeps its global-rank-zero output ownership. - Parallel decode assembles once per sequence-parallel group, on that - group's first rank. ``parallel`` is only true when every group rank will + Both paths produce one output per sequence-parallel group, on that + group's first rank. This retains every data-parallel validation sample + instead of silently keeping only global rank zero. ``parallel`` is only true when every group rank will run the decode body — the collectives inside require uniform participation, so no rank-dependent branch may guard them. """ @@ -51,7 +51,7 @@ def _decode_participation(fastvideo_args: FastVideoArgs, want_parallel: bool) -> sp_group = get_sp_group() if bool(want_parallel) and sp_group.world_size > 1: return sp_group, sp_group.is_first_rank, True - return sp_group, get_world_group().is_first_rank, False + return sp_group, sp_group.is_first_rank, False class MiniMaxH3VideoDecodingStage(PipelineStage): @@ -184,9 +184,9 @@ def verify_output(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> V @torch.no_grad() def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch: """Decode H3 audio latents into a stereo CPU waveform.""" - # Audio decode is sub-second, so preserve the serial path's global - # rank-zero ownership. - if model_parallel_is_initialized() and not get_world_group().is_first_rank: + # Decode once per sequence-parallel group so data-parallel validation + # retains one waveform for every generated sample. + if model_parallel_is_initialized() and not get_sp_group().is_first_rank: batch.extra["audio"] = torch.empty((0, 2), device="cpu", dtype=torch.float32) batch.extra["audio_sample_rate"] = self.audio_vae.sampling_rate self._clear_runtime(batch) diff --git a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py index bce4ea978e..44dfb3d713 100644 --- a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py +++ b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_denoising.py @@ -3,6 +3,8 @@ from __future__ import annotations +import os + from typing import Any import torch @@ -11,6 +13,7 @@ from fastvideo.distributed import get_local_torch_device from fastvideo.fastvideo_args import FastVideoArgs from fastvideo.forward_context import set_forward_context +from fastvideo.logger import init_logger from fastvideo.hooks.activation_trace import trace_step from fastvideo.profiler import nvtx_range, profiler_region from fastvideo.pipelines.basic.minimax_h3.packing import ( @@ -25,6 +28,8 @@ from fastvideo.pipelines.stages.validators import VerificationResult from fastvideo.utils import get_compute_dtype +logger = init_logger(__name__) + def _h3_vsa_metadata_builder(transformer: Any, fastvideo_args: FastVideoArgs) -> Any: """Builder instance when the transformer resolved to VSA-H3, else None. @@ -117,11 +122,22 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward device = get_local_torch_device() dmd_steps = fastvideo_args.pipeline_config.dmd_denoising_steps - if dmd_steps is None: + if not dmd_steps: + env_steps = os.environ.get("FASTVIDEO_DMD_DENOISING_STEPS", "").strip() + if env_steps: + dmd_steps = [int(s) for s in env_steps.split(",") if s.strip()] + # Optionally re-noise x0 between explicit DMD steps. + stochastic_renoise = bool(dmd_steps) and ( + bool(getattr(fastvideo_args.pipeline_config, "dmd_stochastic_renoise", False)) + or os.environ.get("FASTVIDEO_DMD_STOCHASTIC_RENOISE", "0").strip().lower() in ("1", "true", "yes")) + if stochastic_renoise: + logger.info("MiniMax-H3 DMD denoising with stochastic fresh-noise re-noising.") + if dmd_steps: + self._set_dmd_schedule(dmd_steps, batch.num_inference_steps, device) + logger.info("MiniMax-H3 denoising with explicit DMD steps %s.", list(dmd_steps)) + else: self.scheduler.set_timesteps(batch.num_inference_steps, device=device) self.audio_scheduler.set_timesteps(batch.num_inference_steps, device=device) - else: - self._set_dmd_schedule(dmd_steps, batch.num_inference_steps, device) video_timesteps = self.scheduler.timesteps audio_timesteps = self.audio_scheduler.timesteps if video_timesteps is None or audio_timesteps is None: @@ -162,9 +178,16 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward vsa_exempt = vsa_mode == "exempt" vsa_dense_layers = tuple(batch.extra.get("vsa_dense_layers", ())) vsa_dense_first_n = int(batch.extra.get("vsa_dense_first_n_steps", 0)) + # A per-request batch sparsity wins when set; otherwise fall back + # to the run-level args value (training validation and the CLI set + # fastvideo_args.VSA_sparsity — nothing populates the batch field + # on those paths, and the batch's 0.0 default silently sampled + # dense while training ran sparse). + vsa_sparsity_base = (float(batch.VSA_sparsity) + if float(batch.VSA_sparsity) > 0.0 else float(fastvideo_args.VSA_sparsity)) # Run-level tile geometry (256 default, 64 = native Triton path), - # plumbed like the run-level sparsity; the builder validates the - # value against VSA_H3_TILE_SHAPES. + # plumbed like the run-level sparsity above; the builder validates + # the value against VSA_H3_TILE_SHAPES. vsa_tile_size = int(fastvideo_args.VSA_tile_size) try: @@ -184,7 +207,7 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward # Optional schedule: run the first N steps dense (sparsity 0 # selects every tile — parity-proven ≡ dense ≤2e-4); early # steps set global structure and are the most damage-prone. - vsa_sparsity = 0.0 if index < vsa_dense_first_n else float(batch.VSA_sparsity) + vsa_sparsity = 0.0 if index < vsa_dense_first_n else vsa_sparsity_base attn_metadata = vsa_metadata_builder.build( current_timestep=index, raw_latent_shape=(layout.num_video_latent_frames, layout.latent_height, @@ -223,18 +246,51 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward video_start = layout.num_condition_video_rows audio_start = layout.num_condition_audio_rows - batch.latents[video_start:] = self.scheduler.step( - video_velocity[0, video_start:].float(), - video_timestep, - batch.latents[video_start:], - return_dict=False, - )[0] - batch.audio_latents[audio_start:] = self.audio_scheduler.step( - audio_velocity[0, audio_start:].float(), - audio_timestep, - batch.audio_latents[audio_start:], - return_dict=False, - )[0] + if os.environ.get("FASTVIDEO_DMD_DEBUG_STATS", "0") == "1": + with torch.no_grad(): + for tag, lat, vel, st, sig in ( + ("video", batch.latents, video_velocity, video_start, self.scheduler.sigmas), + ("audio", batch.audio_latents, audio_velocity, audio_start, + self.audio_scheduler.sigmas), + ): + s = float(sig[index]) + xin = lat[st:].float() + x0dbg = xin + s * vel[0, st:].float() + logger.info( + "DMD_DEBUG step=%d %s sigma=%.4f in(std=%.4f,mean=%.4f) " + "x0(std=%.4f,mean=%.4f) v(std=%.4f)", index, tag, s, xin.std(), xin.mean(), + x0dbg.std(), x0dbg.mean(), vel[0, st:].float().std()) + if stochastic_renoise: + # Raw H3 output is clean - noise, so + # x0 = sample + sigma * output. + assert self.scheduler.sigmas is not None and self.audio_scheduler.sigmas is not None + for latents, velocity, start, sigmas in ( + (batch.latents, video_velocity, video_start, self.scheduler.sigmas), + (batch.audio_latents, audio_velocity, audio_start, self.audio_scheduler.sigmas), + ): + sigma = float(sigmas[index]) + sigma_next = float(sigmas[index + 1]) + sample = latents[start:].float() + pred_x0 = sample + sigma * velocity[0, start:].float() + if sigma_next > 0.0: + noise = torch.randn_like(pred_x0) + nxt = (1.0 - sigma_next) * pred_x0 + sigma_next * noise + else: + nxt = pred_x0 + latents[start:] = nxt.to(latents.dtype) + else: + batch.latents[video_start:] = self.scheduler.step( + video_velocity[0, video_start:].float(), + video_timestep, + batch.latents[video_start:], + return_dict=False, + )[0] + batch.audio_latents[audio_start:] = self.audio_scheduler.step( + audio_velocity[0, audio_start:].float(), + audio_timestep, + batch.audio_latents[audio_start:], + return_dict=False, + )[0] batch.step_index = index batch.timestep = video_timestep finally: diff --git a/fastvideo/pipelines/pipeline_batch_info.py b/fastvideo/pipelines/pipeline_batch_info.py index 2a7d2533e2..c50abecf8c 100644 --- a/fastvideo/pipelines/pipeline_batch_info.py +++ b/fastvideo/pipelines/pipeline_batch_info.py @@ -323,6 +323,10 @@ class TrainingBatch: # MiniMax H3 reuses the packed row boundaries from batch preparation to # split the transformer's joint sequence back into video and audio outputs. minimax_h3_layout: Any | None = None + # DMD2 flattens video and audio into one tensor. This immutable, batch-local + # layout records the exact native shapes needed to split it again; keeping + # it on the batch avoids mutable geometry state on compiled role models. + minimax_h3_dmd_layout: Any | None = None attn_metadata_vsa: AttentionMetadata | None = None attn_metadata: AttentionMetadata | None = None diff --git a/fastvideo/pipelines/preprocess/preprocess_minimax_h3_overfit.py b/fastvideo/pipelines/preprocess/preprocess_minimax_h3_overfit.py index 33a26ce991..235f723185 100644 --- a/fastvideo/pipelines/preprocess/preprocess_minimax_h3_overfit.py +++ b/fastvideo/pipelines/preprocess/preprocess_minimax_h3_overfit.py @@ -313,21 +313,35 @@ def validate_preprocessed_training_data( print(f"Validated one Crush-Smol H3 training row in {expected_path}") -def main() -> None: - """Encode the selected Crush-Smol record with each H3 component in sequence. +def main( + video: Path | None = None, + prompt: str | None = None, + output_dir: Path = OUTPUT_DIR, + model_path: Path = MODEL_PATH, +) -> None: + """Encode one audio-video-caption record with each H3 component in sequence. + With no arguments this encodes the pinned Crush-Smol record; ``--video`` + + ``--prompt`` preprocess an arbitrary mp4 (with soundtrack) instead. Releasing each component before loading the next component keeps video, audio, and text preprocessing within one GPU's memory. """ _init_single_process_distributed() - resolved_model_path = MODEL_PATH.resolve() + resolved_model_path = model_path.resolve() if not resolved_model_path.is_dir(): raise FileNotFoundError(f"Filtered MiniMax H3 model directory is missing at {resolved_model_path}") model_index = verify_model_config_and_directory(str(resolved_model_path)) - video_path, caption = load_crush_smol_training_sample( - DATA_DIR / "videos2caption.json", - DATA_DIR / "videos", - ) + if video is None: + video_path, caption = load_crush_smol_training_sample( + DATA_DIR / "videos2caption.json", + DATA_DIR / "videos", + ) + else: + if not prompt or not prompt.strip(): + raise ValueError("--prompt is required when --video is given") + video_path, caption = video, prompt.strip() + if not video_path.is_file(): + raise FileNotFoundError(f"Training video is missing at {video_path}") frames, waveform = load_training_media(video_path) pipeline_config = MiniMaxH3PipelineConfig() fastvideo_args = FastVideoArgs( @@ -345,21 +359,34 @@ def main() -> None: audio_latents = encode_audio_latents(waveform, resolved_model_path, model_index, fastvideo_args) text_embedding = encode_text_embedding(caption, resolved_model_path, model_index, fastvideo_args) record = build_parquet_record( - file_name=TRAINING_VIDEO_NAME, + file_name=video_path.name, caption=caption, video_latents=video_latents, audio_latents=audio_latents, text_embedding=text_embedding, ) - output_path = write_parquet(record, OUTPUT_DIR) - print(f"Wrote one Crush-Smol MiniMax H3 T2VA record to {output_path}") + output_path = write_parquet(record, output_dir) + print(f"Wrote one MiniMax H3 T2VA record to {output_path}") if __name__ == "__main__": parser = argparse.ArgumentParser() parser.add_argument("--validate-only", action="store_true") + parser.add_argument("--video", + type=Path, + default=None, + help="mp4 with a soundtrack (>=124 frames at 24 fps); " + "default: the pinned Crush-Smol record") + parser.add_argument("--prompt", type=str, default=None, help="caption for --video") + parser.add_argument("--output-dir", type=Path, default=OUTPUT_DIR) + parser.add_argument("--model-path", type=Path, default=MODEL_PATH) cli_args = parser.parse_args() if cli_args.validate_only: validate_preprocessed_training_data() else: - main() + main( + video=cli_args.video, + prompt=cli_args.prompt, + output_dir=cli_args.output_dir, + model_path=cli_args.model_path, + ) diff --git a/fastvideo/pipelines/preprocess/preprocess_minimax_h3_text_only.py b/fastvideo/pipelines/preprocess/preprocess_minimax_h3_text_only.py new file mode 100644 index 0000000000..a80d9dcdf0 --- /dev/null +++ b/fastvideo/pipelines/preprocess/preprocess_minimax_h3_text_only.py @@ -0,0 +1,213 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Encode a prompt list into MiniMax H3 text-only conditioning rows (data-free DMD2). + +Each line of ``--prompts-file`` is one complete H3 prompt document (for the +VidProM-H3 set: ``integrated_multimodal_description: ... overall_soundscape: +... non_diegetic_music: ...`` fed to the model verbatim as one string). Every +prompt is tokenized raw (no chat template, ``add_special_tokens=False``) and +encoded through ``MiniMaxH3ConditioningStage`` — the same Qwen3-VL layer-50 +path that produced the validated t2va overfit rows — then written as +``pyarrow_schema_text_only`` records for ``rollout_mode: simulate`` training +(``training.data.preprocessed_data_type: text_only``). + +Embeddings are stored as float32 to match the training collate, which decodes +``text_embedding_bytes`` with a hard-coded ``np.float32`` +(``fastvideo/dataset/utils.py``). At ~300 tokens x 5120 dims that is ~6 MB per +prompt — budget ~380 GB for the full 63k VidProM set. + +Sharding: ``--num-shards N`` splits the prompt list round-robin by line index; +run one process per GPU with distinct ``--shard-index``/``CUDA_VISIBLE_DEVICES`` +(and a distinct ``MASTER_PORT`` — each process initializes a one-rank process +group). Each shard writes ``/shard_XX/``; the training dataloader +walks the directory tree, so pointing ``training.data.data_path`` at +``--output-dir`` picks up every shard. Restarting a shard resumes after the +rows already on disk (delete the shard directory to re-encode from scratch). +""" + +from __future__ import annotations + +import argparse +import json +import os +import time +from pathlib import Path +from typing import Any + +import pyarrow.parquet as pq +import torch + +from fastvideo.configs.pipelines.minimax_h3 import MiniMaxH3PipelineConfig +from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter, records_to_table) +from fastvideo.dataset.dataloader.record_schema import text_only_record_creator +from fastvideo.dataset.dataloader.schema import pyarrow_schema_text_only +from fastvideo.fastvideo_args import FastVideoArgs +from fastvideo.models.loader.component_loader import PipelineComponentLoader +from fastvideo.pipelines import ForwardBatch +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_conditioning import MiniMaxH3ConditioningStage +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_input_preparation import MINIMAX_H3_KEYFRAMES_KEY +from fastvideo.utils import verify_model_config_and_directory + + +def _init_single_process_distributed(shard_index: int) -> None: + """Initialize the one-rank process groups required by component loaders. + + Concurrent shards on one host must not share a rendezvous port, so the + default port is offset by the shard index (an explicit ``MASTER_PORT`` + still wins). + """ + os.environ.setdefault("MASTER_ADDR", "127.0.0.1") + os.environ.setdefault("MASTER_PORT", str(29531 + shard_index)) + os.environ.setdefault("RANK", "0") + os.environ.setdefault("WORLD_SIZE", "1") + os.environ.setdefault("LOCAL_RANK", "0") + from fastvideo.distributed import maybe_init_distributed_environment_and_model_parallel + + maybe_init_distributed_environment_and_model_parallel(1, 1) + + +def _load_component( + name: str, + model_path: Path, + model_index: dict[str, Any], + fastvideo_args: FastVideoArgs, +) -> Any: + """Load one checkpoint component through the inference component registry.""" + transformers_or_diffusers, _ = model_index[name][:2] + return PipelineComponentLoader.load_module( + module_name=name, + component_model_path=str(model_path / name), + transformers_or_diffusers=transformers_or_diffusers, + fastvideo_args=fastvideo_args, + ) + + +def _load_shard_prompts( + prompts_file: Path, + shard_index: int, + num_shards: int, + jsonl_field: str | None = None, +) -> list[tuple[int, str, str]]: + """Return this shard's ``(global_line_index, prompt, text_name)`` triples. + + Default mode treats each non-empty line as one complete prompt document. + With ``jsonl_field`` set, each line is a JSON record; the prompt is taken + verbatim from that field (embedded newlines preserved — required by the + h3-t2va-condition-v1 contract, which forbids altering ``prompt_compiled``) + and the record's ``id`` becomes the parquet row id when present. + """ + entries: list[tuple[int, str, str]] = [] + with prompts_file.open(encoding="utf-8") as handle: + for index, line in enumerate(handle): + line = line.strip() + if not line: + continue + if jsonl_field is None: + entries.append((index, line, f"vidprom_{index:05d}")) + continue + record = json.loads(line) + prompt = record.get(jsonl_field) + if not isinstance(prompt, str) or not prompt: + raise ValueError(f"{prompts_file}:{index + 1}: empty or non-string field {jsonl_field!r}") + entries.append((index, prompt, str(record.get("id") or f"jsonl_{index:05d}"))) + return entries[shard_index::num_shards] + + +def _count_existing_rows(shard_dir: Path) -> int: + """Count rows already written under a shard directory (resume offset).""" + total = 0 + for parquet_path in sorted(shard_dir.rglob("*.parquet")): + total += pq.ParquetFile(parquet_path).metadata.num_rows + return total + + +def main(args: argparse.Namespace) -> None: + _init_single_process_distributed(args.shard_index) + + model_path = args.model_path.resolve() + if not model_path.is_dir(): + raise FileNotFoundError(f"MiniMax H3 model directory is missing at {model_path}") + model_index = verify_model_config_and_directory(str(model_path)) + + shard = _load_shard_prompts(args.prompts_file, args.shard_index, args.num_shards, args.jsonl_field) + shard_dir = args.output_dir / f"shard_{args.shard_index:02d}" + already_done = _count_existing_rows(shard_dir) if shard_dir.is_dir() else 0 + if already_done >= len(shard): + print(f"Shard {args.shard_index}/{args.num_shards}: all {len(shard)} rows already encoded") + return + todo = shard[already_done:] + if args.limit is not None: + todo = todo[:args.limit] + print(f"Shard {args.shard_index}/{args.num_shards}: {len(shard)} prompts total, " + f"{already_done} already on disk, encoding {len(todo)} now -> {shard_dir}") + + fastvideo_args = FastVideoArgs( + model_path=str(model_path), + pipeline_config=MiniMaxH3PipelineConfig(), + num_gpus=1, + tp_size=1, + sp_size=1, + hsdp_shard_dim=1, + use_fsdp_inference=False, + vae_cpu_offload=False, + text_encoder_cpu_offload=False, + ) + print("Loading MiniMax H3 tokenizer, processor, and Qwen3-VL encoder") + tokenizer = _load_component("tokenizer", model_path, model_index, fastvideo_args) + processor = _load_component("processor", model_path, model_index, fastvideo_args) + conditioner = _load_component("text_encoder", model_path, model_index, fastvideo_args) + stage = MiniMaxH3ConditioningStage( + conditioner=conditioner, + tokenizer=tokenizer, + processor=processor, + ) + + writer = ParquetDatasetWriter(out_dir=str(shard_dir), samples_per_file=args.samples_per_file) + records: list[dict[str, Any]] = [] + started = time.monotonic() + with torch.inference_mode(): + for done, (global_index, prompt, text_name) in enumerate(todo, start=1): + batch = ForwardBatch(data_type="video", prompt=prompt) + batch.extra[MINIMAX_H3_KEYFRAMES_KEY] = [] + batch = stage.forward(batch, fastvideo_args) + if not batch.prompt_embeds: + raise RuntimeError(f"MiniMax H3 conditioning returned no embedding for line {global_index}") + # float32 to match the training collate's np.frombuffer dtype. + text_embedding = batch.prompt_embeds[0].squeeze(0).float().cpu().contiguous().numpy() + records.append( + text_only_record_creator( + text_name=text_name, + text_embedding=text_embedding, + caption=prompt, + )) + + if len(records) >= args.flush_every or done == len(todo): + writer.append_table(records_to_table(records, pyarrow_schema_text_only)) + records = [] + written = writer.flush(write_remainder=done == len(todo)) + rate = done / (time.monotonic() - started) + remaining = (len(todo) - done) / rate if rate > 0 else float("inf") + print(f"[shard {args.shard_index}] {done}/{len(todo)} encoded " + f"({rate:.2f} prompts/s, ~{remaining / 60:.0f} min left, flushed {written} rows)") + + print(f"Shard {args.shard_index} complete: {already_done + len(todo)}/{len(shard)} rows in {shard_dir}") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--prompts-file", type=Path, required=True, help="one H3 prompt document per line") + parser.add_argument("--model-path", type=Path, required=True, help="MiniMax-H3 checkpoint directory") + parser.add_argument("--output-dir", type=Path, required=True, help="dataset root; shards write shard_XX/ under it") + parser.add_argument("--shard-index", type=int, default=0) + parser.add_argument("--num-shards", type=int, default=1) + parser.add_argument("--samples-per-file", type=int, default=64) + parser.add_argument("--flush-every", type=int, default=256, help="rows buffered between parquet flushes") + parser.add_argument("--limit", type=int, default=None, help="encode at most N prompts this run (smoke tests)") + parser.add_argument("--jsonl-field", + type=str, + default=None, + help="treat --prompts-file as JSONL and take the prompt verbatim from " + "this field (record 'id' becomes the row id when present)") + cli_args = parser.parse_args() + if not 0 <= cli_args.shard_index < cli_args.num_shards: + parser.error(f"--shard-index {cli_args.shard_index} must be in [0, {cli_args.num_shards})") + main(cli_args) diff --git a/fastvideo/tests/attention/test_compile_policy.py b/fastvideo/tests/attention/test_compile_policy.py new file mode 100644 index 0000000000..f9ccb2f4a6 --- /dev/null +++ b/fastvideo/tests/attention/test_compile_policy.py @@ -0,0 +1,19 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Attention compile-boundary policy tests.""" + +from fastvideo.attention.layer import _attention_compile_disabled + + +def test_attention_compile_is_enabled_by_default(monkeypatch) -> None: + monkeypatch.delenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", raising=False) + + assert not _attention_compile_disabled() + + +def test_attention_compile_escape_hatch(monkeypatch) -> None: + monkeypatch.setenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", "1") + + assert _attention_compile_disabled() + + monkeypatch.setenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", "0") + assert not _attention_compile_disabled() diff --git a/fastvideo/tests/attention/test_flash_attn_cute_custom_op.py b/fastvideo/tests/attention/test_flash_attn_cute_custom_op.py index a1d0a5167d..210aad800c 100644 --- a/fastvideo/tests/attention/test_flash_attn_cute_custom_op.py +++ b/fastvideo/tests/attention/test_flash_attn_cute_custom_op.py @@ -173,6 +173,37 @@ def test_flash_attn_func_parity_forward_backward(flash_attn_impls, dtype: torch. _assert_close(dv_test, dv_ref, dtype=dtype, is_grad=True) +def test_flash_attn_func_fullgraph_compile_backward(flash_attn_impls): + custom_flash_attn_func, _, _, _ = flash_attn_impls + + def loss(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor: + return custom_flash_attn_func(q, k, v).square().mean() + + torch.manual_seed(2) + base_inputs = [ + torch.randn( + 1, + 64, + 4, + 64, + device="cuda", + dtype=torch.bfloat16, + requires_grad=False, + ) for _ in range(3) + ] + eager_inputs = [tensor.clone().requires_grad_(True) for tensor in base_inputs] + compiled_inputs = [tensor.clone().requires_grad_(True) for tensor in base_inputs] + + eager_loss = loss(*eager_inputs) + eager_grads = torch.autograd.grad(eager_loss, eager_inputs) + compiled_loss = torch.compile(loss, fullgraph=True)(*compiled_inputs) + compiled_grads = torch.autograd.grad(compiled_loss, compiled_inputs) + + _assert_close(compiled_loss, eager_loss, dtype=torch.bfloat16) + for compiled_grad, eager_grad in zip(compiled_grads, eager_grads): + _assert_close(compiled_grad, eager_grad, dtype=torch.bfloat16, is_grad=True) + + @pytest.mark.parametrize("dtype", [torch.float16, torch.bfloat16]) @pytest.mark.parametrize("causal", [False, True]) def test_flash_attn_varlen_func_parity_forward_backward(flash_attn_impls, dtype: torch.dtype, causal: bool): diff --git a/fastvideo/tests/attention/test_vsa_h3_inference_metadata_parity.py b/fastvideo/tests/attention/test_vsa_h3_inference_metadata_parity.py new file mode 100644 index 0000000000..14c9a868a8 --- /dev/null +++ b/fastvideo/tests/attention/test_vsa_h3_inference_metadata_parity.py @@ -0,0 +1,177 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU checks that the VSA-H3 inference (validation/denoising-stage) metadata +path builds exactly what the training path builds, stays in-bounds, and that +malformed geometry fails synchronously instead of as an async kernel fault. + +Shapes mirror the v7 DMD2 validation request that motivated this test: +768x1344, 124 frames -> video latents (37, 48, 84), 207 audio latents, +patch (1, 2, 2), 3-step DMD ladder, 90% sparsity (jobs 2307/2321).""" + +import math + +import pytest +import torch + +from fastvideo.attention.backends.video_sparse_attn_h3 import (_TILE_ELEMS, MiniMaxH3VSAMetadataBuilder, + _build_block_mask, _h3_tile_geometry, + _validate_h3_tile_geometry) +from fastvideo.pipelines.basic.minimax_h3.packing import (MINIMAX_H3_TEXT_TAG, audio_latent_num_frames, + build_packed_sequence, video_latent_num_frames) +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_denoising import _h3_vsa_prefix_segments + +_CPU = torch.device("cpu") +_PATCH = (1, 2, 2) +_NUM_FRAMES = 124 # -> 37 latent frames, 207 audio latents +_LATENT = (video_latent_num_frames(_NUM_FRAMES), 768 // 16, 1344 // 16) +_NUM_AUDIO = audio_latent_num_frames(_NUM_FRAMES) +_SPARSITY = 0.9 +_DMD_STEPS = 3 + +# text lengths straddle the 256-token tile boundary; "first" adds keyframe +# condition rows (image-conditioned validation requests) +_TEXT_LENS = [7, 100, 255, 256, 257, 500] + + +def _layout(text_len: int, anchors: tuple[str, ...] = ()): + tags = torch.full((text_len, ), MINIMAX_H3_TEXT_TAG, dtype=torch.long) + return build_packed_sequence(tags, *_LATENT, _NUM_AUDIO, _PATCH, anchors) + + +def _inference_metadata(builder, layout, step_index: int, sparsity: float = _SPARSITY): + """Exactly the MiniMaxH3DenoisingStage.forward calling convention.""" + return builder.build( + current_timestep=step_index, + raw_latent_shape=(layout.num_video_latent_frames, layout.latent_height, layout.latent_width), + patch_size=_PATCH, + VSA_sparsity=sparsity, + prefix_segments=_h3_vsa_prefix_segments(layout, _PATCH), + device=_CPU, + exempt=True, + dense_layers=(), + ) + + +def _training_metadata(builder, layout, sparsity: float = _SPARSITY): + """Exactly the MiniMaxH3Model._maybe_build_vsa_metadata calling convention.""" + return builder.build( + current_timestep=0, + raw_latent_shape=(layout.num_video_latent_frames, layout.latent_height, layout.latent_width), + patch_size=_PATCH, + VSA_sparsity=sparsity, + prefix_segments=_h3_vsa_prefix_segments(layout, _PATCH), + device=_CPU, + ) + + +def _assert_in_bounds(meta, layout, tag: str): + n_tiles = meta.num_prefix_tiles + meta.num_video_tiles + sizes = meta.variable_block_sizes + assert sizes.numel() == n_tiles, tag + assert int(sizes.min()) >= 1 and int(sizes.max()) <= _TILE_ELEMS, tag + assert int(sizes.sum()) == meta.total_seq_length == layout.sequence_length, tag + idx = meta.untile_combined_index + assert idx.numel() == meta.total_seq_length, tag + assert int(idx.min()) >= 0 and int(idx.max()) < n_tiles * _TILE_ELEMS, tag + assert idx.unique().numel() == idx.numel(), f"{tag}: untile index must be injective" + # no packed row may land in a pad slot of the padded tile buffer + assert bool((idx % _TILE_ELEMS < sizes[idx // _TILE_ELEMS]).all()), tag + + +@pytest.mark.parametrize("text_len", _TEXT_LENS) +def test_inference_metadata_matches_training(text_len): + layout = _layout(text_len) + assert _h3_vsa_prefix_segments(layout, _PATCH) == (text_len, 0, _NUM_AUDIO * 2) + + meta_train = _training_metadata(MiniMaxH3VSAMetadataBuilder(), layout) + _assert_in_bounds(meta_train, layout, f"train text={text_len}") + + infer_builder = MiniMaxH3VSAMetadataBuilder() # one builder per denoise loop, as the stage does + for step in range(_DMD_STEPS): + meta_inf = _inference_metadata(infer_builder, layout, step) + _assert_in_bounds(meta_inf, layout, f"infer text={text_len} step={step}") + assert meta_inf.total_seq_length == meta_train.total_seq_length + assert meta_inf.num_prefix_tiles == meta_train.num_prefix_tiles + assert meta_inf.num_video_tiles == meta_train.num_video_tiles + assert meta_inf.exempt == meta_train.exempt + assert meta_inf.dense_layers == meta_train.dense_layers + assert meta_inf.VSA_sparsity == meta_train.VSA_sparsity + assert torch.equal(meta_inf.variable_block_sizes, meta_train.variable_block_sizes) + assert torch.equal(meta_inf.untile_combined_index, meta_train.untile_combined_index) + + +def test_keyframe_conditioned_layout_in_bounds(): + """Image-conditioned validation adds condition keyframe rows to the prefix.""" + layout = _layout(100, anchors=("first", )) + rows_per_frame = (_LATENT[1] // _PATCH[1]) * (_LATENT[2] // _PATCH[2]) + assert _h3_vsa_prefix_segments(layout, _PATCH) == (100, rows_per_frame, _NUM_AUDIO * 2) + meta = _inference_metadata(MiniMaxH3VSAMetadataBuilder(), layout, 0) + _assert_in_bounds(meta, layout, "keyframe-conditioned") + + +def test_route_a_expansion_in_bounds(): + """The 256->64 route-A remap the Triton fallback consumes stays in-bounds.""" + try: + from fastvideo_kernel import block_sparse_attn_256 + except Exception as exc: # triton driver probing raises RuntimeError on GPU-less hosts + pytest.skip(f"fastvideo_kernel unavailable here: {exc}") + layout = _layout(100) + meta = _inference_metadata(MiniMaxH3VSAMetadataBuilder(), layout, 0) + n_tiles = meta.variable_block_sizes.numel() + scores = torch.randn(1, 4, n_tiles, n_tiles) + mask = _build_block_mask(scores, meta.num_prefix_tiles, meta.num_video_tiles, _SPARSITY, exempt=True) + mask64, sizes64 = block_sparse_attn_256._expand_mask_and_sizes_256_to_64(mask, meta.variable_block_sizes) + assert mask64.shape[-2:] == (4 * n_tiles, 4 * n_tiles) + assert sizes64.numel() == 4 * n_tiles + assert int(sizes64.min()) >= 0 and int(sizes64.max()) <= 64 + assert int(sizes64.sum()) == meta.total_seq_length + # empty 64-blocks only ever pad the tail of a logical 256 tile: each + # tile's children must be its size chopped into non-increasing 64-strides + per_tile = sizes64.view(n_tiles, 4) + assert bool((per_tile[:, 0] > 0).all()), "every logical tile keeps at least one valid 64-block" + assert bool((per_tile[:, :-1] >= per_tile[:, 1:]).all()), "child sizes must be non-increasing" + assert torch.equal(per_tile.sum(dim=1), meta.variable_block_sizes.to(per_tile.dtype)) + + +def test_geometry_guard_rejects_corruption(): + """The synchronous guard must catch what would otherwise be an async fault.""" + prefix = (100, _NUM_AUDIO * 2) + dit_shape = tuple(d // p for d, p in zip(_LATENT, _PATCH, strict=True)) + (_, sizes, untile, _, _) = _h3_tile_geometry(prefix, dit_shape, _CPU) + + with pytest.raises(ValueError, match="tile sizes out of bounds"): + bad = sizes.clone() + bad[0] = _TILE_ELEMS + 1 + _validate_h3_tile_geometry(prefix, dit_shape, bad, untile) + with pytest.raises(ValueError, match="tile sizes out of bounds"): + bad = sizes.clone() + bad[-1] += 1 # sum mismatch + _validate_h3_tile_geometry(prefix, dit_shape, bad, untile) + with pytest.raises(ValueError, match="untile index"): + _validate_h3_tile_geometry(prefix, dit_shape, sizes, untile[:-1]) + with pytest.raises(ValueError, match="injective"): + bad = untile.clone() + bad[1] = int(bad[0]) # duplicate slot + _validate_h3_tile_geometry(prefix, dit_shape, sizes, bad) + with pytest.raises(ValueError, match="injective"): + bad = untile.clone() + bad[0] = sizes.numel() * _TILE_ELEMS # beyond the padded buffer + _validate_h3_tile_geometry(prefix, dit_shape, sizes, bad) + # a slot inside a partial tile's pad region is also out of bounds + partial = int((sizes < _TILE_ELEMS).nonzero()[0]) + with pytest.raises(ValueError, match="injective"): + bad = untile.clone() + bad[0] = partial * _TILE_ELEMS + int(sizes[partial]) # first pad slot + _validate_h3_tile_geometry(prefix, dit_shape, sizes, bad) + + # the untampered geometry passes + _validate_h3_tile_geometry(prefix, dit_shape, sizes, untile) + assert int(sizes.sum()) == sum(prefix) + math.prod(dit_shape) + + +if __name__ == "__main__": + for _text_len in _TEXT_LENS: + test_inference_metadata_matches_training(_text_len) + test_keyframe_conditioned_layout_in_bounds() + test_route_a_expansion_in_bounds() + test_geometry_guard_rejects_corruption() + print("all VSA-H3 inference-metadata parity checks passed") diff --git a/fastvideo/tests/attention/test_vsa_h3_metadata.py b/fastvideo/tests/attention/test_vsa_h3_metadata.py index 24292f26b4..e4c26968a1 100644 --- a/fastvideo/tests/attention/test_vsa_h3_metadata.py +++ b/fastvideo/tests/attention/test_vsa_h3_metadata.py @@ -143,6 +143,29 @@ def test_prefix_queries_stay_dense_at_high_sparsity(): "video rows should actually be sparse at 75%" +def test_tile_under_grad_does_not_reuse_shared_buffer(): + """Training forwards need fresh tile buffers; in-place reuse of the shared + holder would trip autograd's saved-tensor version check at backward.""" + meta = _build(_TINY) + impl = _impl() + seq = meta.total_seq_length + + x1 = torch.randn(1, seq, 2, 8, requires_grad=True) + x2 = torch.randn(1, seq, 2, 8, requires_grad=True) + buf1 = impl.tile(x1, meta) + saved = (buf1 * buf1).sum() # saves buf1 for backward, like the kernel + buf2 = impl.tile(x2, meta) + assert buf1 is not buf2 + saved.backward() # raises "modified by an inplace operation" on reuse + assert x1.grad is not None and x2.grad is None + + # no-grad paths (inference / rollout) keep the single-buffer reuse + with torch.no_grad(): + y1 = impl.tile(x1.detach(), meta) + y2 = impl.tile(x2.detach(), meta) + assert y1 is y2 + + # --------------------------------------------------------------------------- # 64-token (4,4,4) tile geometry # --------------------------------------------------------------------------- @@ -251,6 +274,7 @@ def test_builder_rejects_unknown_tile_size(): test_mask_policy() test_sparsity_zero_matches_dense_sdpa() test_prefix_queries_stay_dense_at_high_sparsity() + test_tile_under_grad_does_not_reuse_shared_buffer() test_geometry_tile64_ragged_tails() test_geometry_tile64_production_shape() test_sparsity_zero_matches_dense_sdpa_tile64() diff --git a/fastvideo/tests/dataset/test_exact_shape_bucket_sampler.py b/fastvideo/tests/dataset/test_exact_shape_bucket_sampler.py new file mode 100644 index 0000000000..9cdb0e5443 --- /dev/null +++ b/fastvideo/tests/dataset/test_exact_shape_bucket_sampler.py @@ -0,0 +1,115 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Exact-shape global microbatch scheduling contracts.""" + +from pathlib import Path + +import pytest + +from fastvideo.dataset.parquet_dataset_map_style import ( + DP_SP_BatchSampler, + shape_bucket_ids_from_parquet_files, +) +from fastvideo.dataset.shape_bucket import parse_video_shape_bucket_id + + +@pytest.mark.parametrize( + "bucket_id", + [ + "bucket=0x768-124f", + "bucket=1344x0-124f", + "bucket=1344x768-0f", + "bucket=1344X768-124f", + "bucket=1344x768-124F", + "bucket=01344x768-124f", + "1344x768-124f", + "bucket=1344x768", + "bucket=1344x768-124f-extra", + ], +) +def test_portable_bucket_id_rejects_noncanonical_spelling(bucket_id: str) -> None: + with pytest.raises(ValueError, match="bucket=x"): + parse_video_shape_bucket_id(bucket_id) + + +def test_bucket_ids_expand_from_canonical_parquet_ancestors() -> None: + files = [ + "/shared/source-a/data/bucket=1344x768-124f/part-0.parquet", + "/shared/source-b/data/bucket=768x1344-362f/part-1.parquet", + ] + assert shape_bucket_ids_from_parquet_files(files, [2, 1]) == [ + "bucket=1344x768-124f", + "bucket=1344x768-124f", + "bucket=768x1344-362f", + ] + + with pytest.raises(ValueError, match="exactly one ancestor"): + shape_bucket_ids_from_parquet_files(["/shared/data/part.parquet"], [1]) + malformed = str(Path("/shared/data/bucket=1344-768-124f/part.parquet")) + with pytest.raises(ValueError, match="positive decimal"): + shape_bucket_ids_from_parquet_files([malformed], [1]) + + +def _rank_samplers(bucket_ids: list[str], *, seed: int = 17) -> list[DP_SP_BatchSampler]: + return [ + DP_SP_BatchSampler( + batch_size=1, + dataset_size=len(bucket_ids), + num_sp_groups=32, + sp_world_size=1, + global_rank=rank, + drop_last=True, + seed=seed, + sample_bucket_ids=bucket_ids, + ) for rank in range(32) + ] + + +def test_rare_bucket_is_padded_not_dropped_and_all_ranks_share_schedule() -> None: + rare = "bucket=480x832-124f" + common = "bucket=1344x768-362f" + bucket_ids = [rare] * 7 + [common] * 64 + samplers = _rank_samplers(bucket_ids) + + assert all(sampler.bucket_schedule == samplers[0].bucket_schedule for sampler in samplers) + assert samplers[0].bucket_schedule is not None + assert samplers[0].bucket_schedule.count(rare) == 1 + assert samplers[0].bucket_schedule.count(common) == 2 + assert samplers[0].bucket_padding == {rare: 25} + assert samplers[0].num_padded_samples == 25 + + batches_by_rank = [list(sampler) for sampler in samplers] + observed_originals: set[int] = set() + for step, scheduled_bucket in enumerate(samplers[0].bucket_schedule): + step_indices = [batches_by_rank[rank][step][0] for rank in range(32)] + assert {bucket_ids[index] for index in step_indices} == {scheduled_bucket} + observed_originals.update(step_indices) + if scheduled_bucket == rare: + assert set(step_indices) == set(range(7)) + assert len(step_indices) == 32 + + # Padding may repeat rows, but every frozen row participates at least once. + assert observed_originals == set(range(len(bucket_ids))) + + +def test_bucket_schedule_is_seeded_and_sp_ranks_share_indices() -> None: + bucket_ids = (["bucket=1344x768-124f"] * 16 + ["bucket=768x1344-362f"] * 16) + + def sampler(rank: int, seed: int) -> DP_SP_BatchSampler: + return DP_SP_BatchSampler( + batch_size=1, + dataset_size=len(bucket_ids), + num_sp_groups=4, + sp_world_size=2, + global_rank=rank, + drop_last=True, + seed=seed, + sample_bucket_ids=bucket_ids, + ) + + first = [list(sampler(rank, 9)) for rank in range(8)] + replay = [list(sampler(rank, 9)) for rank in range(8)] + changed = [list(sampler(rank, 10)) for rank in range(8)] + assert first == replay + assert first != changed + for sp_leader in range(0, 8, 2): + assert first[sp_leader] == first[sp_leader + 1] diff --git a/fastvideo/tests/dataset/test_parquet_dataset_map_style.py b/fastvideo/tests/dataset/test_parquet_dataset_map_style.py index e899cd0ae6..5413fccc7c 100644 --- a/fastvideo/tests/dataset/test_parquet_dataset_map_style.py +++ b/fastvideo/tests/dataset/test_parquet_dataset_map_style.py @@ -5,10 +5,17 @@ import pickle +import pyarrow as pa +import pyarrow.parquet as pq import pytest from fastvideo.dataset import parquet_dataset_map_style as parquet_dataset -from fastvideo.dataset.parquet_dataset_map_style import _parse_data_path_specs +from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2va, pyarrow_schema_text_only +from fastvideo.dataset.parquet_dataset_map_style import ( + _parse_data_path_specs, + LatentsParquetMapStyleDataset, + read_row_from_parquet_file, +) def test_parse_data_path_specs_accepts_old_repeat_string() -> None: @@ -119,3 +126,68 @@ def test_get_parquet_files_and_length_raises_when_all_repeats_dropped() -> None: # to read, the mix branch fails loudly instead of yielding an empty dataset. with pytest.raises(FileNotFoundError): parquet_dataset.get_parquet_files_and_length({"data/a": 0}) + + +def test_read_row_projects_text_columns_from_t2va_superset(tmp_path) -> None: + parquet_path = tmp_path / "sample.parquet" + row = { + "id": ["sample-0"], + "vae_latent_bytes": [b"video-must-not-be-read"], + "vae_latent_shape": [[24, 2, 4, 4]], + "vae_latent_dtype": ["float32"], + "audio_latent_bytes": [b"audio-must-not-be-read"], + "audio_latent_shape": [[2, 32, 8]], + "audio_latent_dtype": ["float32"], + "text_embedding_bytes": [b"text"], + "text_embedding_shape": [[1, 4]], + "text_embedding_dtype": ["float32"], + "file_name": ["sample.mp4"], + "caption": ["prompt"], + "media_type": ["video"], + "width": [64], + "height": [64], + "num_frames": [5], + "duration_sec": [5.0 / 24.0], + "fps": [24.0], + "audio_sample_rate": [32_000], + } + pq.write_table(pa.Table.from_pydict(row, schema=pyarrow_schema_t2va), parquet_path) + text_columns = [ + "id", + "text_embedding_bytes", + "text_embedding_shape", + "text_embedding_dtype", + "caption", + ] + + projected = read_row_from_parquet_file([str(parquet_path)], 0, [1], columns=text_columns) + + assert projected == { + "id": "sample-0", + "text_embedding_bytes": b"text", + "text_embedding_shape": [1, 4], + "text_embedding_dtype": "float32", + "caption": "prompt", + } + + +def test_dataset_projects_its_declared_schema_columns(monkeypatch) -> None: + observed = {} + + def fake_read(parquet_files, global_row_idx, lengths, columns=None): + observed["columns"] = columns + return {"id": "sample-0"} + + monkeypatch.setattr(parquet_dataset, "read_row_from_parquet_file", fake_read) + monkeypatch.setattr(parquet_dataset, "collate_rows_from_parquet_schema", lambda rows, *args, **kwargs: rows[0]) + dataset = LatentsParquetMapStyleDataset.__new__(LatentsParquetMapStyleDataset) + dataset.parquet_files = ("unused.parquet", ) + dataset.lengths = (1, ) + dataset.parquet_schema = pyarrow_schema_text_only + dataset.text_padding_length = 512 + dataset.cfg_rate = 0.0 + dataset.seed = 42 + dataset.sample_bucket_ids = None + + assert dataset.__getitems__([0]) == {"id": "sample-0", "_sample_index": 0} + assert observed["columns"] == pyarrow_schema_text_only.names diff --git a/fastvideo/tests/dataset/test_validation_dataset.py b/fastvideo/tests/dataset/test_validation_dataset.py new file mode 100644 index 0000000000..e511164331 --- /dev/null +++ b/fastvideo/tests/dataset/test_validation_dataset.py @@ -0,0 +1,33 @@ +# SPDX-License-Identifier: Apache-2.0 +import json + +import pytest + +from fastvideo.dataset.validation_dataset import ValidationDataset + + +@pytest.mark.parametrize("wrapped", [False, True]) +def test_validation_json_accepts_array_and_legacy_data_wrapper( + tmp_path, + monkeypatch: pytest.MonkeyPatch, + wrapped: bool, +) -> None: + rows = [{ + "caption": "A held-out prompt", + "source": "source-a", + "sample_id": "sample-1", + "width": 128, + "height": 80, + "num_frames": 39, + }] + document = {"data": rows} if wrapped else rows + manifest = tmp_path / "validation.json" + manifest.write_text(json.dumps(document), encoding="utf-8") + monkeypatch.setattr("fastvideo.dataset.validation_dataset.get_world_rank", lambda: 0) + monkeypatch.setattr("fastvideo.dataset.validation_dataset.get_world_size", lambda: 1) + monkeypatch.setattr("fastvideo.dataset.validation_dataset.get_sp_world_size", lambda: 1) + + dataset = ValidationDataset(str(manifest)) + + assert dataset.all_samples == rows + assert list(dataset)[0]["prompt"] == rows[0]["caption"] diff --git a/fastvideo/tests/loader/test_shard_cache.py b/fastvideo/tests/loader/test_shard_cache.py new file mode 100644 index 0000000000..a907b078db --- /dev/null +++ b/fastvideo/tests/loader/test_shard_cache.py @@ -0,0 +1,220 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU contract tests for the sharded base-weight cache. + +Runs on a single-process gloo group with a (1, 1) CPU device mesh: DTensor +round-trip through write/load, the FQN reconciliation matrix (allowed +zero-init params, disallowed extras, shape mismatches), and the +never-fail-the-run contract. +""" + +import os + +import pytest +import torch +import torch.distributed as dist +import torch.nn as nn +from torch.distributed.device_mesh import init_device_mesh +from torch.distributed.tensor import Replicate, Shard, distribute_tensor + +from fastvideo.models.loader.shard_cache import ( + ShardCacheContext, + try_load_from_shard_cache, + write_shard_cache, +) + + +@pytest.fixture(scope="module") +def cpu_mesh(): + os.environ.setdefault("MASTER_ADDR", "127.0.0.1") + os.environ.setdefault("MASTER_PORT", "29581") + if not dist.is_initialized(): + dist.init_process_group("gloo", rank=0, world_size=1) + return init_device_mesh("cpu", (1, 1), mesh_dim_names=("replicate", "shard")) + + +def _make_model( + cpu_mesh, + *, + extra_param: str | None = None, + weight_rows: int = 8, + dtype: torch.dtype = torch.float32, +) -> nn.Module: + model = nn.Module() + placements = (Replicate(), Shard(0)) + weight = distribute_tensor(torch.randn(weight_rows, 4, dtype=dtype), cpu_mesh, placements) + bias = distribute_tensor(torch.randn(weight_rows, dtype=dtype), cpu_mesh, placements) + model.register_parameter("weight", nn.Parameter(weight)) + model.register_parameter("bias", nn.Parameter(bias)) + model.register_buffer("scale", torch.full((1, ), 2.0)) + if extra_param is not None: + extra = distribute_tensor(torch.randn(4, 4), cpu_mesh, placements) + model.register_parameter(extra_param.replace(".", "_"), nn.Parameter(extra)) + # register under the dotted name via a child module for realism + model.reverse_param_names_mapping = {"weight": ("hf.weight", None, None)} + return model + + +def _ctx(tmp_path) -> ShardCacheContext: + return ShardCacheContext(entry_dir=tmp_path / "entry", key="testkey", shard_index=0, num_shards=1, is_writer=True) + + +def test_round_trip_restores_tensors_and_reverse_mapping(cpu_mesh, tmp_path): + src = _make_model(cpu_mesh) + ctx = _ctx(tmp_path) + write_shard_cache(src, ctx) + + dst = _make_model(cpu_mesh) + with torch.no_grad(): + dst.weight.mul_(0) + dst.bias.mul_(0) + dst.reverse_param_names_mapping = {} + assert try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + assert torch.equal(dst.weight.to_local(), src.weight.to_local()) + assert torch.equal(dst.bias.to_local(), src.bias.to_local()) + assert torch.equal(dst.scale, src.scale) + assert dst.reverse_param_names_mapping == {"weight": ("hf.weight", None, None)} + + +def test_allowed_new_param_zero_inits_on_hit(cpu_mesh, tmp_path): + src = _make_model(cpu_mesh) + ctx = _ctx(tmp_path) + write_shard_cache(src, ctx) + + dst = _make_model(cpu_mesh) + gate = distribute_tensor(torch.randn(4, 4), cpu_mesh, (Replicate(), Shard(0))) + dst.register_parameter("to_gate_compress", nn.Parameter(gate)) + assert try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + assert torch.equal(dst.to_gate_compress.to_local(), torch.zeros(4, 4)) + + +def test_disallowed_missing_param_misses_without_mutation(cpu_mesh, tmp_path): + src = _make_model(cpu_mesh) + ctx = _ctx(tmp_path) + write_shard_cache(src, ctx) + + dst = _make_model(cpu_mesh) + mystery = distribute_tensor(torch.randn(4, 4), cpu_mesh, (Replicate(), Shard(0))) + dst.register_parameter("mystery", nn.Parameter(mystery)) + before = dst.weight.to_local().clone() + assert not try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + assert torch.equal(dst.weight.to_local(), before) + + +def test_shape_mismatch_misses(cpu_mesh, tmp_path): + src = _make_model(cpu_mesh) + ctx = _ctx(tmp_path) + write_shard_cache(src, ctx) + + dst = _make_model(cpu_mesh, weight_rows=9) + assert not try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + + +def test_missing_entry_misses_cleanly(cpu_mesh, tmp_path): + dst = _make_model(cpu_mesh) + ctx = ShardCacheContext(entry_dir=tmp_path / "absent", key="k2", shard_index=0, num_shards=1, is_writer=True) + assert not try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + + +def test_ac_wrapped_buffer_stays_buffer_on_cache_hit(cpu_mesh, tmp_path): + """Warm-path counterpart of the cold-path AC-prefix fix in fsdp_load. + + Under pre-FSDP activation checkpointing the model handed to + ``try_load_from_shard_cache`` is already checkpoint-wrapped: + ``state_dict()`` (and manifest) keys are clean, but raw + ``named_buffers()`` keys carry the ``_checkpoint_wrapped_module.`` + segment. The buffer-membership test must compare canonical names, or a + cached persistent buffer inside a wrapped block is reassigned as an + ``nn.Parameter`` by ``load_state_dict(assign=True)`` — on warm boots + only, silently diverging from the (already fixed) cold-boot path. + """ + from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( + checkpoint_wrapper, ) + + def _block_model() -> nn.Module: + model = nn.Module() + block = nn.Module() + weight = distribute_tensor(torch.randn(8, 4), cpu_mesh, (Replicate(), Shard(0))) + block.register_parameter("weight", nn.Parameter(weight)) + block.register_buffer("gain", torch.full((4, ), 3.0)) + model.block = block + model.reverse_param_names_mapping = {} + return model + + # Cold boot writes the cache from the same (wrapped) model shape; keys in + # the manifest are clean either way because state_dict strips the prefix. + src = _block_model() + src.block = checkpoint_wrapper(src.block) + ctx = _ctx(tmp_path) + write_shard_cache(src, ctx) + assert "block.gain" in src.state_dict() + + dst = _block_model() + dst.block = checkpoint_wrapper(dst.block) + with torch.no_grad(): + dst.block.weight.mul_(0) + dst.block.gain.mul_(0) + assert try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + + buffer_names = {name for name, _ in dst.named_buffers()} + parameter_names = {name for name, _ in dst.named_parameters()} + assert "block._checkpoint_wrapped_module.gain" in buffer_names + assert not any(name.endswith("gain") for name in parameter_names) + assert torch.equal(dst.block.gain, torch.full((4, ), 3.0)) + assert torch.equal(dst.block.weight.to_local(), src.block.weight.to_local()) + + +def test_nonzero_rank_writer_emits_manifest_for_per_node_roots(cpu_mesh, tmp_path, monkeypatch): + """Every shard writer must write the manifest, not just global rank 0. + + With FASTVIDEO_WEIGHT_SHARD_CACHE_PER_NODE=1 each node keeps its own copy + of the entry; ranks on non-head nodes write their shard files into their + node's tmpfs but (pre-fix) never a manifest, so try_load_from_shard_cache + failed its `manifest.json` existence check there and the all-rank vote + turned every multi-node warm boot into a full load (observed on the + h3-compile-ab job 2592 b2_on_warm leg: shard4-7 present on the second + tray, manifest.json absent). + """ + import fastvideo.models.loader.shard_cache as sc + + src = _make_model(cpu_mesh) + ctx = _ctx(tmp_path) + # Simulate a rank on a non-head node: still a writer (per-node root), but + # dist.get_rank() != 0. + monkeypatch.setattr(sc.dist, "get_rank", lambda: 4) + write_shard_cache(src, ctx) + assert (ctx.entry_dir / "manifest.json").is_file() + + dst = _make_model(cpu_mesh) + with torch.no_grad(): + dst.weight.mul_(0) + dst.reverse_param_names_mapping = {} + assert try_load_from_shard_cache(dst, ctx, torch.device("cpu")) + assert torch.equal(dst.weight.to_local(), src.weight.to_local()) + + +def test_model_selected_dtype_rejects_stale_cache_and_hits_fresh_cache(cpu_mesh, tmp_path): + stale = _make_model(cpu_mesh, dtype=torch.bfloat16) + stale_ctx = _ctx(tmp_path / "stale") + write_shard_cache(stale, stale_ctx) + + def select_dtype(name: str, default: torch.dtype) -> torch.dtype: + return torch.float32 if name == "weight" else default + + destination = _make_model(cpu_mesh, dtype=torch.bfloat16) + destination._get_parameter_dtype = select_dtype + assert not try_load_from_shard_cache(destination, stale_ctx, torch.device("cpu")) + + fresh = _make_model(cpu_mesh, dtype=torch.bfloat16) + fresh_weight = distribute_tensor( + torch.randn(8, 4, dtype=torch.float32), + cpu_mesh, + (Replicate(), Shard(0)), + ) + fresh.weight = nn.Parameter(fresh_weight) + fresh_ctx = _ctx(tmp_path / "fresh") + write_shard_cache(fresh, fresh_ctx) + + destination = _make_model(cpu_mesh, dtype=torch.bfloat16) + destination._get_parameter_dtype = select_dtype + assert try_load_from_shard_cache(destination, fresh_ctx, torch.device("cpu")) + assert destination.weight.dtype == torch.float32 diff --git a/fastvideo/tests/ops/quantization/test_allowlist_mirror.py b/fastvideo/tests/ops/quantization/test_allowlist_mirror.py new file mode 100644 index 0000000000..6057f994aa --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_allowlist_mirror.py @@ -0,0 +1,43 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Guard against drift between the two copies of the zero-init allowlist. + +``fsdp_load`` owns the canonical ``ALLOWED_NEW_PARAM_PATTERNS`` and uses it to +decide which model parameters may be zero-initialized when a checkpoint does not +supply them. ``shard_cache`` keeps a literal mirror of that tuple because +``fsdp_load`` imports ``shard_cache`` -- importing back would be circular. + +A mismatch does not fail a run; it silently keeps the weight shard cache cold +for any quantization config that registers new parameters (for example +``AbsMaxFP8``, whose ``scale_weight`` / ``scale_input`` appear in the model but +never in the checkpoint). That is an invisible performance regression, so this +test pins the two tuples together. +""" + +import unittest + +from fastvideo.models.loader.fsdp_load import ALLOWED_NEW_PARAM_PATTERNS +from fastvideo.models.loader.shard_cache import _ALLOWED_NEW_PARAM_PATTERNS + + +class TestAllowlistMirror(unittest.TestCase): + """``shard_cache`` must mirror ``fsdp_load``'s allowlist exactly.""" + + def test_mirror_matches_canonical_source(self): + self.assertEqual( + _ALLOWED_NEW_PARAM_PATTERNS, + ALLOWED_NEW_PARAM_PATTERNS, + "shard_cache._ALLOWED_NEW_PARAM_PATTERNS has drifted from " + "fsdp_load.ALLOWED_NEW_PARAM_PATTERNS. Update the mirror in " + "fastvideo/models/loader/shard_cache.py to match.", + ) + + def test_quant_scale_params_are_shared(self): + """The specific names a quantizing loader depends on are present in both.""" + for name in ("scale_weight", "scale_input"): + with self.subTest(param=name): + self.assertIn(name, ALLOWED_NEW_PARAM_PATTERNS) + self.assertIn(name, _ALLOWED_NEW_PARAM_PATTERNS) + + +if __name__ == "__main__": + unittest.main() diff --git a/fastvideo/tests/ops/quantization/test_int8_affine_config.py b/fastvideo/tests/ops/quantization/test_int8_affine_config.py new file mode 100644 index 0000000000..32e8406f3c --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_int8_affine_config.py @@ -0,0 +1,339 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU-only tests for the affine group-64 INT8 quantization config. + +No CUDA, no model weights, no inference stack: every test here runs on a +laptop. The three things being pinned are (1) the module imports without +flashinfer/CUDA, (2) the quantizer reproduces the validated MLX affine math +and round-trips within a stated tolerance, and (3) layer selection picks +H3's attention/FFN GEMMs and *excludes* ``attn.to_gate_compress``. + +Test 3 is the load-bearing one: ``to_gate_compress`` is H3's VSA sparse +routing gate, and quantizing it changes a discrete routing decision rather +than perturbing an output. +""" + +from __future__ import annotations + +import importlib +import math +import sys + +import pytest +import torch + +from fastvideo.layers.quantization.int8_affine_config import ( + INT8AffineConfig, + MINIMAX_H3_INT8_AFFINE_SUFFIXES, + int8_affine_dequantize, + int8_affine_quantize, + minimax_h3_int8_affine_prefixes, +) + +# --------------------------------------------------------------------------- +# (a) importability without CUDA / inference deps +# --------------------------------------------------------------------------- + + +def test_config_module_imports_without_cuda_or_flashinfer(): + """The config must import on a CPU-only host with no flashinfer.""" + assert "flashinfer" not in sys.modules, "importing the config must not pull in flashinfer" + module = importlib.import_module("fastvideo.layers.quantization.int8_affine_config") + assert module is not None + assert "flashinfer" not in sys.modules + + +def test_config_metadata_is_well_formed(): + cfg = INT8AffineConfig() + assert cfg.get_name() == "INT8Affine" + assert cfg.group_size == 64 + assert cfg.bits == 8 + assert torch.bfloat16 in cfg.get_supported_act_dtypes() + assert cfg.get_config_filenames() == [] + assert INT8AffineConfig.get_min_capability() >= 70 + # from_config round-trips the constructor surface the loader may pass. + rebuilt = INT8AffineConfig.from_config({"group_size": 32, "bits": 4}) + assert (rebuilt.group_size, rebuilt.bits) == (32, 4) + + +def test_invalid_bits_and_group_size_are_rejected(): + with pytest.raises(ValueError): + INT8AffineConfig(bits=16) + with pytest.raises(ValueError): + INT8AffineConfig(group_size=0) + + +# --------------------------------------------------------------------------- +# (b) quantizer correctness +# --------------------------------------------------------------------------- + + +def _independent_affine_reference(w: torch.Tensor, group_size: int = 64, bits: int = 8): + """A from-scratch restatement of the MLX affine algorithm. + + Written against the description in ``mlx_affine_qat.py``'s module + docstring (per-group min/max, magnitude-anchored sign, exact-integer + anchor re-expression, rint rounding, clamp to ``[0, 2**bits-1]``) rather + than by copying the implementation, so it is a real cross-check of the + transcription and not a tautology. + """ + n_bins = float((1 << bits) - 1) + flat = w.reshape(-1, w.shape[-1] // group_size, group_size).float() + lo = flat.amin(dim=-1) + hi = flat.amax(dim=-1) + # Anchor at whichever endpoint has the larger magnitude. + anchor = torch.where(lo.abs() > hi.abs(), lo, hi) + other = torch.where(lo.abs() > hi.abs(), hi, lo) + step = ((hi - lo) / n_bins).clamp_min(1e-7) + step = torch.where(lo.abs() > hi.abs(), step, -step) + # Re-express the anchor as an exact integer multiple of the step so the + # extreme value round-trips exactly. + q0 = torch.round(anchor / step) + use = q0 != 0 + step = torch.where(use, anchor / torch.where(use, q0, torch.ones_like(q0)), step) + zero = torch.where(use, anchor, torch.zeros_like(anchor)) + del other + codes = torch.round((flat - zero.unsqueeze(-1)) / step.unsqueeze(-1)).clamp(0.0, n_bins) + return codes.to(torch.int64), step, zero + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16, torch.float16]) +def test_quantize_matches_independent_reference(dtype): + torch.manual_seed(0) + w = torch.randn(37, 256, dtype=dtype) + codes, scales, biases = int8_affine_quantize(w, group_size=64, bits=8) + ref_codes, ref_scales, ref_bias = _independent_affine_reference(w, group_size=64, bits=8) + + assert codes.dtype == torch.uint8 + # Grouped shape, matching the reference contract. + assert codes.shape == (*w.shape[:-1], w.shape[-1] // 64, 64) + assert scales.shape == (*w.shape[:-1], w.shape[-1] // 64) + assert torch.equal(codes.reshape(ref_codes.shape).to(torch.int64), ref_codes) + # The quantizer casts scales/biases to the input dtype (the reference does + # the same), so compare at that dtype: the fp32 solve is identical, the + # bf16/fp16 store is the only lossy step. + torch.testing.assert_close(scales, ref_scales.to(dtype), rtol=0, atol=0) + torch.testing.assert_close(biases, ref_bias.to(dtype), rtol=0, atol=0) + + +def test_quantize_matches_in_tree_mlx_reference_when_available(): + """Parity against the real ``mlx_affine_qat`` transcription, if present. + + That module lives in a sibling worktree today, so this is skipped unless + it has landed next to us; when it does, the two transcriptions must agree + bit-for-bit on codes and exactly on the fp32 scales/biases. + """ + mlx = pytest.importorskip("fastvideo.layers.quantization.mlx_affine_qat", + reason="mlx_affine_qat.py is not in this tree yet") + torch.manual_seed(1) + w = torch.randn(16, 128, dtype=torch.float32) + codes, scales, biases = int8_affine_quantize(w, group_size=64, bits=8) + ref_codes, ref_scales, ref_bias = mlx.mlx_affine_quantize_reference(w, group_size=64, bits=8) + assert torch.equal(codes.reshape(ref_codes.shape).to(torch.int64), ref_codes.to(torch.int64)) + torch.testing.assert_close(scales, ref_scales, rtol=0, atol=0) + torch.testing.assert_close(biases, ref_bias, rtol=0, atol=0) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_roundtrip_error_is_within_documented_tolerance(dtype): + """Group-64 8-bit affine on N(0,1) weights. + + Tolerances are measured over seeds 0-5, not guessed, then given ~2x + headroom. Group-64 8-bit spends ~1/255 of the per-group range per step + (step ~= 0.016 for a 64-sample N(0,1) group), so the reconstruction error + is ~step/sqrt(12) RMS and ~step/2 at worst: + + - fp32 source: max/max|w| <= 0.0055, rms/max|w| <= 0.0013 + - bf16 source: max/max|w| <= 0.0093, rms/max|w| <= 0.0021 (the source is + itself bf16-rounded before quantizing) + """ + torch.manual_seed(2) + w = torch.randn(64, 1024, dtype=dtype) + codes, scales, biases = int8_affine_quantize(w, group_size=64, bits=8) + # Dequant returns the scales' dtype (the reference's contract) — fp32 for an + # fp32 source, bf16 for a bf16 one — so measure in fp32. + deq = int8_affine_dequantize(codes, scales, biases, out_shape=w.shape).float() + + assert deq.shape == w.shape + top = w.float().abs().max().item() + rel_max = (w.float() - deq).abs().max().item() / top + rel_rms = math.sqrt(((w.float() - deq)**2).mean().item()) / top + max_limit, rms_limit = (0.02, 0.005) if dtype is torch.bfloat16 else (0.01, 0.003) + assert rel_max < max_limit, f"max relative error {rel_max:.5f} exceeded {max_limit} for {dtype}" + assert rel_rms < rms_limit, f"rms relative error {rel_rms:.5f} exceeded {rms_limit} for {dtype}" + + # The extreme value of each group is the quantizer's anchor: it is + # re-expressed as an exact integer multiple of the scale, so it must + # round-trip to the precision the stored scales allow. Exact for an fp32 + # source; bf16-level for a bf16 source, whose scales are themselves bf16. + grouped_w = w.float().reshape(-1, 64) + grouped_deq = deq.reshape(-1, 64) + extreme = grouped_w.abs().max(dim=-1).values + at_extreme = grouped_w.abs() == extreme.unsqueeze(-1) + rtol, atol = (1e-5, 1e-5) if dtype is torch.float32 else (1e-2, 1e-2) + assert torch.allclose(grouped_w[at_extreme], grouped_deq[at_extreme], rtol=rtol, atol=atol) + + +def test_group_size_must_divide_input_dim(): + with pytest.raises(ValueError, match="not divisible"): + int8_affine_quantize(torch.randn(4, 100), group_size=64) + + +def test_codes_never_exceed_uint8_range(): + torch.manual_seed(3) + for _ in range(5): + w = torch.randn(8, 256) * torch.rand(8, 1) * 10 + codes, _, _ = int8_affine_quantize(w, group_size=64, bits=8) + assert codes.max().item() <= 255 + + +# --------------------------------------------------------------------------- +# (c) H3 layer selection — the important test +# --------------------------------------------------------------------------- + +# Every linear name H3's DiT actually builds, from fastvideo/models/dits/minimax_h3.py. +_H3_INCLUDED = [ + "minimax_h3.transformer_blocks.0.attn.to_q", + "minimax_h3.transformer_blocks.0.attn.to_k", + "minimax_h3.transformer_blocks.0.attn.to_v", + "minimax_h3.transformer_blocks.0.attn.to_out", + "minimax_h3.transformer_blocks.0.ff.fc_in", + "minimax_h3.transformer_blocks.0.ff.fc_out", + "minimax_h3.transformer_blocks.49.attn.to_q", + "minimax_h3.transformer_blocks.49.ff.fc_out", + "minimax_h3.token_refiner.refiner_blocks.0.attn.to_q", + "minimax_h3.token_refiner.refiner_blocks.0.attn.to_out", + "minimax_h3.token_refiner.refiner_blocks.1.ff.fc_in", + "minimax_h3.transformer_blocks.0.adaln_proj.linear", +] + +_H3_EXCLUDED = [ + # THE critical exclusion: the VSA sparse-attention gate. + "minimax_h3.transformer_blocks.0.attn.to_gate_compress", + "minimax_h3.transformer_blocks.49.attn.to_gate_compress", + # Global timestep-basis projector. + "minimax_h3.adaln_basis", + # Modules H3 itself pins to fp32. + "minimax_h3.proj_in", + "minimax_h3.audio_proj_in", + "minimax_h3.proj_out", + "minimax_h3.audio_proj_out", + "minimax_h3.time_embedder.fc_in", + "minimax_h3.time_embedder.fc_out", + # Text input projection (excluded by default; see the config docstring). + "minimax_h3.context_embedder", + # Non-linear / unrelated names must not be swept in. + "minimax_h3.transformer_blocks.0.norm1", + "minimax_h3.rope", +] + + +def test_h3_layer_selection_include_and_exclude_sets(): + cfg = INT8AffineConfig.for_minimax_h3() + for prefix in _H3_INCLUDED: + assert cfg.is_target_layer(prefix), f"expected {prefix!r} to be quantized" + for prefix in _H3_EXCLUDED: + assert not cfg.is_target_layer(prefix), f"expected {prefix!r} to be EXCLUDED" + + +def test_to_gate_compress_is_excluded_even_by_a_broad_allowlist(): + """The deny list is fail-closed: no constructor argument re-enables it. + + This is the regression guard for the failure mode the recon flagged — + a name that matches no "norm"/"scale_shift_table"-style heuristic being + silently swept into a broad suffix rule. + """ + hostile = INT8AffineConfig( + layer_suffixes=("to_q", "to_k", "to_v", "to_out", "to_gate_compress"), + exclude_substrings=(), # caller tries to clear the deny list + ) + assert not hostile.is_target_layer("minimax_h3.transformer_blocks.7.attn.to_gate_compress") + # The same broad rule still reaches the real projections. + assert hostile.is_target_layer("minimax_h3.transformer_blocks.7.attn.to_q") + + # And via an explicit target_layers set, which bypasses suffix matching. + explicit = INT8AffineConfig( + target_layers=("minimax_h3.transformer_blocks.7.attn.to_gate_compress", ), + ) + assert not explicit.is_target_layer("minimax_h3.transformer_blocks.7.attn.to_gate_compress") + + +def test_enumerated_h3_prefixes_agree_with_selection(): + """The literal enumerated set and the suffix rule must pick the same layers.""" + cfg = INT8AffineConfig.for_minimax_h3() + enumerated = minimax_h3_int8_affine_prefixes() + assert len(enumerated) == 50 * 7 + 2 * 6 # 50 blocks x 7 suffixes, 2 refiner blocks x 6 + for prefix in enumerated: + assert cfg.is_target_layer(prefix), f"enumerated prefix {prefix!r} not selected by the suffix rule" + # And nothing the suffix rule selects in the H3 blocks is missing from the + # enumeration: walk the two block scopes and compare. + selected = { + f"minimax_h3.{scope}.{i}.{suffix}" + for scope, count in (("transformer_blocks", 50), ("token_refiner.refiner_blocks", 2)) + for i in range(count) + for suffix in MINIMAX_H3_INT8_AFFINE_SUFFIXES + if not (scope != "transformer_blocks" and suffix.startswith("adaln_proj")) + if cfg.is_target_layer(f"minimax_h3.{scope}.{i}.{suffix}") + } + assert selected == set(enumerated) + + +def test_non_linear_layers_get_no_quant_method(): + from fastvideo.layers.linear import ReplicatedLinear + + cfg = INT8AffineConfig.for_minimax_h3() + linear = ReplicatedLinear(64, 64, bias=False, quant_config=cfg, prefix="minimax_h3.transformer_blocks.0.attn.to_q") + assert linear.quant_method is not None + assert linear.quant_method.__class__.__name__ == "INT8AffineQuantizeMethod" + + gate = ReplicatedLinear(64, 64, bias=False, quant_config=cfg, + prefix="minimax_h3.transformer_blocks.0.attn.to_gate_compress") + assert gate.quant_method.__class__.__name__ == "UnquantizedLinearMethod" + + +# --------------------------------------------------------------------------- +# conversion + apply round-trip on a real ReplicatedLinear (CPU) +# --------------------------------------------------------------------------- + + +def test_conversion_and_apply_match_dense_linear_within_tolerance(): + """End-to-end: load-time conversion then apply() dequantizes and matches. + + Uses ``retain_original_weight=False`` to exercise the purge path too. + """ + from fastvideo.layers.linear import ReplicatedLinear + from fastvideo.layers.quantization.int8_affine_config import convert_model_to_int8_affine + + torch.manual_seed(4) + cfg = INT8AffineConfig.for_minimax_h3(retain_original_weight=False) + layer = ReplicatedLinear(128, 256, bias=False, quant_config=cfg, + prefix="minimax_h3.transformer_blocks.0.attn.to_q") + with torch.no_grad(): + layer.weight.copy_(torch.randn(256, 128)) + + convert_model_to_int8_affine(layer) + assert layer._int8_affine_codes.dtype == torch.uint8 + assert layer._int8_affine_codes.shape == (256, 128) + assert layer._int8_affine_scales.shape == (256, 2) + assert layer.weight is None, "retain_original_weight=False should purge the bf16 weight" + + x = torch.randn(4, 128, dtype=torch.bfloat16) + out, _ = layer(x) + assert out.shape == (4, 256) + assert torch.isfinite(out).all() + + +def test_apply_falls_back_to_dense_under_grad(): + """A training step must see the master weight, not a frozen dequant copy.""" + from fastvideo.layers.linear import ReplicatedLinear + + torch.manual_seed(5) + cfg = INT8AffineConfig.for_minimax_h3() + layer = ReplicatedLinear(64, 64, bias=False, quant_config=cfg, + prefix="minimax_h3.transformer_blocks.0.attn.to_q") + with torch.no_grad(): + layer.weight.copy_(torch.randn(64, 64)) + x = torch.randn(2, 64, dtype=torch.bfloat16) + with torch.enable_grad(): + out, _ = layer(x) + assert out.shape == (2, 64) + assert not hasattr(layer, "_int8_affine_codes"), "grad-enabled forward must not quantize in place" diff --git a/fastvideo/tests/ops/quantization/test_int8_dispatch.py b/fastvideo/tests/ops/quantization/test_int8_dispatch.py new file mode 100644 index 0000000000..ca80166478 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_int8_dispatch.py @@ -0,0 +1,70 @@ +# SPDX-License-Identifier: Apache-2.0 +"""The loader must dispatch INT8 affine configs to their conversion hook. + +``_maybe_quantize_model`` walks the module tree and converts weights for +whichever quantization method it finds attached. It is an explicit +``isinstance`` chain, so a new quantization config is silently inert until it +gains a branch here -- the model keeps its dense weights and quietly produces +unquantized output with no error. + +This test pins the INT8Affine branch so that adding or reordering the chain +cannot silently drop it. +""" + +import unittest + +import torch +import torch.nn as nn + +from fastvideo.layers.quantization import int8_affine_config +from fastvideo.layers.quantization.int8_affine_config import ( + INT8AffineQuantizeMethod, + convert_model_to_int8_affine, +) +from fastvideo.models.loader import fsdp_load + + +class _FakeLinear(nn.Module): + """Minimal stand-in: a dense weight plus an attached quant method.""" + + def __init__(self, prefix: str = "minimax_h3.transformer_blocks.0.attn.to_q"): + super().__init__() + self.weight = nn.Parameter(torch.randn(64, 64, dtype=torch.bfloat16)) + self.quant_method = INT8AffineQuantizeMethod(layer_prefix=prefix) + + +class TestInt8Dispatch(unittest.TestCase): + + def setUp(self): + # The loader imports its conversion hook lazily *inside* the call, so it + # binds from the config module each time -- patch the source module. + self._calls = [] + self._orig = int8_affine_config.convert_model_to_int8_affine + + def _spy(model): + self._calls.append(model) + return self._orig(model) + + int8_affine_config.convert_model_to_int8_affine = _spy + self.addCleanup(setattr, int8_affine_config, "convert_model_to_int8_affine", self._orig) + + def test_int8_layer_triggers_conversion(self): + model = _FakeLinear() + fsdp_load._maybe_quantize_model(model) + self.assertEqual(len(self._calls), 1, "INT8AffineQuantizeMethod did not reach its conversion hook") + self.assertTrue(hasattr(model, "_int8_affine_codes"), "quantized buffers were not registered") + + def test_unquantized_model_is_untouched(self): + """A plain layer must not trigger any conversion.""" + model = nn.Linear(8, 8) + fsdp_load._maybe_quantize_model(model) + self.assertEqual(self._calls, [], "conversion ran on a model with no quant method") + + def test_conversion_is_exported_by_the_config_module(self): + """The hook the loader imports must be the one the config defines.""" + self.assertTrue(callable(convert_model_to_int8_affine)) + self.assertTrue(callable(self._orig)) + + +if __name__ == "__main__": + unittest.main() diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_h3_prefixes.py b/fastvideo/tests/ops/quantization/test_nvfp4_h3_prefixes.py new file mode 100644 index 0000000000..4a1063e596 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_nvfp4_h3_prefixes.py @@ -0,0 +1,161 @@ +# SPDX-License-Identifier: Apache-2.0 +"""NVFP4 layer-prefix selection: LTX-2 default, MiniMax-H3 opt-in. + +``NVFP4Config`` used to hardcode the LTX-2 layer paths, so the same config +handed to MiniMax-H3 attached no quant methods at all and the model ran dense +without an error. These CPU-only tests pin both halves of the fix: the default +still selects exactly the LTX-2 set, and an H3-configured instance selects the +300 block linears while refusing to quantize the VSA compression gate. + +No flashinfer and no CUDA are required — the tests use real +``ReplicatedLinear`` layers (which is where ``get_quant_method`` is called from) +with only ``NVFP4QuantizeMethod.__init__`` replaced, because the real one +allocates ``x_global_sf`` on ``cuda``. The same shim is used by +``test_nvfp4_purge.py``. +""" +from __future__ import annotations + +import sys + +import pytest +import torch + +import fastvideo.layers.quantization.nvfp4_config as nv +from fastvideo.layers.linear import ReplicatedLinear, UnquantizedLinearMethod + +_H3_BLOCK = "minimax_h3.transformer_blocks.{idx}.{suffix}" +_GATE_PREFIX = _H3_BLOCK.format(idx=0, suffix="attn.to_gate_compress") + + +@pytest.fixture(autouse=True) +def _cpu_quantize_method(monkeypatch): + """Build NVFP4QuantizeMethod without its cuda-allocated ``x_global_sf``.""" + + def _init(self, layer_prefix: str = ""): + self.weight_fp4 = None + self.weight_scale = None + self.x_global_sf = torch.tensor(1.0, dtype=torch.float32) + self.layer_prefix = layer_prefix + self._is_refine_only_layer = nv._is_ltx2_refine_only_prefix(layer_prefix) + self._retain_original_weights = None + + monkeypatch.setattr(nv.NVFP4QuantizeMethod, "__init__", _init) + + +def _linear(quant_config, prefix: str) -> ReplicatedLinear: + return ReplicatedLinear(8, 8, bias=False, quant_config=quant_config, prefix=prefix) + + +def test_default_config_keeps_the_ltx2_layer_set() -> None: + """No regression: the historical default is still the LTX-2 set.""" + config = nv.NVFP4Config() + assert config.layer_prefixes == nv._LTX2_NVFP4_LINEAR_PREFIXES + assert len(config.layer_prefixes) == 577 + assert config.exclude_prefixes == frozenset() + assert config.is_nvfp4_linear_prefix("ltx2.blocks.0.attn1.to_q") + assert config.is_nvfp4_linear_prefix("ltx2.adaln_single.linear") + assert not config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=0, suffix="attn.to_q")) + + +def test_h3_config_selects_the_300_block_linears() -> None: + config = nv.NVFP4Config.for_minimax_h3() + assert len(config.layer_prefixes) == 300 + assert config.layer_prefixes == nv.MINIMAX_H3_NVFP4_LINEAR_PREFIXES + # Every block, every suffix. + for idx in range(nv.MINIMAX_H3_NUM_LAYERS): + for suffix in nv.MINIMAX_H3_BLOCK_LINEAR_SUFFIXES: + assert config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=idx, suffix=suffix)) + # Edge blocks are really in the set, and one past the last block is not. + assert config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=49, suffix="ff.fc_out")) + assert not config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=50, suffix="ff.fc_out")) + # Nothing outside the block set is selected. + assert not config.is_nvfp4_linear_prefix("minimax_h3.token_refiner.blocks.0.attn.to_q") + assert not config.is_nvfp4_linear_prefix("minimax_h3.proj_in") + assert not config.is_nvfp4_linear_prefix("minimax_h3.transformer_blocks.0.adaln_proj.linear") + assert not config.is_nvfp4_linear_prefix("ltx2.blocks.0.attn1.to_q") + + +def test_h3_config_excludes_the_vsa_gate() -> None: + """``attn.to_gate_compress`` must never be quantized.""" + config = nv.NVFP4Config.for_minimax_h3() + assert not config.is_nvfp4_linear_prefix(_GATE_PREFIX) + assert not config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=37, suffix="attn.to_gate_compress")) + assert _GATE_PREFIX not in config.layer_prefixes + assert config.exclude_prefixes == frozenset(nv.MINIMAX_H3_NVFP4_EXCLUDED_LINEAR_SUFFIXES) + + +def test_gate_is_excluded_even_if_a_caller_allowlists_it() -> None: + """The gate exclusion is unconditional, not just absent from the set. + + A caller who builds the prefix set with a glob (or hand-lists every linear + in the block) must still not get a quantized VSA gate. + """ + config = nv.NVFP4Config(layer_prefixes=frozenset({_GATE_PREFIX, _H3_BLOCK.format(idx=0, suffix="attn.to_q")})) + assert not config.is_nvfp4_linear_prefix(_GATE_PREFIX) + assert config.is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=0, suffix="attn.to_q")) + # The same holds when the caller passes no exclusion list at all. + assert _GATE_PREFIX not in nv._LTX2_NVFP4_LINEAR_PREFIXES + assert not nv.NVFP4Config(layer_prefixes=[_GATE_PREFIX]).is_nvfp4_linear_prefix(_GATE_PREFIX) + + +def test_exclude_prefixes_match_a_suffix_or_a_full_path() -> None: + full_path = "custom.blocks.3.attn.to_out" + config = nv.NVFP4Config(layer_prefixes=[full_path, "custom.other"], exclude_prefixes=["attn.to_out"]) + assert config.is_nvfp4_linear_prefix("custom.other") + assert not config.is_nvfp4_linear_prefix(full_path) + # A dot boundary is required, so a sibling name is not excluded by accident. + nested = nv.NVFP4Config(layer_prefixes=["custom.cross_ff.fc_in"], exclude_prefixes=["ff.fc_in"]) + assert nested.is_nvfp4_linear_prefix("custom.cross_ff.fc_in") + + +def test_from_config_round_trips_layer_prefixes() -> None: + config = nv.NVFP4Config.from_config({ + "layer_profile": "base", + "layer_prefixes": ["minimax_h3.transformer_blocks.0.attn.to_q"], + "exclude_prefixes": ["attn.to_gate_compress"], + }) + assert config.layer_profile == "base" + assert config.layer_prefixes == frozenset({"minimax_h3.transformer_blocks.0.attn.to_q"}) + assert config.exclude_prefixes == frozenset({"attn.to_gate_compress"}) + # Absent keys keep the LTX-2 default. + assert nv.NVFP4Config.from_config({}).layer_prefixes == nv._LTX2_NVFP4_LINEAR_PREFIXES + + +def test_for_minimax_h3_forwards_kwargs() -> None: + config = nv.NVFP4Config.for_minimax_h3(retain_original_weights=True) + assert config.retain_original_weights is True + assert config.layer_prefixes == nv.MINIMAX_H3_NVFP4_LINEAR_PREFIXES + + +def test_get_quant_method_attaches_for_h3_and_skips_the_gate() -> None: + """End-to-end through ``ReplicatedLinear``, which is the real call site.""" + h3 = nv.NVFP4Config.for_minimax_h3() + selected = _linear(h3, _H3_BLOCK.format(idx=0, suffix="attn.to_q")) + assert isinstance(selected.quant_method, nv.NVFP4QuantizeMethod) + assert selected.quant_method.layer_prefix == _H3_BLOCK.format(idx=0, suffix="attn.to_q") + + gate = _linear(h3, _GATE_PREFIX) + assert type(gate.quant_method) is UnquantizedLinearMethod + + # The pre-fix bug, pinned: the default config attaches nothing on H3. + default = nv.NVFP4Config() + untagged = _linear(default, _H3_BLOCK.format(idx=0, suffix="attn.to_q")) + assert type(untagged.quant_method) is UnquantizedLinearMethod + # ... and still tags LTX-2. + ltx2 = _linear(default, "ltx2.blocks.0.attn1.to_q") + assert isinstance(ltx2.quant_method, nv.NVFP4QuantizeMethod) + + +def test_module_imports_without_flashinfer(monkeypatch) -> None: + """The H3 surface must import on hosts with no flashinfer (only the + kernels fail, at use time).""" + monkeypatch.setitem(sys.modules, "flashinfer", None) + # delitem (not a bare pop) so monkeypatch puts the original module object + # back on teardown: leaving a re-imported copy in sys.modules would give + # later tests a second, non-identical NVFP4Config class. + monkeypatch.delitem(sys.modules, "fastvideo.layers.quantization.nvfp4_config", raising=False) + import importlib + + reloaded = importlib.import_module("fastvideo.layers.quantization.nvfp4_config") + assert len(reloaded.MINIMAX_H3_NVFP4_LINEAR_PREFIXES) == 300 + assert reloaded.NVFP4Config.for_minimax_h3().is_nvfp4_linear_prefix(_H3_BLOCK.format(idx=0, suffix="attn.to_k")) diff --git a/fastvideo/tests/ops/quantization/test_nvfp4_sidecar.py b/fastvideo/tests/ops/quantization/test_nvfp4_sidecar.py new file mode 100644 index 0000000000..3161976123 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_nvfp4_sidecar.py @@ -0,0 +1,334 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU-only round-trip tests for the compact NVFP4 checkpoint sidecar. + +``convert_model_to_nvfp4`` registers its packed tensors with +``persistent=False``, so a saved checkpoint carries dense bf16 weights only and +the FP4 payload is rebuilt at every load. ``save_nvfp4_checkpoint`` / +``load_nvfp4_checkpoint`` write and restore that payload directly. + +flashinfer, the FP4 GEMM and ``NVFP4QuantizeMethod.__init__``'s cuda allocation +are all stubbed, so this runs on any host. The conversion, the buffer +registration and the retention policy under test are the real ones. +""" +from __future__ import annotations + +import json +import logging + +import pytest +import torch +import torch.nn as nn +from safetensors import safe_open +from safetensors.torch import save_file + +import fastvideo.layers.quantization.nvfp4_config as nv +from fastvideo.layers.linear import ReplicatedLinear + +_BLOCK = "minimax_h3.transformer_blocks.{idx}" +_IN_DIM = 8 +_OUT_DIM = 4 + + +@pytest.fixture(autouse=True) +def _stub_flashinfer(monkeypatch): + """Patch the three flashinfer touch points the real path would use.""" + import types + + fake_sf_layout = types.SimpleNamespace(layout_128x4=None) + monkeypatch.setattr(nv, "_require_flashinfer", lambda: (fake_sf_layout, None, None)) + monkeypatch.setattr(nv, "_nvfp4_quantize", _fake_quantize) + + def _init(self, layer_prefix: str = ""): + self.weight_fp4 = None + self.weight_scale = None + self.x_global_sf = torch.tensor(1.0, dtype=torch.float32) + self.layer_prefix = layer_prefix + self._is_refine_only_layer = nv._is_ltx2_refine_only_prefix(layer_prefix) + self._retain_original_weights = None + + monkeypatch.setattr(nv.NVFP4QuantizeMethod, "__init__", _init) + + +def _fake_quantize(weight, global_sf, sfLayout=None, do_shuffle=False): + """Deterministic stand-in for ``nvfp4_quantize`` with the real shapes.""" + out_dim, in_dim = weight.shape[0], weight.shape[-1] + packed = torch.arange(out_dim * ((in_dim + 1) // 2), dtype=torch.uint8).view(out_dim, (in_dim + 1) // 2) + scales = torch.arange(out_dim * ((in_dim + 15) // 16), dtype=torch.uint8).view(out_dim, (in_dim + 15) // 16) + return packed, scales + + +def _module() -> nn.Module: + """An empty submodule, so FQNs can be assembled explicitly.""" + return nn.Module() + + +def _build_model(*, num_blocks: int = 2, seed: int = 0, quant_config=None) -> nn.Module: + """A miniature H3-shaped DiT: ``minimax_h3.transformer_blocks.{i}.{attn.to_q,ff.fc_in}``.""" + config = quant_config if quant_config is not None else nv.NVFP4Config.for_minimax_h3() + generator = torch.Generator().manual_seed(seed) + root = _module() + dit = _module() + blocks = nn.ModuleList() + for idx in range(num_blocks): + block = _module() + attn = _module() + attn.to_q = ReplicatedLinear(_IN_DIM, + _OUT_DIM, + bias=False, + quant_config=config, + prefix=f"{_BLOCK.format(idx=idx)}.attn.to_q") + block.attn = attn + ff = _module() + ff.fc_in = ReplicatedLinear(_IN_DIM, + _OUT_DIM, + bias=False, + quant_config=config, + prefix=f"{_BLOCK.format(idx=idx)}.ff.fc_in") + block.ff = ff + blocks.append(block) + dit.transformer_blocks = blocks + root.minimax_h3 = dit + # ``create_weights`` allocates uninitialized storage; give it real values so + # the global scale (and therefore _nvfp4_alpha) is meaningful. + for _, param in root.named_parameters(): + param.data.copy_(torch.randn(param.shape, generator=generator)) + return root + + +def _buffers(model: nn.Module) -> dict[str, torch.Tensor]: + return nv.nvfp4_sidecar_state_dict(model) + + +def _rewrite_sidecar(src, dst, *, metadata: dict | None = None, tensors: dict | None = None) -> str: + """Copy a sidecar, optionally replacing its manifest or one tensor.""" + with safe_open(src, framework="pt", device="cpu") as handle: + payload = {key: handle.get_tensor(key) for key in handle.keys()} + manifest = json.loads(handle.metadata()[nv._NVFP4_SIDECAR_METADATA_KEY]) + if metadata is not None: + manifest.update(metadata) + if tensors is not None: + payload.update(tensors) + save_file(payload, str(dst), metadata={nv._NVFP4_SIDECAR_METADATA_KEY: json.dumps(manifest)}) + return str(dst) + + +def test_convert_purges_weights_and_sidecar_shrinks_the_payload(tmp_path) -> None: + model = _build_model() + nv.convert_model_to_nvfp4(model) + # The always-FP4 layers are purged, so the sidecar is the only copy left. + assert model.minimax_h3.transformer_blocks[0].attn.to_q.weight is None + + path = tmp_path / "nvfp4.safetensors" + receipt = nv.save_nvfp4_checkpoint(model, path) + assert path.exists() + assert receipt["num_layers"] == 4 + assert receipt["num_tensors"] == 16 # 4 buffers x 4 layers + assert receipt["quantized_bytes"] == sum(t.numel() * t.element_size() for t in _buffers(model).values()) + assert receipt["quantized_bytes"] < receipt["dense_bfloat16_bytes"] + assert receipt["compression_ratio"] > 1.0 + + manifest = nv.read_nvfp4_sidecar_metadata(path) + assert manifest["format"] == nv._NVFP4_SIDECAR_FORMAT + assert manifest["version"] == nv._NVFP4_SIDECAR_VERSION + assert manifest["sf_layout"] == "layout_128x4" + assert manifest["do_shuffle"] is False + assert manifest["block_size"] == 16 + assert manifest["layers"][f"{_BLOCK.format(idx=1)}.ff.fc_in"] == [_OUT_DIM, _IN_DIM] + assert manifest["quant_prefixes"][f"{_BLOCK.format(idx=0)}.attn.to_q"] == f"{_BLOCK.format(idx=0)}.attn.to_q" + + +def test_sidecar_state_dict_uses_module_fqn_keys_and_cpu_tensors() -> None: + model = _build_model() + nv.convert_model_to_nvfp4(model) + state = _buffers(model) + key = f"{_BLOCK.format(idx=0)}.attn.to_q::{nv._NVFP4_SIDECAR_BUFFERS[0]}" + assert key in state + assert state[key].device.type == "cpu" + # Only the four registered buffers are serialized; nothing else leaks in. + assert {name for _, name in (k.split("::") for k in state)} == set(nv._NVFP4_SIDECAR_BUFFERS) + assert state[f"{_BLOCK.format(idx=0)}.attn.to_q::_nvfp4_weight"].dtype is torch.uint8 + + +def test_load_restores_buffers_without_reconverting(tmp_path) -> None: + """The restored tensors must be identical to a fresh conversion's.""" + source = _build_model(seed=1) + nv.convert_model_to_nvfp4(source) + expected = {key: value.clone() for key, value in _buffers(source).items()} + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + # A different model (different weights) restored purely from the sidecar. + target = _build_model(seed=2) + restored = nv.load_nvfp4_checkpoint(target, path) + assert restored == 4 + actual = _buffers(target) + assert set(actual) == set(expected) + for key, value in expected.items(): + assert torch.equal(actual[key], value), key + assert actual[key].dtype == value.dtype + # The dense bf16 weights are gone under the default retention policy, so the + # sidecar really is the only copy of the quantized weights. + assert target.minimax_h3.transformer_blocks[0].attn.to_q.weight is None + + +def test_load_does_not_need_flashinfer(monkeypatch, tmp_path) -> None: + """Serving a pre-quantized checkpoint must not require the FP4 kernels.""" + source = _build_model(seed=3) + nv.convert_model_to_nvfp4(source) + expected = {key: value.clone() for key, value in _buffers(source).items()} + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + def _boom(): + raise ImportError("NVFP4 quantization requires flashinfer.") + + monkeypatch.setattr(nv, "_require_flashinfer", _boom) + target = _build_model(seed=4) + assert nv.load_nvfp4_checkpoint(target, path) == 4 + assert all(torch.equal(_buffers(target)[key], value) for key, value in expected.items()) + + +def test_load_into_a_model_whose_dense_weights_were_never_loaded(tmp_path) -> None: + """The compact case: the checkpoint carries no bf16 weights at all.""" + source = _build_model(seed=5) + nv.convert_model_to_nvfp4(source) + expected = {key: value.clone() for key, value in _buffers(source).items()} + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + target = _build_model(seed=6) + for module in target.modules(): + if getattr(module, "quant_method", None) is not None and hasattr(module, "weight"): + module.register_parameter("weight", None) + assert nv.load_nvfp4_checkpoint(target, path) == 4 + assert all(torch.equal(_buffers(target)[key], value) for key, value in expected.items()) + + +def test_load_purges_dense_weights_unless_asked_not_to(tmp_path) -> None: + source = _build_model(seed=7) + nv.convert_model_to_nvfp4(source) + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + keep = _build_model(seed=8) + nv.load_nvfp4_checkpoint(keep, path, purge_dense_weights=False) + assert keep.minimax_h3.transformer_blocks[0].attn.to_q.weight is not None + + +def _sidecar_of(model: nn.Module, path) -> str: + nv.convert_model_to_nvfp4(model) + nv.save_nvfp4_checkpoint(model, path) + return str(path) + + +def test_layer_set_mismatch_strict_and_lenient(tmp_path, caplog) -> None: + source = _build_model(num_blocks=1, seed=10) + nv.convert_model_to_nvfp4(source) + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + wider = _build_model(num_blocks=3, seed=13) + bigger = _build_model(num_blocks=2, seed=11) + with pytest.raises(ValueError, match="missing from the sidecar"): + nv.load_nvfp4_checkpoint(bigger, path) + + with caplog.at_level(logging.WARNING): + assert nv.load_nvfp4_checkpoint(bigger, path, strict=False) == 2 + assert any("does not match this model" in record.message for record in caplog.records) + # The unmatched layer keeps whatever it had (nothing) rather than silently + # becoming a dense layer with a quant_method attached. + assert getattr(bigger.minimax_h3.transformer_blocks[1].attn.to_q, "_nvfp4_weight", None) is None + + # The other direction: a sidecar with layers this model does not have. + smaller = _build_model(num_blocks=1, seed=12) + with pytest.raises(ValueError, match="not in the model"): + nv.load_nvfp4_checkpoint(smaller, _sidecar_of(wider, tmp_path / "wide.safetensors")) + + +def test_layout_and_version_mismatches_are_fatal(tmp_path) -> None: + source = _build_model(seed=13) + nv.convert_model_to_nvfp4(source) + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + target = _build_model(seed=14) + bad_layout = _rewrite_sidecar(path, tmp_path / "layout.safetensors", metadata={"sf_layout": "layout_linear"}) + with pytest.raises(ValueError, match="sf_layout"): + nv.load_nvfp4_checkpoint(target, bad_layout) + # Never downgraded by strict=False: mis-read nibbles are silent corruption. + with pytest.raises(ValueError, match="sf_layout"): + nv.load_nvfp4_checkpoint(target, bad_layout, strict=False) + + bad_shuffle = _rewrite_sidecar(path, tmp_path / "shuffle.safetensors", metadata={"do_shuffle": True}) + with pytest.raises(ValueError, match="do_shuffle"): + nv.load_nvfp4_checkpoint(target, bad_shuffle) + + bad_version = _rewrite_sidecar(path, tmp_path / "version.safetensors", metadata={"version": 99}) + with pytest.raises(ValueError, match="version"): + nv.load_nvfp4_checkpoint(target, bad_version) + + bad_format = _rewrite_sidecar(path, tmp_path / "format.safetensors", metadata={"format": "something.else"}) + with pytest.raises(ValueError, match="format"): + nv.load_nvfp4_checkpoint(target, bad_format) + + +def test_tensor_shape_mismatch_is_rejected(tmp_path) -> None: + source = _build_model(seed=15) + nv.convert_model_to_nvfp4(source) + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + key = f"{_BLOCK.format(idx=0)}.attn.to_q::_nvfp4_weight" + bad = _rewrite_sidecar(path, tmp_path / "shape.safetensors", tensors={key: torch.zeros(3, 7, dtype=torch.uint8)}) + with pytest.raises(ValueError, match="expected"): + nv.load_nvfp4_checkpoint(_build_model(seed=16), bad) + + +def test_block_scales_padded_to_the_128_row_tile_are_accepted(tmp_path) -> None: + """``nvfp4_quantize`` returns block scales padded to the 128-row tile. + + ``_nvfp4_quantize`` narrows the packed weight back to the logical row count + but not the scales, so a model whose output dim is not a multiple of 128 + can legitimately hold a scale tensor with more rows than the weight. A + sidecar carrying that must load rather than fail shape validation. + """ + source = _build_model(seed=19) + nv.convert_model_to_nvfp4(source) + path = tmp_path / "nvfp4.safetensors" + nv.save_nvfp4_checkpoint(source, path) + + scale_key = f"{_BLOCK.format(idx=0)}.attn.to_q::_nvfp4_weight_scale" + padded = _rewrite_sidecar(path, + tmp_path / "padded.safetensors", + tensors={scale_key: torch.zeros(128, (_IN_DIM + 15) // 16, dtype=torch.uint8)}) + assert nv.load_nvfp4_checkpoint(_build_model(seed=20), padded) == 4 + + +def test_saving_an_unconverted_model_raises(tmp_path) -> None: + model = _build_model(seed=17) + with pytest.raises(RuntimeError, match="convert_model_to_nvfp4"): + nv.save_nvfp4_checkpoint(model, tmp_path / "nvfp4.safetensors") + + +def test_a_model_with_no_nvfp4_layers_raises_with_the_prefix_hint(tmp_path) -> None: + """The silent-dense failure mode must be reported, not ignored.""" + empty = _build_model(seed=18, quant_config=nv.NVFP4Config()) + with pytest.raises(RuntimeError, match="for_minimax_h3"): + nv.save_nvfp4_checkpoint(empty, tmp_path / "nvfp4.safetensors") + with pytest.raises(RuntimeError, match="for_minimax_h3"): + nv.load_nvfp4_checkpoint(empty, tmp_path / "nvfp4.safetensors") + + +def test_sidecar_path_helper(tmp_path) -> None: + assert nv.nvfp4_sidecar_path_for("/models/h3/transformer.safetensors") == "/models/h3/transformer.nvfp4.safetensors" + assert nv.nvfp4_sidecar_path_for("/models/h3") == "/models/h3.nvfp4.safetensors" + assert nv.nvfp4_sidecar_path_for(str(tmp_path)) == str(tmp_path / "nvfp4.safetensors") + + +def test_read_metadata_rejects_a_foreign_file(tmp_path) -> None: + from safetensors.torch import save_file as _save + + plain = tmp_path / "plain.safetensors" + _save({"w": torch.zeros(2, 2)}, str(plain)) + with pytest.raises(ValueError, match="not a FastVideo NVFP4 sidecar"): + nv.read_nvfp4_sidecar_metadata(plain) diff --git a/fastvideo/tests/ops/quantization/test_quant_param_allowlist.py b/fastvideo/tests/ops/quantization/test_quant_param_allowlist.py new file mode 100644 index 0000000000..86e27077c5 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_quant_param_allowlist.py @@ -0,0 +1,139 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU-only contract tests for the ``fsdp_load`` zero-init allowlist. + +A quantization config registers scale tensors that no checkpoint carries +(``AbsMaxFP8`` -> ``scale_weight`` / ``scale_input``; see +``fastvideo/layers/quantization/absmax_fp8.py``). ``fsdp_load`` rejects every +model parameter that is absent from the incoming state dict, so enabling +``engine.quantization.transformer_quant`` used to abort the load with:: + + Unsupported new parameter: transformer_blocks.0.attn.to_out.scale_input + +These tests pin both halves of that contract, on CPU and without any +distributed process group: + +* quant scale parameters are accepted and zero-initialized, +* a parameter that is genuinely missing from the checkpoint still raises, so + the allowlist is not a blanket exemption. + +See ``docs/quantization/loader_quant_params.md`` before adding a quantization +config that registers new parameters. +""" + +import unittest + +import torch +import torch.nn as nn + +from fastvideo.layers.quantization.absmax_fp8 import AbsMaxFP8LinearMethod +from fastvideo.models.loader.fsdp_load import ( + ALLOWED_NEW_PARAM_PATTERNS, + is_allowed_new_param, + load_model_from_full_model_state_dict, +) + +DTYPE = torch.float32 +# The FQN from the original failure, used verbatim so the regression is +# recognizable if it ever comes back. +REPORTED_FQN = "transformer_blocks.0.attn.to_out.scale_input" + + +class _QuantToOut(nn.Module): + """A module holding an AbsMaxFP8 linear's parameters under their real names.""" + + def __init__(self, in_features: int = 3, out_features: int = 2) -> None: + super().__init__() + AbsMaxFP8LinearMethod().create_weights( + self, + input_size_per_partition=in_features, + output_partition_sizes=[out_features], + input_size=in_features, + output_size=out_features, + params_dtype=DTYPE, + ) + + +class _Block(nn.Module): + + def __init__(self) -> None: + super().__init__() + self.attn = _Attn() + + +class _Attn(nn.Module): + + def __init__(self) -> None: + super().__init__() + self.to_out = _QuantToOut() + + +def _model() -> nn.Module: + model = nn.Module() + model.add_module("transformer_blocks", nn.ModuleList([_Block()])) + return model + + +def _load(model: nn.Module, checkpoint: dict[str, torch.Tensor]): + return load_model_from_full_model_state_dict( + model, + ((name, tensor) for name, tensor in checkpoint.items()), + device=torch.device("cpu"), + param_dtype=DTYPE, + strict=False, + cpu_offload=False, + param_names_mapping=lambda name: (name, None, None), + training_mode=False, + ) + + +class TestQuantParamAllowlist(unittest.TestCase): + + def test_reported_fqn_is_allowed(self): + # The exact name from the failed AbsMaxFP8 run. + self.assertTrue(is_allowed_new_param(REPORTED_FQN)) + + def test_quant_scale_params_are_zero_initialized(self): + # The checkpoint holds only `weight`; `scale_weight` / `scale_input` + # exist in the model but never in the checkpoint. + model = _model() + checkpoint = {"transformer_blocks.0.attn.to_out.weight": torch.ones(2, 3, dtype=DTYPE)} + _load(model, checkpoint) + + to_out = model.transformer_blocks[0].attn.to_out + # Per-tensor scales for a plain (non-merged) linear, zero-initialized. + self.assertEqual(to_out.scale_weight.shape, (1, )) + self.assertEqual(to_out.scale_input.shape, (1, )) + self.assertTrue(torch.equal(to_out.scale_weight, torch.zeros(1, dtype=DTYPE))) + self.assertTrue(torch.equal(to_out.scale_input, torch.zeros(1, dtype=DTYPE))) + self.assertTrue(torch.equal(to_out.weight, torch.ones(2, 3, dtype=DTYPE))) + + def test_missing_real_weight_still_raises(self): + # No quant scale involved: a plain model weight absent from the + # checkpoint is a mapping bug and must keep failing loudly. + model = _model() + with self.assertRaisesRegex(ValueError, "is not supported"): + _load(model, {}) + + def test_bare_scale_param_is_not_admitted(self): + # Guard against widening the allowlist to a bare "scale" token: a + # learned parameter that merely ends in `scale` must stay rejected. + self.assertFalse(is_allowed_new_param("transformer_blocks.0.attn.to_out.scale")) + + def test_every_absmax_fp8_registered_param_is_allowed(self): + # Derived from the quant method itself, so a new scale tensor added by + # AbsMaxFP8 fails here instead of at checkpoint-load time. + to_out = _QuantToOut() + new_params = [name for name, _ in to_out.named_parameters() if name != "weight"] + self.assertEqual(sorted(new_params), ["scale_input", "scale_weight"]) + for name in new_params: + fqn = f"transformer_blocks.0.attn.to_out.{name}" + self.assertTrue(is_allowed_new_param(fqn), f"{fqn} is not in {ALLOWED_NEW_PARAM_PATTERNS}") + + def test_existing_attention_patterns_still_allowed(self): + # The pre-existing entries must survive the change. + for name in ("transformer_blocks.0.attn.to_gate_compress.weight", "blocks.0.attn1.attn_impl.proj_l.weight"): + self.assertTrue(is_allowed_new_param(name)) + + +if __name__ == "__main__": + unittest.main() diff --git a/fastvideo/tests/ops/quantization/test_quant_sidecar_roundtrip.py b/fastvideo/tests/ops/quantization/test_quant_sidecar_roundtrip.py new file mode 100644 index 0000000000..910fa3e516 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_quant_sidecar_roundtrip.py @@ -0,0 +1,697 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU-only round-trip tests for the INT8 affine and W4A16 checkpoint sidecars. + +``convert_model_to_int8_affine`` / ``convert_model_to_w4a16`` register their +quantized tensors with ``persistent=False``, so a saved checkpoint carries +dense bf16 weights only and the low-bit payload is rebuilt at every load — by +re-running a conversion that starts from those dense weights. On the hardware +these lanes target (24-48 GB Ada parts, a 32 GB RTX 5090) H3's bf16 DiT does +not fit, so that rebuild is impossible and a pre-quantized checkpoint is the +only thing that can be served. ``save_*_checkpoint`` / ``load_*_checkpoint`` +write and restore that payload directly. + +No CUDA, no GPU, no flashinfer: INT8 affine and W4A16 are pure PyTorch, and +every test here runs on a laptop. The quantizers, the buffer registration, the +manifest and the load-time validation under test are the real ones. + +The load-bearing assertions are the ones about *silent* corruption: +``torch.equal`` (not ``allclose``) for the round trip, an exact dtype check on +the codes, and a shape check per tensor — a mis-read or mis-unpacked code +buffer produces plausible-looking garbage with no error anywhere. +""" + +from __future__ import annotations + +import json +import logging +import sys +from collections.abc import Callable +from typing import Any, NamedTuple + +import pytest +import torch +import torch.nn as nn +from safetensors import safe_open +from safetensors.torch import save_file + +from fastvideo.layers.linear import ReplicatedLinear +from fastvideo.layers.quantization import int8_affine_config as i8 +from fastvideo.layers.quantization import w4a16_config as w4 + +IN_DIM = 128 +OUT_DIM = 32 +GROUP_SIZE = 64 +_BLOCK = "minimax_h3.transformer_blocks.{idx}" +_Q = f"{_BLOCK.format(idx=0)}.attn.to_q" +_FF = f"{_BLOCK.format(idx=1)}.ff.fc_in" + + +class _Scheme(NamedTuple): + """Everything the two lanes do not share, addressed by name.""" + + name: str + module: Any + make_config: Callable[..., Any] + buffers: tuple[str, ...] + format_name: str + metadata_key: str + suffix: str + dir_name: str + convert: Callable[[nn.Module], None] + state_dict: Callable[[nn.Module], dict[str, torch.Tensor]] + save: Callable[..., dict[str, Any]] + load: Callable[..., int] + read_metadata: Callable[[Any], dict[str, Any]] + path_for: Callable[[Any], str] + # The sidecar's error text when a model has no tagged layers at all. + wiring_hint: str + + +INT8 = _Scheme( + name="int8_affine", + module=i8, + make_config=i8.INT8AffineConfig, + buffers=i8._INT8_AFFINE_SIDECAR_BUFFERS, + format_name=i8._INT8_AFFINE_SIDECAR_FORMAT, + metadata_key=i8._INT8_AFFINE_SIDECAR_METADATA_KEY, + suffix=i8.INT8_AFFINE_SIDECAR_SUFFIX, + dir_name=i8.INT8_AFFINE_DIR_SIDECAR_NAME, + convert=i8.convert_model_to_int8_affine, + state_dict=i8.int8_affine_sidecar_state_dict, + save=i8.save_int8_affine_checkpoint, + load=i8.load_int8_affine_checkpoint, + read_metadata=i8.read_int8_affine_sidecar_metadata, + path_for=i8.int8_affine_sidecar_path_for, + wiring_hint="for_minimax_h3", +) + +W4A16 = _Scheme( + name="w4a16", + module=w4, + make_config=w4.W4A16Config, + buffers=w4._W4A16_SIDECAR_BUFFERS, + format_name=w4._W4A16_SIDECAR_FORMAT, + metadata_key=w4._W4A16_SIDECAR_METADATA_KEY, + suffix=w4.W4A16_SIDECAR_SUFFIX, + dir_name=w4.W4A16_DIR_SIDECAR_NAME, + convert=w4.convert_model_to_w4a16, + state_dict=w4.w4a16_sidecar_state_dict, + save=w4.save_w4a16_checkpoint, + load=w4.load_w4a16_checkpoint, + read_metadata=w4.read_w4a16_sidecar_metadata, + path_for=w4.w4a16_sidecar_path_for, + wiring_hint="for_minimax_h3", +) + +_ALL = (INT8, W4A16) +_by_name = pytest.mark.parametrize("scheme", _ALL, ids=[s.name for s in _ALL]) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _build(scheme: _Scheme, *, num_blocks: int = 2, seed: int = 0, config=None) -> nn.Module: + """A miniature H3-shaped DiT: ``minimax_h3.transformer_blocks.{i}.{attn.to_q,ff.fc_in}``.""" + cfg = config if config is not None else scheme.make_config() + generator = torch.Generator().manual_seed(seed) + root = nn.Module() + dit = nn.Module() + blocks = nn.ModuleList() + for idx in range(num_blocks): + block = nn.Module() + attn = nn.Module() + attn.to_q = ReplicatedLinear(IN_DIM, + OUT_DIM, + bias=False, + quant_config=cfg, + prefix=f"{_BLOCK.format(idx=idx)}.attn.to_q") + ff = nn.Module() + ff.fc_in = ReplicatedLinear(IN_DIM, + OUT_DIM, + bias=False, + quant_config=cfg, + prefix=f"{_BLOCK.format(idx=idx)}.ff.fc_in") + block.attn = attn + block.ff = ff + blocks.append(block) + dit.transformer_blocks = blocks + root.minimax_h3 = dit + # ``create_weights`` allocates uninitialized storage; real values are what + # make the codes meaningful (and the manifest's byte counts non-trivial). + for _, param in root.named_parameters(): + param.data.copy_(torch.randn(param.shape, generator=generator)) + return root + + +def _tagged(model: nn.Module) -> dict[str, nn.Module]: + return {fqn: mod for fqn, mod in model.named_modules() if getattr(mod, "quant_method", None) is not None} + + +def _rewrite_sidecar(scheme: _Scheme, src, dst, *, metadata: dict | None = None, tensors: dict | None = None) -> str: + """Copy a sidecar, optionally replacing its manifest or some tensors.""" + with safe_open(src, framework="pt", device="cpu") as handle: + payload = {key: handle.get_tensor(key) for key in handle.keys()} + manifest = json.loads(handle.metadata()[scheme.metadata_key]) + if metadata is not None: + manifest.update(metadata) + if tensors is not None: + payload.update(tensors) + save_file(payload, str(dst), metadata={scheme.metadata_key: json.dumps(manifest)}) + return str(dst) + + +def _sidecar_of(scheme: _Scheme, model: nn.Module, path) -> str: + scheme.convert(model) + scheme.save(model, path) + return str(path) + + +# --------------------------------------------------------------------------- +# (d) Why this function has to exist: state_dict() does not carry the payload +# --------------------------------------------------------------------------- + + +@_by_name +def test_state_dict_does_not_carry_the_quantized_buffers(scheme: _Scheme) -> None: + """The whole reason for the sidecar: the buffers are non-persistent. + + A plain ``state_dict()`` is the dense bf16 weights and nothing else, so a + checkpoint written the ordinary way cannot be served on a host that cannot + hold the dense weights and re-convert. + """ + model = _build(scheme) + scheme.convert(model) + + keys = list(model.state_dict()) + assert keys, "the dense weights are persistent, so this must not be empty" + for key in keys: + for buffer_name in scheme.buffers: + assert buffer_name not in key, f"{key} leaked a non-persistent buffer" + assert set(keys) == {f"{fqn}.weight" for fqn in _tagged(model)} + # Every tagged layer did convert, so it is the persistence flag, not an + # empty conversion, that keeps the payload out of the checkpoint. + for fqn, mod in _tagged(model).items(): + assert getattr(mod, scheme.buffers[0]) is not None, fqn + assert len(scheme.state_dict(model)) == 4 * len(scheme.buffers) + + +@_by_name +def test_sidecar_state_dict_uses_module_fqn_keys_and_cpu_tensors(scheme: _Scheme) -> None: + model = _build(scheme) + scheme.convert(model) + state = scheme.state_dict(model) + + assert f"{_Q}::{scheme.buffers[0]}" in state + assert all(value.device.type == "cpu" for value in state.values()) + # Only the registered buffers are serialized; nothing else leaks in. + assert {key.split("::", 1)[1] for key in state} == set(scheme.buffers) + assert state[f"{_Q}::{scheme.buffers[0]}"].dtype is torch.uint8 + for name in scheme.buffers[1:]: + assert state[f"{_Q}::{name}"].dtype is torch.float32 + + +def test_int8_codes_are_uint8_because_they_exceed_the_int8_range() -> None: + """Affine bits=8 codes span [0, 255]; int8 storage would wrap 255 to -1. + + Pinned on a crafted weight rather than a random one so the assertion is + about the scheme, not about a lucky seed. + """ + model = _build(INT8) + linear = _tagged(model)[_Q] + weight = torch.full((OUT_DIM, IN_DIM), -1.0) + weight[:, 0] = 3.0 # one large positive value per row + linear.weight.data.copy_(weight) + INT8.convert(model) + + codes = linear._int8_affine_codes + assert codes.dtype is torch.uint8 + assert codes.max().item() == 255 + # The same bytes read as int8 are a different number: 255 wraps to -1. A + # load that accepted int8 codes would dequantize a wrong weight with no + # error anywhere, which is why the dtype is asserted rather than cast. + as_int8 = codes.to(torch.int8) + assert as_int8.min().item() == -1 + assert not torch.equal(as_int8.to(torch.int64), codes.to(torch.int64)) + + +# --------------------------------------------------------------------------- +# (a) Round trip into a fresh module is bit-identical +# --------------------------------------------------------------------------- + + +@_by_name +def test_save_then_load_into_a_fresh_module_is_bit_identical(scheme: _Scheme, tmp_path) -> None: + source = _build(scheme, seed=1) + scheme.convert(source) + expected = {key: value.clone() for key, value in scheme.state_dict(source).items()} + + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + # A different model (different dense weights) restored purely from the + # sidecar, with no conversion ever running on it. + target = _build(scheme, seed=2) + assert scheme.load(target, path) == 4 + actual = scheme.state_dict(target) + + assert set(actual) == set(expected) + for key, value in expected.items(): + assert actual[key].dtype == value.dtype, key + assert torch.equal(actual[key], value), f"{key} is not bit-identical (allclose would hide this)" + # The dense weights were never touched: the target keeps its own. + assert not torch.equal(target.minimax_h3.transformer_blocks[0].attn.to_q.weight, + source.minimax_h3.transformer_blocks[0].attn.to_q.weight) + + +@_by_name +def test_loaded_buffers_reproduce_the_source_forward(scheme: _Scheme, tmp_path) -> None: + """A loaded layer must compute what the converted source layer computes. + + This is the end-to-end payoff: for W4A16 it only holds if the load also + restored ``_w4a16_weight_shape``, which is a plain attribute and not a + buffer, so nothing else would carry it across. + """ + source = _build(scheme, seed=3) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + target = _build(scheme, seed=4) + scheme.load(target, path) + + x = torch.randn(2, IN_DIM, generator=torch.Generator().manual_seed(5)) + with torch.no_grad(): + want, _ = _tagged(source)[_Q](x) + got, _ = _tagged(target)[_Q](x) + assert torch.equal(got, want) + + +def test_w4a16_load_restores_the_logical_weight_shape_attribute(tmp_path) -> None: + """``_w4a16_weight_shape`` is not a buffer, so the manifest must carry it.""" + source = _build(W4A16, seed=6) + W4A16.convert(source) + path = tmp_path / "w4a16.safetensors" + W4A16.save(source, path) + + target = _build(W4A16, seed=7) + linear = _tagged(target)[_Q] + assert getattr(linear, "_w4a16_weight_shape", None) is None + W4A16.load(target, path) + assert tuple(linear._w4a16_weight_shape) == (OUT_DIM, IN_DIM) + # Without it, apply() could not dequantize: the packed shape (32, 64) does + # not encode the logical (32, 128). + assert linear._w4a16_codes.shape == (OUT_DIM, IN_DIM // 2) + + +@_by_name +def test_save_and_load_with_no_dense_weights_at_all(scheme: _Scheme, tmp_path) -> None: + """The target-hardware case: the bf16 weights are never present. + + With ``retain_original_weight=False`` the converted model keeps no dense + copy, and on a card that cannot hold them the loading model has none + either (``weight`` is ``None``). Save and load must both work from the + quantized buffers alone — this is the path that exercises the manifest's + per-layer ``weight_shape`` for the shape it cannot read off a weight. + """ + config = scheme.make_config(retain_original_weight=False) + source = _build(scheme, seed=34, config=config) + scheme.convert(source) + assert _tagged(source)[_Q].weight is None + + path = tmp_path / "sidecar.safetensors" + receipt = scheme.save(source, path) + assert receipt["num_layers"] == 4 + expected = {key: value.clone() for key, value in scheme.state_dict(source).items()} + + target = _build(scheme, seed=35, config=scheme.make_config(retain_original_weight=False)) + for mod in _tagged(target).values(): + mod.register_parameter("weight", None) + assert scheme.load(target, path) == 4 + + actual = scheme.state_dict(target) + assert set(actual) == set(expected) + assert all(torch.equal(actual[key], value) for key, value in expected.items()) + # And the model is actually usable: no dense weight anywhere, yet forward + # runs off the restored buffers. + x = torch.randn(2, IN_DIM, generator=torch.Generator().manual_seed(36)) + with torch.no_grad(): + assert torch.equal(_tagged(target)[_Q](x)[0], _tagged(source)[_Q](x)[0]) + + +def test_w4a16_bits_8_sidecar_uses_the_unpacked_code_shape(tmp_path) -> None: + """``bits=8`` stores one code per byte, so the code shape is the weight shape. + + The manifest and the load-time shape check both have to follow ``bits`` + rather than assume the 4-bit packed layout. + """ + config = W4A16.make_config(bits=8) + source = _build(W4A16, seed=42, config=config) + W4A16.convert(source) + assert _tagged(source)[_Q]._w4a16_codes.shape == (OUT_DIM, IN_DIM) + + path = tmp_path / "w4a16_bits8.safetensors" + W4A16.save(source, path) + manifest = W4A16.read_metadata(path) + assert manifest["bits"] == 8 + assert manifest["layers"][_Q]["tensors"][W4A16.buffers[0]] == [OUT_DIM, IN_DIM] + + target = _build(W4A16, seed=43, config=W4A16.make_config(bits=8)) + assert W4A16.load(target, path) == 4 + assert torch.equal(W4A16.state_dict(target)[f"{_Q}::{W4A16.buffers[0]}"], + W4A16.state_dict(source)[f"{_Q}::{W4A16.buffers[0]}"]) + x = torch.randn(2, IN_DIM, generator=torch.Generator().manual_seed(44)) + with torch.no_grad(): + assert torch.equal(_tagged(target)[_Q](x)[0], _tagged(source)[_Q](x)[0]) + + +def test_w4a16_codes_preserve_the_low_nibble_first_packing(tmp_path) -> None: + """The packed bytes must round-trip the packer's nibble order exactly. + + The convention is the module's own (``_pack_4bit``): the low nibble is the + lower K index. Re-deriving the expected byte from the source layer's own + quantizer output is what makes this a check of the *serialized* bytes + rather than of the packing function. + """ + source = _build(W4A16, seed=8) + W4A16.convert(source) + linear = _tagged(source)[_Q] + codes = linear._w4a16_codes + weight = linear.weight.detach().float() + unpacked = w4._unpack_4bit(codes) + expected = w4._pack_4bit(unpacked) + assert torch.equal(codes, expected) + assert codes.shape == (OUT_DIM, IN_DIM // 2) + assert unpacked.shape == (OUT_DIM, IN_DIM) + # The unpacked stream is the quantizer's own code stream, not a re-quantize. + raw, _, _ = w4.w4a16_quantize(weight, group_size=GROUP_SIZE, bits=4) + assert torch.equal(unpacked, w4._unpack_4bit(raw)) + + path = tmp_path / "w4a16.safetensors" + W4A16.save(source, path) + target = _build(W4A16, seed=9) + W4A16.load(target, path) + assert torch.equal(_tagged(target)[_Q]._w4a16_codes, codes) + + +# --------------------------------------------------------------------------- +# (c) The manifest round-trips +# --------------------------------------------------------------------------- + + +@_by_name +def test_manifest_round_trips_scheme_and_layer_inventory(scheme: _Scheme, tmp_path) -> None: + model = _build(scheme) + scheme.convert(model) + path = tmp_path / "sidecar.safetensors" + receipt = scheme.save(model, path) + + manifest = scheme.read_metadata(path) + assert manifest["format"] == scheme.format_name + assert manifest["version"] == 1 + assert manifest["group_size"] == GROUP_SIZE + assert manifest["bits"] == (8 if scheme is INT8 else 4) + assert manifest["num_layers"] == 4 + # The quantized module fqns are the manifest's layer keys. + assert set(manifest["layers"]) == {_Q, _FF, f"{_BLOCK.format(idx=0)}.ff.fc_in", + f"{_BLOCK.format(idx=1)}.attn.to_q"} + assert manifest["quant_prefixes"][_Q] == _Q + assert manifest["model_class"] == "Module" + + entry = manifest["layers"][_Q] + assert entry["weight_shape"] == [OUT_DIM, IN_DIM] + assert entry["group_size"] == GROUP_SIZE + assert entry["bits"] == (8 if scheme is INT8 else 4) + assert set(entry["tensors"]) == set(scheme.buffers) + + # The declared per-layer shapes are the shapes actually written. + with safe_open(path, framework="pt", device="cpu") as handle: + for fqn, layer in manifest["layers"].items(): + for name, shape in layer["tensors"].items(): + assert list(handle.get_tensor(f"{fqn}::{name}").shape) == shape + + # The receipt adds the byte accounting, and the sidecar is the smaller copy. + assert receipt["num_layers"] == 4 + assert receipt["num_tensors"] == 4 * len(scheme.buffers) + assert receipt["quantized_bytes"] == sum(t.numel() * t.element_size() for t in scheme.state_dict(model).values()) + assert receipt["quantized_bytes"] < receipt["dense_bfloat16_bytes"] + assert receipt["compression_ratio"] > 1.0 + + +@_by_name +def test_save_logs_a_receipt_with_the_module_count_and_bytes(scheme: _Scheme, tmp_path, caplog) -> None: + model = _build(scheme) + scheme.convert(model) + with caplog.at_level(logging.INFO): + receipt = scheme.save(model, tmp_path / "sidecar.safetensors") + message = "\n".join(record.message for record in caplog.records) + assert "4 quantized modules" in message + assert str(receipt["quantized_bytes"]) in message + assert "dense bf16" in message + + +@_by_name +def test_extra_metadata_is_merged_into_the_manifest(scheme: _Scheme, tmp_path) -> None: + model = _build(scheme) + scheme.convert(model) + path = tmp_path / "sidecar.safetensors" + scheme.save(model, path, extra_metadata={"source_checkpoint": "h3-bf16"}) + assert scheme.read_metadata(path)["source_checkpoint"] == "h3-bf16" + + +# --------------------------------------------------------------------------- +# (b) Corrupt or foreign payloads are rejected loudly +# --------------------------------------------------------------------------- + + +@_by_name +def test_wrong_code_dtype_is_rejected_at_load(scheme: _Scheme, tmp_path) -> None: + """Codes cast to int8 must fail, not silently dequantize to garbage.""" + source = _build(scheme, seed=10) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + key = f"{_Q}::{scheme.buffers[0]}" + with safe_open(path, framework="pt", device="cpu") as handle: + bad_codes = handle.get_tensor(key).to(torch.int8) + bad = _rewrite_sidecar(scheme, path, tmp_path / "int8codes.safetensors", tensors={key: bad_codes}) + + with pytest.raises(ValueError, match="dtype"): + scheme.load(_build(scheme, seed=11), bad) + # Never downgraded: a wrong-typed code buffer is silent corruption. + with pytest.raises(ValueError, match="dtype"): + scheme.load(_build(scheme, seed=11), bad, strict=False) + + +@_by_name +def test_wrong_float_dtype_is_rejected_at_load(scheme: _Scheme, tmp_path) -> None: + """A bf16 scale store is not a bit-exact restore of the fp32 constants.""" + source = _build(scheme, seed=12) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + key = f"{_Q}::{scheme.buffers[1]}" + with safe_open(path, framework="pt", device="cpu") as handle: + bad_scales = handle.get_tensor(key).to(torch.bfloat16) + bad = _rewrite_sidecar(scheme, path, tmp_path / "bf16scales.safetensors", tensors={key: bad_scales}) + with pytest.raises(ValueError, match="dtype"): + scheme.load(_build(scheme, seed=13), bad) + + +@_by_name +def test_wrong_shape_is_rejected_at_load(scheme: _Scheme, tmp_path) -> None: + source = _build(scheme, seed=14) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + key = f"{_Q}::{scheme.buffers[0]}" + bad = _rewrite_sidecar(scheme, + path, + tmp_path / "shape.safetensors", + tensors={key: torch.zeros(OUT_DIM, IN_DIM * 2, dtype=torch.uint8)}) + with pytest.raises(ValueError, match="expected"): + scheme.load(_build(scheme, seed=15), bad) + with pytest.raises(ValueError, match="expected"): + scheme.load(_build(scheme, seed=15), bad, strict=False) + + +@_by_name +def test_scale_shape_mismatch_is_rejected_at_load(scheme: _Scheme, tmp_path) -> None: + """A regrouped scales tensor would regroup every code in the row.""" + source = _build(scheme, seed=16) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + key = f"{_Q}::{scheme.buffers[1]}" + bad = _rewrite_sidecar(scheme, + path, + tmp_path / "scales.safetensors", + tensors={key: torch.zeros(OUT_DIM, IN_DIM, dtype=torch.float32)}) + with pytest.raises(ValueError, match="expected"): + scheme.load(_build(scheme, seed=17), bad) + + +@_by_name +def test_group_size_and_bits_mismatches_are_fatal(scheme: _Scheme, tmp_path) -> None: + """A different scheme means different dequantize arithmetic over the bytes.""" + source = _build(scheme, seed=18) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + other_group = scheme.make_config(group_size=GROUP_SIZE * 2) + with pytest.raises(ValueError, match="group_size"): + scheme.load(_build(scheme, seed=19, config=other_group), path) + with pytest.raises(ValueError, match="group_size"): + scheme.load(_build(scheme, seed=19, config=other_group), path, strict=False) + + other_bits = scheme.make_config(bits=4 if scheme is INT8 else 8) + with pytest.raises(ValueError, match="bits"): + scheme.load(_build(scheme, seed=19, config=other_bits), path) + + +@_by_name +def test_weight_shape_that_disagrees_with_the_model_is_fatal(scheme: _Scheme, tmp_path) -> None: + """A sidecar describing a differently-shaped layer must not load into this one.""" + source = _build(scheme, seed=20) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + layers = dict(scheme.read_metadata(path)["layers"]) + layers[_Q] = {**layers[_Q], "weight_shape": [OUT_DIM, IN_DIM * 2]} + reshaped = _rewrite_sidecar(scheme, path, tmp_path / "ws.safetensors", metadata={"layers": layers}) + with pytest.raises(ValueError, match="shape"): + scheme.load(_build(scheme, seed=21), reshaped) + + +@_by_name +def test_malformed_layer_entry_is_reported_not_a_key_error(scheme: _Scheme, tmp_path) -> None: + """A hand-edited manifest must fail with a message naming the layer.""" + source = _build(scheme, seed=37) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + layers = dict(scheme.read_metadata(path)["layers"]) + layers[_Q] = {key: value for key, value in layers[_Q].items() if key != "group_size"} + malformed = _rewrite_sidecar(scheme, path, tmp_path / "malformed.safetensors", metadata={"layers": layers}) + with pytest.raises(ValueError, match="malformed"): + scheme.load(_build(scheme, seed=38), malformed) + + +@_by_name +def test_format_and_version_mismatches_are_fatal(scheme: _Scheme, tmp_path) -> None: + source = _build(scheme, seed=22) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + bad_format = _rewrite_sidecar(scheme, path, tmp_path / "format.safetensors", metadata={"format": "other"}) + with pytest.raises(ValueError, match="format"): + scheme.load(_build(scheme, seed=23), bad_format) + + bad_version = _rewrite_sidecar(scheme, path, tmp_path / "version.safetensors", metadata={"version": 99}) + with pytest.raises(ValueError, match="version"): + scheme.load(_build(scheme, seed=23), bad_version) + + +@_by_name +def test_incomplete_layer_entry_is_rejected(scheme: _Scheme, tmp_path) -> None: + """A layer missing its code buffer must not silently become dense.""" + source = _build(scheme, seed=24) + scheme.convert(source) + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + with safe_open(path, framework="pt", device="cpu") as handle: + payload = {key: handle.get_tensor(key) for key in handle.keys()} + manifest = json.loads(handle.metadata()[scheme.metadata_key]) + del payload[f"{_Q}::{scheme.buffers[0]}"] + stripped = tmp_path / "stripped.safetensors" + save_file(payload, str(stripped), metadata={scheme.metadata_key: json.dumps(manifest)}) + + with pytest.raises(ValueError, match="incomplete"): + scheme.load(_build(scheme, seed=25), stripped) + + +@_by_name +def test_layer_set_mismatch_strict_and_lenient(scheme: _Scheme, tmp_path, caplog) -> None: + source = _build(scheme, num_blocks=1, seed=26) + path = tmp_path / "one.safetensors" + scheme.convert(source) + scheme.save(source, path) + + bigger = _build(scheme, num_blocks=2, seed=27) + with pytest.raises(ValueError, match="missing from the sidecar"): + scheme.load(bigger, path) + + with caplog.at_level(logging.WARNING): + assert scheme.load(bigger, path, strict=False) == 2 + assert any("does not match this model" in record.message for record in caplog.records) + # The unmatched layer keeps whatever it had (nothing) rather than becoming + # a dense layer with a quant_method attached. + assert getattr(bigger.minimax_h3.transformer_blocks[1].attn.to_q, scheme.buffers[0], None) is None + + wider = _build(scheme, num_blocks=3, seed=28) + with pytest.raises(ValueError, match="not in the model"): + scheme.load(_build(scheme, num_blocks=1, seed=29), _sidecar_of(scheme, wider, + tmp_path / "wide.safetensors")) + + +@_by_name +def test_read_metadata_rejects_a_foreign_file(scheme: _Scheme, tmp_path) -> None: + plain = tmp_path / "plain.safetensors" + save_file({"w": torch.zeros(2, 2)}, str(plain)) + with pytest.raises(ValueError, match="not a FastVideo"): + scheme.read_metadata(plain) + + +@_by_name +def test_saving_an_unconverted_model_raises(scheme: _Scheme, tmp_path) -> None: + model = _build(scheme, seed=30) + with pytest.raises(RuntimeError, match="convert_model_to"): + scheme.save(model, tmp_path / "sidecar.safetensors") + + +@_by_name +def test_a_model_with_no_tagged_layers_raises_with_the_wiring_hint(scheme: _Scheme, tmp_path) -> None: + """The silent-dense failure mode must be reported, not ignored.""" + empty = _build(scheme, seed=31, config=scheme.make_config(target_layers=["nothing.matches.this"])) + with pytest.raises(RuntimeError, match=scheme.wiring_hint): + scheme.save(empty, tmp_path / "sidecar.safetensors") + with pytest.raises(RuntimeError, match=scheme.wiring_hint): + scheme.load(empty, tmp_path / "sidecar.safetensors") + + +# --------------------------------------------------------------------------- +# Deployment constraints: no GPU, no flashinfer +# --------------------------------------------------------------------------- + + +@_by_name +def test_load_needs_no_gpu_and_no_flashinfer(scheme: _Scheme, tmp_path, monkeypatch) -> None: + """Serving a pre-quantized checkpoint on the target host must not import kernels.""" + source = _build(scheme, seed=32) + scheme.convert(source) + expected = {key: value.clone() for key, value in scheme.state_dict(source).items()} + path = tmp_path / "sidecar.safetensors" + scheme.save(source, path) + + assert "flashinfer" not in sys.modules, "neither lane may pull in flashinfer" + target = _build(scheme, seed=33) + assert scheme.load(target, path) == 4 + assert "flashinfer" not in sys.modules + assert all(torch.equal(scheme.state_dict(target)[key], value) for key, value in expected.items()) + + +def test_sidecar_path_helpers(tmp_path) -> None: + for scheme in _ALL: + assert scheme.path_for("/models/h3/transformer.safetensors") == f"/models/h3/transformer{scheme.suffix}" + assert scheme.path_for("/models/h3") == f"/models/h3{scheme.suffix}" + assert scheme.path_for(str(tmp_path)) == str(tmp_path / scheme.dir_name) diff --git a/fastvideo/tests/ops/quantization/test_w4a16_config.py b/fastvideo/tests/ops/quantization/test_w4a16_config.py new file mode 100644 index 0000000000..c3691de261 --- /dev/null +++ b/fastvideo/tests/ops/quantization/test_w4a16_config.py @@ -0,0 +1,462 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU-only tests for the W4A16 (4-bit weight / 16-bit activation) config. + +Scope: the quantizer round-trip, the packed-code layout, the MiniMax-H3 layer +selection (including the ``attn.to_gate_compress`` exclusion) and the load-time +conversion path. Nothing here needs CUDA, a GPU kernel, or flashinfer — the +compute path under test is the documented ``dequantize then dense GEMM`` +reference, which is pure PyTorch. + +Tolerance note: W4A16's 4-bit codes give a per-element reconstruction error +bounded by half the group's quantizer step, +``scale = (max - min) / 15``. Tests assert against that *derived* bound rather +than a hand-picked constant, so the assertion stays meaningful if the group +size changes. +""" + +from __future__ import annotations + +import pytest +import torch +import torch.nn as nn +import torch.nn.functional as F + +from fastvideo.layers.linear import ReplicatedLinear +from fastvideo.layers.quantization.w4a16_config import ( + DEFAULT_BITS, + DEFAULT_GROUP_SIZE, + MINIMAX_H3_MAIN_STACK_LINEAR_SUFFIXES, + W4A16Config, + W4A16QuantizeMethod, + convert_model_to_w4a16, + minimax_h3_w4a16_prefixes, + w4a16_dequantize, + w4a16_quantize, +) + + +def _random_weight(out_dim: int, in_dim: int, *, generator: torch.Generator) -> torch.Tensor: + """Deterministic, roughly-linear-layer-shaped weights (non-uniform rows).""" + return torch.randn(out_dim, in_dim, generator=generator) * 0.02 + + +# --------------------------------------------------------------------------- +# Config surface +# --------------------------------------------------------------------------- + + +def test_config_imports_without_cuda_dependencies(): + """Importing the config must not require CUDA, a kernel, or flashinfer.""" + config = W4A16Config() + assert config.get_name() == "W4A16" + assert torch.bfloat16 in config.get_supported_act_dtypes() + assert config.get_config_filenames() == [] + # Declared contract only -- the reference path is a dense 16-bit GEMM, so + # no 4-bit tensor-core class is required to *load* it. + assert W4A16Config.get_min_capability() >= 70 + assert DEFAULT_BITS == 4 + + +def test_config_rejects_unsupported_bits(): + with pytest.raises(ValueError): + W4A16Config(bits=3) + with pytest.raises(ValueError): + W4A16Config(group_size=0) + + +def test_from_config_round_trips_fields(): + config = W4A16Config.from_config({ + "group_size": 32, + "bits": 4, + "target_layers": ["a.b"], + "exclude_substrings": ["keep_me_dense"], + "retain_original_weight": False, + }) + assert config.group_size == 32 + assert config.target_layers == frozenset({"a.b"}) + assert "keep_me_dense" in config.exclude_substrings + assert config.retain_original_weight is False + + +# --------------------------------------------------------------------------- +# Quantizer round-trip +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("group_size", [32, 64, 128]) +def test_quantize_dequantize_round_trip_within_step_bound(group_size: int): + """Every reconstructed element is within half a quantizer step of the source.""" + generator = torch.Generator().manual_seed(0) + weight = _random_weight(48, 4 * group_size, generator=generator) + + codes, scales, zeros = w4a16_quantize(weight, group_size=group_size, bits=4) + restored = w4a16_dequantize(codes, scales, zeros, group_size=group_size, bits=4, out_shape=weight.shape) + + # The quantizer returns codes in the group layout; the logical weight shape + # is the caller's to supply. Codes are packed two-per-byte for bits=4. + assert codes.dtype == torch.uint8 + assert codes.shape == (weight.shape[0], weight.shape[1] // 2) + assert scales.shape == (weight.shape[0], weight.shape[1] // group_size) + assert zeros.shape == scales.shape + assert restored.shape == weight.shape + + error = (restored - weight).abs() + # Bound per group: ``code = round(w / scale + zero)`` is off by at most half + # a code, so the reconstruction is off by at most half a step. The zero + # point's own rounding is absorbed by that same ``round``, so it does not + # widen the bound. + per_element_bound = (scales.unsqueeze(-1) / 2 + 1e-6) + assert torch.all(error.reshape(weight.shape[0], -1, group_size) <= per_element_bound) + + +def test_round_trip_is_exact_for_uniform_group(): + """A group that lands exactly on the code grid reconstructs exactly. + + ``zero`` is a rounded integer, so the anchor is only exact when + ``-min / scale`` is already integral -- here the ladder spans codes 0..15 + with a step of 0.1, giving ``zero = 8`` exactly. + """ + group_size = 64 + ladder = ((torch.arange(16).float() - 8) * 0.1).repeat(group_size // 16) + weight = ladder.repeat(2, 1) + + codes, scales, zeros = w4a16_quantize(weight, group_size=group_size, bits=4) + restored = w4a16_dequantize(codes, scales, zeros, group_size=group_size, bits=4, out_shape=weight.shape) + + assert torch.allclose(restored, weight, atol=1e-6) + + +def test_round_trip_holds_across_activation_dtypes(): + """bf16/fp16 sources still reconstruct within the step bound.""" + group_size = 64 + generator = torch.Generator().manual_seed(1) + base = _random_weight(16, group_size * 2, generator=generator) + for dtype in (torch.bfloat16, torch.float16, torch.float32): + weight = base.to(dtype) + codes, scales, zeros = w4a16_quantize(weight, group_size=group_size, bits=4) + restored = w4a16_dequantize(codes, scales, zeros, group_size=group_size, bits=4, out_shape=weight.shape) + error = (restored - weight.float()).abs() + # A 16-bit source rounds before quantizing, so the tolerance over the + # half-step bound is widened by that dtype's own resolution. + bound = (scales.unsqueeze(-1) / 2 + 1e-2) + assert torch.all(error.reshape(weight.shape[0], -1, group_size) <= bound), dtype + + +def test_quantize_rejects_indivisible_group_size(): + weight = torch.randn(4, 100) + with pytest.raises(ValueError, match="not divisible"): + w4a16_quantize(weight, group_size=64) + + +def test_packed_codes_store_four_bits_per_weight(): + """The 4-bit storage claim is real: 4 bits/weight for the codes themselves. + + (Excludes the per-group scales/zeros, which add ``2 * 32 / group_size`` + bits per weight -- 1 bit/weight at ``group_size=64``.) + """ + weight = torch.randn(32, 128) + codes, scales, zeros = w4a16_quantize(weight, group_size=64, bits=4) + assert codes.numel() == weight.numel() // 2 + assert codes.element_size() == 1 + assert codes.numel() * 8 == weight.numel() * 4 + assert scales.numel() == zeros.numel() == weight.numel() // 64 + + +def test_packed_nibble_layout_is_low_first(): + """Low nibble holds the lower K index -- the documented packing order. + + Asserted structurally rather than against hand-computed codes: a strictly + increasing row must unpack to a non-decreasing code sequence. Swapping the + nibble order would make the unpacked sequence zig-zag within every pair, + so this distinguishes the two conventions without depending on where the + zero-point rounding lands. + """ + from fastvideo.layers.quantization.w4a16_config import _unpack_4bit + + group_size = 64 + weight = torch.linspace(-1.0, 1.0, group_size).unsqueeze(0) + codes, _, _ = w4a16_quantize(weight, group_size=group_size, bits=4) + unpacked = _unpack_4bit(codes) + + assert unpacked.shape == weight.shape + assert unpacked.dtype == torch.uint8 + deltas = unpacked[0, 1:].to(torch.int16) - unpacked[0, :-1].to(torch.int16) + assert torch.all(deltas >= 0), unpacked[0] + + +def test_nan_does_not_poison_the_group(): + """nan_to_num matches the other configs: one NaN must not nuke its group.""" + group_size = 64 + weight = torch.randn(1, group_size) + weight[0, 3] = float("nan") + codes, scales, zeros = w4a16_quantize(weight, group_size=group_size, bits=4) + restored = w4a16_dequantize(codes, scales, zeros, group_size=group_size, bits=4, out_shape=weight.shape) + assert torch.isfinite(restored).all() + + +# --------------------------------------------------------------------------- +# MiniMax-H3 layer selection +# --------------------------------------------------------------------------- + + +def test_h3_allowlist_is_362_linears(): + config = W4A16Config.for_minimax_h3() + assert len(config.target_layers) == 362 + # 50 main blocks x (4 attn + 2 ff + 1 adaln) + 2 refiner blocks x (4 attn + 2 ff) + assert len(config.target_layers) == 50 * 7 + 2 * 6 + assert len(minimax_h3_w4a16_prefixes()) == 362 + + +def test_h3_selection_includes_the_intended_linears(): + config = W4A16Config.for_minimax_h3() + included = [ + "minimax_h3.transformer_blocks.0.attn.to_q", + "minimax_h3.transformer_blocks.0.attn.to_k", + "minimax_h3.transformer_blocks.0.attn.to_v", + "minimax_h3.transformer_blocks.0.attn.to_out", + "minimax_h3.transformer_blocks.0.ff.fc_in", + "minimax_h3.transformer_blocks.0.ff.fc_out", + "minimax_h3.transformer_blocks.0.adaln_proj.linear", + "minimax_h3.transformer_blocks.49.attn.to_q", + "minimax_h3.transformer_blocks.49.ff.fc_out", + "minimax_h3.token_refiner.refiner_blocks.0.attn.to_q", + "minimax_h3.token_refiner.refiner_blocks.1.ff.fc_out", + ] + for prefix in included: + assert config.is_target_layer(prefix), prefix + assert "adaln_proj.linear" in MINIMAX_H3_MAIN_STACK_LINEAR_SUFFIXES + + +def test_h3_selection_excludes_to_gate_compress(): + """The VSA gate steers discrete sparse routing -- it must never be quantized.""" + config = W4A16Config.for_minimax_h3() + for index in (0, 17, 49): + assert not config.is_target_layer(f"minimax_h3.transformer_blocks.{index}.attn.to_gate_compress") + assert not config.is_target_layer("minimax_h3.token_refiner.refiner_blocks.0.attn.to_gate_compress") + + +def test_to_gate_compress_exclusion_cannot_be_widened_away(): + """A caller cannot opt the gate back in, even by naming it in the allowlist.""" + gate = "minimax_h3.transformer_blocks.0.attn.to_gate_compress" + config = W4A16Config(target_layers=[gate, "minimax_h3.transformer_blocks.0.attn.to_q"]) + assert not config.is_target_layer(gate) + assert config.is_target_layer("minimax_h3.transformer_blocks.0.attn.to_q") + + +def test_h3_gate_module_is_built_and_left_dense(monkeypatch): + """The strongest form of the gate check: the real H3 attention module. + + H3 only builds ``attn.to_gate_compress`` when the VSA backend resolves + (``use_vsa`` guards the construction), which does not happen on a CPU-only + host. Forcing the backend resolution is enough to get the module built and + observe which quant method it receives -- this is the wiring the exclusion + actually has to protect, not just the string predicate. + """ + from fastvideo.layers.linear import UnquantizedLinearMethod + + import fastvideo.models.dits.minimax_h3 as h3_module + from fastvideo.platforms import AttentionBackendEnum + + class _FakeVSABackend: + + def get_name(self) -> str: + return "VIDEO_SPARSE_ATTN_H3" + + monkeypatch.setattr(h3_module, "get_attn_backend", lambda *args, **kwargs: _FakeVSABackend()) + + prefix = "minimax_h3.transformer_blocks.0.attn" + attention = h3_module.MiniMaxH3Attention( + 64, + 2, + 32, + 1e-5, + (AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3, ), + W4A16Config.for_minimax_h3(), + prefix=prefix, + ) + assert attention.to_gate_compress is not None + assert isinstance(attention.to_gate_compress.quant_method, UnquantizedLinearMethod) + assert isinstance(attention.to_q.quant_method, W4A16QuantizeMethod) + assert isinstance(attention.to_out.quant_method, W4A16QuantizeMethod) + + +def test_h3_selection_excludes_fp32_pinned_modules(): + config = W4A16Config.for_minimax_h3() + for prefix in ( + "minimax_h3.proj_in", + "minimax_h3.audio_proj_in", + "minimax_h3.time_embedder.fc_in", + "minimax_h3.proj_out", + "minimax_h3.audio_proj_out", + "minimax_h3.adaln_basis", + "minimax_h3.norm_out.linear", + "minimax_h3.context_embedder", + ): + assert not config.is_target_layer(prefix), prefix + + +def test_generic_suffix_matching_respects_dot_boundaries(): + """``ff.fc_in`` must not match a hypothetical ``cross_ff.fc_in``.""" + config = W4A16Config() + assert config.is_target_layer("model.blocks.0.ff.fc_in") + assert not config.is_target_layer("model.blocks.0.cross_ff.fc_in") + + +def test_get_quant_method_skips_indivisible_and_odd_dims(): + """A group must fit inside one row and 4-bit codes need an even last dim.""" + config = W4A16Config(target_layers=["blk.a", "blk.b", "blk.c"], group_size=64) + assert isinstance(config.get_quant_method(ReplicatedLinear(64, 8, bias=False, prefix="blk.a"), "blk.a"), + W4A16QuantizeMethod) + # 100 is not divisible by 64 -> dense. + assert config.get_quant_method(ReplicatedLinear(100, 8, bias=False, prefix="blk.b"), "blk.b") is None + # Divisible by group_size but odd -> 4-bit packing impossible -> dense. + odd = W4A16Config(target_layers=["blk.c"], group_size=5) + assert odd.get_quant_method(ReplicatedLinear(25, 8, bias=False, prefix="blk.c"), "blk.c") is None + + +def test_get_quant_method_ignores_non_linear_layers(): + config = W4A16Config() + assert config.get_quant_method(nn.RMSNorm(64), "minimax_h3.transformer_blocks.0.attn.norm_q") is None + + +# --------------------------------------------------------------------------- +# Load-time conversion path +# --------------------------------------------------------------------------- + + +def _tiny_model(prefix: str, quant_config: W4A16Config, in_dim: int = 64, out_dim: int = 32): + linear = ReplicatedLinear(in_dim, out_dim, bias=False, quant_config=quant_config, prefix=prefix) + module = nn.Module() + module.add_module("linear", linear) + return module, linear + + +def test_convert_registers_non_persistent_buffers_and_keeps_weight(): + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + model, linear = _tiny_model("block.linear", config) + linear.weight.data.normal_() + + convert_model_to_w4a16(model) + + assert linear._w4a16_codes.dtype == torch.uint8 + assert tuple(linear._w4a16_weight_shape) == tuple(linear.weight.shape) + # Non-persistent: the quantized payload must not leak into checkpoints. + assert list(model.state_dict().keys()) == ["linear.weight"] + assert linear.weight is not None # retained by default + + +def test_apply_matches_the_documented_dequantize_then_gemm_reference(): + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + model, linear = _tiny_model("block.linear", config) + linear.weight.data.normal_() + convert_model_to_w4a16(model) + + x = torch.randn(4, 64) + reference = F.linear( + x, + w4a16_dequantize(linear._w4a16_codes, + linear._w4a16_scales, + linear._w4a16_zeros, + out_shape=linear._w4a16_weight_shape, + out_dtype=x.dtype), + ) + out, _ = linear(x) + assert torch.equal(out, reference) + + +def test_apply_converts_lazily_when_the_loader_hook_never_ran(): + """No conversion hook -> still correct, just later and noisier.""" + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + linear = ReplicatedLinear(64, 32, bias=False, quant_config=config, prefix="block.linear") + linear.weight.data.normal_() + assert getattr(linear, "_w4a16_codes", None) is None + + x = torch.randn(2, 64) + with torch.no_grad(): + out, _ = linear(x) + assert out.shape == (2, 32) + assert getattr(linear, "_w4a16_codes", None) is not None + + +def test_grad_enabled_forward_stays_dense(): + """Under grad the master weight must be used -- no frozen dequantized copy.""" + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + linear = ReplicatedLinear(64, 32, bias=False, quant_config=config, prefix="block.linear") + linear.weight.data.normal_() + + x = torch.randn(2, 64) + out, _ = linear(x) # grad enabled by default in pytest + assert out.shape == (2, 32) + assert getattr(linear, "_w4a16_codes", None) is None + + +def test_purging_frees_the_dense_weight_when_opted_in(): + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"], retain_original_weight=False) + model, linear = _tiny_model("block.linear", config) + linear.weight.data.normal_() + + convert_model_to_w4a16(model) + + assert linear.weight is None + assert "linear.weight" not in model.state_dict() + assert linear._w4a16_codes is not None + # apply() must still work off the buffers alone. + out, _ = linear(torch.randn(2, 64)) + assert out.shape == (2, 32) + + +def test_untargeted_layer_is_untouched_by_convert(): + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + model, linear = _tiny_model("block.linear", config) + dense = ReplicatedLinear(64, 32, bias=False, prefix="block.other") + model.add_module("other", dense) + linear.weight.data.normal_() + dense.weight.data.normal_() + + before = dense.weight.clone() + convert_model_to_w4a16(model) + + assert getattr(dense, "_w4a16_codes", None) is None + assert torch.equal(dense.weight, before) + + +def test_bias_is_preserved_on_the_quantized_path(): + torch.manual_seed(0) + config = W4A16Config(target_layers=["block.linear"]) + linear = ReplicatedLinear(64, 32, bias=True, quant_config=config, prefix="block.linear") + model = nn.Module() + model.add_module("linear", linear) + linear.weight.data.normal_() + + convert_model_to_w4a16(model) + x = torch.randn(3, 64) + out, out_bias = linear(x) + # ``skip_bias_add`` is False, so the bias is folded into the output. + assert out_bias is None + reference = F.linear( + x, + w4a16_dequantize(linear._w4a16_codes, + linear._w4a16_scales, + linear._w4a16_zeros, + out_shape=linear._w4a16_weight_shape, + out_dtype=x.dtype), linear.bias) + assert torch.equal(out, reference) + + +def test_default_group_size_divides_every_h3_targeted_input_dim(): + """The H3 profile's group size is only valid if it divides all its K dims. + + Documented H3 dims (``MiniMaxH3ArchConfig``): hidden 5376, attention inner + 50 * 128 = 7168, ffn 14336, adaln 2688. This test pins that arithmetic so a + config change that breaks divisibility fails here rather than at load time. + """ + h3_input_dims = (5376, 7168, 14336, 2688) + for dim in h3_input_dims: + assert dim % DEFAULT_GROUP_SIZE == 0, dim + assert dim % 2 == 0, dim diff --git a/fastvideo/tests/train/callbacks/test_callback.py b/fastvideo/tests/train/callbacks/test_callback.py index 13a022a06b..1d6f85cb30 100644 --- a/fastvideo/tests/train/callbacks/test_callback.py +++ b/fastvideo/tests/train/callbacks/test_callback.py @@ -96,6 +96,7 @@ def test_default_hooks_return_none(self) -> None: assert (cb.on_training_step_end(method=None, loss_dict={}) is None) assert cb.on_before_optimizer_step(method=None) is None assert cb.on_validation_begin(method=None) is None + assert cb.will_run_validation() is False assert cb.on_validation_end(method=None) is None assert cb.on_train_end(method=None) is None diff --git a/fastvideo/tests/train/callbacks/test_ema.py b/fastvideo/tests/train/callbacks/test_ema.py index 3e84079aa3..c2b906bcf1 100644 --- a/fastvideo/tests/train/callbacks/test_ema.py +++ b/fastvideo/tests/train/callbacks/test_ema.py @@ -47,6 +47,18 @@ def __init__( self.tracker = tracker +class _CriticOnlyMethod(_Method): + + def __init__(self, transformer: torch.nn.Module) -> None: + super().__init__(transformer) + self._student_optimizer = object() + self._critic_optimizer = object() + + def get_optimizers(self, iteration: int) -> list[object]: + del iteration + return [self._critic_optimizer] + + def _tiny_transformer(*, fill: float = 0.0) -> torch.nn.Module: m = torch.nn.Linear(4, 2, bias=False) with torch.no_grad(): @@ -108,6 +120,19 @@ def test_skipped_until_start_iter(self) -> None: torch.full((2, 4), 1.0), ) + def test_skips_iterations_without_student_optimizer(self) -> None: + transformer = _tiny_transformer(fill=1.0) + cb = EMACallback(decay=0.5, start_iter=0) + method = _CriticOnlyMethod(transformer) + cb.on_train_start(method, iteration=0) + + with torch.no_grad(): + transformer.weight.fill_(7.0) + cb.on_training_step_end(method, loss_dict={}, iteration=1) + + assert not cb._ema_started + assert torch.allclose(cb.student_ema.shadow["weight"], torch.full((2, 4), 1.0)) + def test_first_active_step_reinits_then_updates(self) -> None: transformer = _tiny_transformer(fill=1.0) cb = EMACallback(decay=0.9, start_iter=10) diff --git a/fastvideo/tests/train/callbacks/test_latent_vis_shape_context.py b/fastvideo/tests/train/callbacks/test_latent_vis_shape_context.py new file mode 100644 index 0000000000..2f3ec3d9ac --- /dev/null +++ b/fastvideo/tests/train/callbacks/test_latent_vis_shape_context.py @@ -0,0 +1,50 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Latent visualization keeps native packed-shape context.""" + +from types import SimpleNamespace + +import torch + +from fastvideo.train.callbacks.latent_vis import LatentVisCallback + + +def test_latent_vis_passes_batch_local_layout_to_decoder(monkeypatch) -> None: + layout = object() + latent = torch.ones(1, 9) + decoded: list[tuple[torch.Tensor, object]] = [] + + class Student: + + @staticmethod + def decode_vis_latents(value, *, layout): + decoded.append((value, layout)) + return torch.zeros(1, 1, 3, 2, 2, dtype=torch.uint8).numpy() + + class Tracker: + + def video(self, clip, *, fps, format): + assert fps == 24 + assert format == "mp4" + return clip + + def log_artifacts(self, artifacts, iteration): + assert set(artifacts) == {"latent_vis/generator_pred_video"} + assert iteration == 8 + + monkeypatch.setattr( + "fastvideo.train.callbacks.latent_vis.get_world_group", + lambda: SimpleNamespace(rank=0), + ) + callback = LatentVisCallback(every_steps=1, keys=["generator_pred_video"]) + callback.tracker = Tracker() + method = SimpleNamespace( + student=Student(), + latent_vis={ + "generator_pred_video": latent, + "_fv_latent_layout": layout, + }, + ) + + callback.on_training_step_end(method, {}, iteration=8) + + assert decoded == [(latent, layout)] diff --git a/fastvideo/tests/train/callbacks/test_validation.py b/fastvideo/tests/train/callbacks/test_validation.py index e8baf8d1df..448cfbcbdd 100644 --- a/fastvideo/tests/train/callbacks/test_validation.py +++ b/fastvideo/tests/train/callbacks/test_validation.py @@ -21,6 +21,7 @@ import torch from fastvideo.api.sampling_param import SamplingParam +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_input_preparation import resolve_target_num_frames from fastvideo.train.callbacks.callback import CallbackDict from fastvideo.train.callbacks.ema import EMACallback from fastvideo.train.callbacks.validation import ( @@ -44,6 +45,8 @@ def _make_callback( sampling_steps: list[int] | None = None, guidance_scale: float | None = None, num_frames: int | None = None, + use_record_dimensions: bool = False, + max_record_num_frames: int | None = None, num_videos_per_prompt: int = 1, use_validation_media_conditioning: bool = True, sampling_timesteps: list[int] | None = None, @@ -59,6 +62,8 @@ def _make_callback( sampling_steps=sampling_steps, guidance_scale=guidance_scale, num_frames=num_frames, + use_record_dimensions=use_record_dimensions, + max_record_num_frames=max_record_num_frames, num_videos_per_prompt=num_videos_per_prompt, use_validation_media_conditioning=use_validation_media_conditioning, sampling_timesteps=sampling_timesteps, @@ -84,6 +89,8 @@ def test_defaults(self) -> None: assert cb.sampling_steps == [40] assert cb.guidance_scale is None assert cb.num_frames is None + assert cb.use_record_dimensions is False + assert cb.max_record_num_frames is None assert cb.num_videos_per_prompt == 1 assert cb.use_validation_media_conditioning is True assert cb.run_at_start is True @@ -110,6 +117,8 @@ def test_string_inputs_are_coerced(self) -> None: sampling_steps=["20", "40"], # type: ignore[arg-type] guidance_scale="4.5", # type: ignore[arg-type] num_frames="77", # type: ignore[arg-type] + use_record_dimensions="true", # type: ignore[arg-type] + max_record_num_frames="345", # type: ignore[arg-type] num_videos_per_prompt="2", # type: ignore[arg-type] use_validation_media_conditioning="false", # type: ignore[arg-type] sampling_timesteps=["1000", "500"], @@ -123,6 +132,8 @@ def test_string_inputs_are_coerced(self) -> None: assert cb.sampling_steps == [20, 40] assert cb.guidance_scale == 4.5 assert cb.num_frames == 77 + assert cb.use_record_dimensions is True + assert cb.max_record_num_frames == 345 assert cb.num_videos_per_prompt == 2 assert cb.use_validation_media_conditioning is False assert cb.run_at_start is False @@ -137,6 +148,12 @@ def test_init_rejects_nonpositive_video_count(self) -> None: with pytest.raises(ValueError, match="num_videos_per_prompt must be positive"): _make_callback(num_videos_per_prompt=0) + @pytest.mark.parametrize("max_record_num_frames", [0, -1]) + def test_init_rejects_nonpositive_record_frame_cap(self, max_record_num_frames: int) -> None: + """A configured native-record cap must describe a usable request.""" + with pytest.raises(ValueError, match="max_record_num_frames must be positive"): + _make_callback(max_record_num_frames=max_record_num_frames) + def test_pipeline_kwargs_collected(self) -> None: cb = ValidationCallback( pipeline_target=_PIPE_TARGET, @@ -220,6 +237,7 @@ class TestOnValidationBegin: def test_skipped_when_every_steps_zero(self) -> None: cb = _make_recording(every_steps=0) + assert cb.will_run_validation(0) is False cb.on_validation_begin(method=None, iteration=0) cb.on_validation_begin(method=None, iteration=1000) assert cb.run_calls == [] @@ -232,6 +250,8 @@ def test_skipped_on_off_iter(self) -> None: def test_runs_on_match(self) -> None: cb = _make_recording(every_steps=50) + assert cb.will_run_validation(50) is True + assert cb.will_run_validation(51) is False cb.on_validation_begin(method=None, iteration=50) cb.on_validation_begin(method=None, iteration=100) assert cb.run_calls == [50, 100] @@ -257,6 +277,8 @@ def test_on_validation_begin_skips_step_zero_when_disabled(self) -> None: cb.on_validation_begin(method=None, iteration=0) cb.on_validation_begin(method=None, iteration=20) assert cb.run_calls == [20] + assert cb.will_run_validation(0) is False + assert cb.will_run_validation(20) is True class TestH3ValidationContract: @@ -292,6 +314,122 @@ def test_prepare_validation_batch_forwards_video_count( assert batch.num_videos_per_prompt == 3 + def test_prepare_validation_batch_honors_complete_record_dimensions( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + """V10 samples each held-out prompt at its native spatial/temporal shape.""" + cb = _make_callback( + num_frames=77, + use_record_dimensions=True, + ) + cb.training_config = SimpleNamespace( + data=SimpleNamespace( + num_height=64, + num_width=96, + num_latent_t=2, + ), + pipeline_config=SimpleNamespace(vae_config=SimpleNamespace( + arch_config=SimpleNamespace(temporal_compression_ratio=4), ), ), + model_path="unused", + vsa_sparsity=0.0, + ) + monkeypatch.setattr( + "fastvideo.train.callbacks.validation.make_inference_args", + lambda *args, **kwargs: SimpleNamespace(), + ) + + batch = cb._prepare_validation_batch( + SamplingParam(), + { + "prompt": "Generate synchronized media.", + "width": 128, + "height": 80, + "num_frames": 39, + }, + num_inference_steps=4, + ) + + assert (batch.width, batch.height, batch.num_frames) == (128, 80, 39) + assert batch.n_tokens == 10 * 10 * 16 + + def test_record_dimensions_require_complete_triplet(self) -> None: + cb = _make_callback(use_record_dimensions=True) + cb.training_config = SimpleNamespace( + data=SimpleNamespace( + num_height=64, + num_width=96, + num_latent_t=2, + ), + pipeline_config=SimpleNamespace(vae_config=SimpleNamespace( + arch_config=SimpleNamespace(temporal_compression_ratio=4), ), ), + ) + + with pytest.raises(ValueError, match="must provide width, height, and num_frames together"): + cb._validation_sampling_dimensions({ + "prompt": "partial", + "width": 128, + }) + + @pytest.mark.parametrize( + ("record_num_frames", "max_record_num_frames", "expected_num_frames"), + [ + (328, 345, 328), + (362, 345, 345), + (362, None, 362), + ], + ) + def test_record_frame_cap_preserves_shorter_and_legacy_geometry( + self, + record_num_frames: int, + max_record_num_frames: int | None, + expected_num_frames: int, + ) -> None: + """The opt-in cap affects only over-limit native record lengths.""" + cb = _make_callback( + use_record_dimensions=True, + max_record_num_frames=max_record_num_frames, + ) + cb.training_config = SimpleNamespace( + data=SimpleNamespace( + num_height=64, + num_width=96, + num_latent_t=2, + ), + pipeline_config=SimpleNamespace(vae_config=SimpleNamespace( + arch_config=SimpleNamespace(temporal_compression_ratio=4), ), ), + ) + + assert cb._validation_sampling_dimensions({ + "width": 1344, + "height": 768, + "num_frames": record_num_frames, + }) == (768, 1344, expected_num_frames) + + def test_record_frame_cap_matches_h3_inference_boundary(self) -> None: + """The v10 cap is the largest released H3 geometry below 15 seconds.""" + assert resolve_target_num_frames(345) == 345 + with pytest.raises(ValueError, match="aligned num_frames=362"): + resolve_target_num_frames(362) + + def test_record_dimensions_are_ignored_without_opt_in(self) -> None: + cb = _make_callback(num_frames=77, use_record_dimensions=False) + cb.training_config = SimpleNamespace( + data=SimpleNamespace( + num_height=64, + num_width=96, + num_latent_t=2, + ), + pipeline_config=SimpleNamespace(vae_config=SimpleNamespace( + arch_config=SimpleNamespace(temporal_compression_ratio=4), ), ), + ) + + assert cb._validation_sampling_dimensions({ + "width": 128, + "height": 80, + "num_frames": 39, + }) == (64, 96, 77) + def test_prepare_validation_batch_ignores_media_for_text_only_generation( self, monkeypatch: pytest.MonkeyPatch, @@ -374,8 +512,12 @@ def forward(self, batch, inference_args): monkeypatch.setattr( cb, "_prepare_validation_batch", - lambda sampling_param, validation_batch, num_inference_steps: SimpleNamespace(prompt=validation_batch[ - "caption"], ), + lambda sampling_param, validation_batch, num_inference_steps: SimpleNamespace( + prompt=validation_batch["caption"], + width=96, + height=64, + num_frames=5, + ), ) result = cb._run_validation_for_steps(50, transformer=torch.nn.Identity()) @@ -384,6 +526,13 @@ def forward(self, batch, inference_args): assert len(result.videos) == 8 assert result.audio_sample_rates == [32_000] * 8 assert len(result.audio_waveforms) == 8 + assert result.metadata == [{ + "source": "unknown", + "sample_id": None, + "width": 96, + "height": 64, + "num_frames": 5, + }] * 8 for prompt_index, waveform in enumerate(result.audio_waveforms): assert torch.is_tensor(waveform) torch.testing.assert_close(waveform, torch.full((32, 2), float(prompt_index))) @@ -431,6 +580,64 @@ def log_artifacts(self, artifacts, step): assert artifacts["validation_videos_50_steps"] == filenames assert {key: artifacts[key] for key in scalar_metrics} == scalar_metrics + def test_log_validation_artifacts_keeps_references_and_metadata_in_one_event(self) -> None: + """Generated/reference pairs and source/shape receipts remain aligned.""" + + class FakeWandbTracker: + + def __init__(self) -> None: + self.artifact_calls = [] + + def video(self, filename, *, caption, fps): + return (filename, caption, fps) + + def log_artifacts(self, artifacts, step): + self.artifact_calls.append((artifacts, step)) + + cb = _make_callback() + cb.tracker = FakeWandbTracker() + metadata = [{ + "source": "nuva/50k", + "sample_id": "sample-1", + "width": 1344, + "height": 768, + "num_frames": 345, + "reference_num_frames": 362, + }] + caption = cb._validation_artifact_caption("A prompt", metadata[0]) + ref_caption = cb._validation_artifact_caption( + "A prompt", + metadata[0], + prefix="held-out reference", + use_reference_num_frames=True, + ) + scalar_metrics = cb._validation_metadata_scalar_metrics( + metadata, + num_inference_steps=4, + ) + + cb._log_validation_video_artifacts( + ["generated.mp4"], + [caption], + key="validation_videos_4_steps", + step=100, + fps=24, + reference_video_filenames=["reference.mp4"], + reference_captions=[ref_caption], + reference_key="validation_references_4_steps", + scalar_metrics=scalar_metrics, + ) + + artifacts, step = cb.tracker.artifact_calls[0] + assert step == 100 + assert artifacts["validation_videos_4_steps"][0][0] == "generated.mp4" + assert artifacts["validation_references_4_steps"][0][0] == "reference.mp4" + assert "source=nuva/50k" in artifacts["validation_videos_4_steps"][0][1] + assert "shape=1344x768x345f" in artifacts["validation_videos_4_steps"][0][1] + assert "shape=1344x768x362f" in artifacts["validation_references_4_steps"][0][1] + assert artifacts["validation/4_steps/source/nuva_50k_count"] == 1.0 + assert artifacts["validation/4_steps/shape/1344x768x345f_count"] == 1.0 + class TestAttnQatInferValidation: diff --git a/fastvideo/tests/train/callbacks/test_validation_sampling_contract.py b/fastvideo/tests/train/callbacks/test_validation_sampling_contract.py new file mode 100644 index 0000000000..826f547da6 --- /dev/null +++ b/fastvideo/tests/train/callbacks/test_validation_sampling_contract.py @@ -0,0 +1,95 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Validation must sample at the operating point training teaches. + +Two independent knobs have to survive the training-config -> validation hop: +the few-step denoising ladder and the VSA attention contract (sparsity AND +tile geometry). v8 shipped 2400 steps of validation that missed both — three +forwards on the scheduler's native grid at tile 256, while training ran four +forwards on ``[999, 749, 500, 250]`` at tile 64 — with nothing raising. +""" + +from __future__ import annotations + +import types + +import pytest + +from fastvideo.train.callbacks.validation import ValidationCallback + +LADDER = [999, 749, 500, 250] + + +def _callback(**kwargs): + """A callback instance without touching distributed state.""" + return ValidationCallback( + pipeline_target="fastvideo.pipelines.basic.minimax_h3." + "minimax_h3_pipeline.MiniMaxH3Pipeline", + dataset_file="unused.json", + **kwargs, + ) + + +def _method(ladder=LADDER): + cfg = {} if ladder is None else {"dmd_denoising_steps": list(ladder)} + return types.SimpleNamespace(method_config=cfg) + + +def test_ladder_is_inherited_when_unset() -> None: + cb = _callback(sampling_steps=[4]) + assert cb.sampling_timesteps is None + cb._adopt_training_sampling_contract(_method()) + assert cb.sampling_timesteps == LADDER + + +def test_matching_ladder_is_accepted() -> None: + cb = _callback(sampling_steps=[4], sampling_timesteps=LADDER) + cb._adopt_training_sampling_contract(_method()) + assert cb.sampling_timesteps == LADDER + + +def test_diverging_ladder_raises() -> None: + cb = _callback(sampling_steps=[4], sampling_timesteps=[1000, 667, 333]) + with pytest.raises(ValueError, match="disagrees with the trained ladder"): + cb._adopt_training_sampling_contract(_method()) + + +@pytest.mark.parametrize("ladder", [None, []]) +def test_non_dmd_methods_are_left_alone(ladder) -> None: + cb = _callback(sampling_steps=[40]) + cb._adopt_training_sampling_contract(_method(ladder)) + assert cb.sampling_timesteps is None + + +def test_method_without_config_is_tolerated() -> None: + cb = _callback(sampling_steps=[40]) + cb._adopt_training_sampling_contract(types.SimpleNamespace()) + assert cb.sampling_timesteps is None + + +def test_attention_contract_accepts_matching_args() -> None: + tc = types.SimpleNamespace(vsa_sparsity=0.9, vsa_tile_size=64) + args = types.SimpleNamespace(VSA_sparsity=0.9, VSA_tile_size=64) + ValidationCallback._assert_attention_contract(args, tc) + + +def test_attention_contract_catches_tile_size_drift() -> None: + """The exact v8 failure: sparsity propagated, tile size left at default.""" + tc = types.SimpleNamespace(vsa_sparsity=0.9, vsa_tile_size=64) + args = types.SimpleNamespace(VSA_sparsity=0.9, VSA_tile_size=256) + with pytest.raises(ValueError, match="VSA_tile_size=256"): + ValidationCallback._assert_attention_contract(args, tc) + + +def test_attention_contract_catches_sparsity_drift() -> None: + tc = types.SimpleNamespace(vsa_sparsity=0.9, vsa_tile_size=64) + args = types.SimpleNamespace(VSA_sparsity=0.0, VSA_tile_size=64) + with pytest.raises(ValueError, match="VSA_sparsity=0.0"): + ValidationCallback._assert_attention_contract(args, tc) + + +def test_make_inference_args_propagates_both_vsa_knobs() -> None: + """The missing line that caused the drift, pinned at its source.""" + from fastvideo.train.utils.moduleloader import make_inference_args + src = __import__("inspect").getsource(make_inference_args) + assert "args.VSA_sparsity = tc.vsa_sparsity" in src + assert "args.VSA_tile_size = tc.vsa_tile_size" in src diff --git a/fastvideo/tests/train/fixtures/minimax_h3_dmd2_min.yaml b/fastvideo/tests/train/fixtures/minimax_h3_dmd2_min.yaml new file mode 100644 index 0000000000..a29e4ae288 --- /dev/null +++ b/fastvideo/tests/train/fixtures/minimax_h3_dmd2_min.yaml @@ -0,0 +1,61 @@ +# Minimal H3 DMD2 trio and method configuration for CPU contract tests. +models: + student: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: MiniMaxAI/MiniMax-H3 + trainable: true + teacher: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: MiniMaxAI/MiniMax-H3 + trainable: false + critic: + _target_: fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel + init_from: MiniMaxAI/MiniMax-H3 + trainable: true + +method: + _target_: fastvideo.train.methods.distribution_matching.dmd2.DMD2Method + rollout_mode: data_latent + generator_update_interval: 1 + real_score_guidance_scale: 3.5 + dmd_denoising_steps: [1000, 757, 522] + min_timestep_ratio: 0.02 + max_timestep_ratio: 0.98 + cfg_uncond: + text: zero + fake_score_learning_rate: 1.0e-3 + fake_score_betas: [0.0, 0.999] + fake_score_lr_scheduler: constant + +training: + distributed: + num_gpus: 1 + sp_size: 1 + tp_size: 1 + hsdp_replicate_dim: 1 + hsdp_shard_dim: 1 + data: + data_path: /tmp/minimax_h3_t2va + preprocessed_data_type: t2va + train_batch_size: 1 + training_cfg_rate: 0.0 + seed: 42 + # Tiny geometry: [1, 24, 2, 4, 4] video latents, [1, 2, 32, 8] audio. + num_latent_t: 2 + num_height: 64 + num_width: 64 + num_frames: 5 + optimizer: + learning_rate: 1.0e-3 + betas: [0.0, 0.999] + weight_decay: 0.0 + lr_scheduler: constant + lr_warmup_steps: 0 + loop: + max_train_steps: 4 + model: + precondition_outputs: false + dit_precision: bf16 + +callbacks: {} +pipeline: {} diff --git a/fastvideo/tests/train/methods/test_dmd2_data_forcing.py b/fastvideo/tests/train/methods/test_dmd2_data_forcing.py new file mode 100644 index 0000000000..59ca37a594 --- /dev/null +++ b/fastvideo/tests/train/methods/test_dmd2_data_forcing.py @@ -0,0 +1,463 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Per-batch data forcing for carried DMD2 (v9, FastGen data-driven regime). + +CPU contract tests for ``rollout_data_forcing``: routing by latent +presence under mixed loading, forced-input noising math (uniform grid-rung +draw, per-modality shifts), walk pausing, uniform first-call seeding, and +knob validation. ``rollout_data_forcing: false`` (the default) keeps the +carried walk byte-identical — ``test_dmd2_rollout_carry.py`` covers that +path unchanged. +""" +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from fastvideo.tests.train.methods.test_dmd2_rollout_carry import ( + _GRID, + _LATENT_SHAPE, + _CarryStudent, + _make_method, + _stub_losses, +) +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method +from fastvideo.train.models.minimax_h3.minimax_h3 import shift_noise_amount + +_VIDEO_NUMEL = 6 +_AUDIO_NUMEL = 2 +assert _VIDEO_NUMEL + _AUDIO_NUMEL == _LATENT_SHAPE[1] + + +class _ForcingStudent(_CarryStudent): + """Carry fake that also packs real latents for ``latents_source='data'``.""" + + def __init__(self) -> None: + super().__init__() + self.prepare_sources: list[str] = [] + self.add_noise_calls: list[dict] = [] + + def prepare_batch(self, raw_batch, *, generator, latents_source): + self.prepare_sources.append(latents_source) + self.prepare_calls.append(raw_batch) + if latents_source == "data": + latents = torch.cat( + ( + raw_batch["vae_latent"].reshape(1, -1), + raw_batch["audio_latent"].reshape(1, -1), + ), + dim=1, + ) + else: + latents = torch.zeros(_LATENT_SHAPE) + batch = SimpleNamespace( + latents=latents, + timesteps=torch.tensor([0.0]), + attn_metadata=None, + attn_metadata_vsa="vsa-metadata", + dmd_latent_vis_dict={}, + fake_score_latent_vis_dict={}, + ) + self.last_batch = batch + return batch + + def add_noise(self, clean, noise, timestep): + noisy = super().add_noise(clean, noise, timestep) + self.add_noise_calls.append({ + "clean": clean, + "noise": noise, + "timestep": float(timestep.reshape(-1)[0]), + "noisy": noisy, + }) + return noisy + + +def _latent_batch(seed: int = 1) -> dict[str, torch.Tensor]: + generator = torch.Generator().manual_seed(seed) + return { + "vae_latent": torch.randn(1, _VIDEO_NUMEL, generator=generator), + "audio_latent": torch.randn(1, _AUDIO_NUMEL, generator=generator), + "text_embedding": torch.randn(1, 4, generator=generator), + } + + +def _text_only_batch() -> dict[str, torch.Tensor]: + """Mixed loading under the t2va schema: latent columns come through empty.""" + return { + "vae_latent": torch.zeros(1, 0), + "audio_latent": torch.zeros(1, 0), + "text_embedding": torch.ones(1, 4), + } + + +# ---------------------------------------------------------------------- +# (1) Off by default: latent-bearing batches still walk the carry +# ---------------------------------------------------------------------- + + +def test_data_forcing_defaults_off_and_latent_batches_walk() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1, student=_ForcingStudent()) + _stub_losses(method) + student = method.student + + assert method._rollout_data_forcing is False + _, _, metrics = method.single_train_step(_latent_batch(), iteration=0) + + # The batch's latents are ignored: the carried walk prepares with zeros + # and emits its usual metric set with no routing marker. + assert student.prepare_sources == ["zeros"] + assert "data_forced" not in metrics + assert metrics["rollout_step"] == 0.0 + + +# ---------------------------------------------------------------------- +# (2) Routing: latent presence picks the branch, emptiness means text-only +# ---------------------------------------------------------------------- + + +def test_routing_by_latent_presence_under_mixed_loading() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1, student=_ForcingStudent(), data_forcing=True) + _stub_losses(method) + student = method.student + + _, _, walk_metrics = method.single_train_step(_text_only_batch(), iteration=0) + _, _, forced_metrics = method.single_train_step(_latent_batch(), iteration=1) + + assert student.prepare_sources == ["zeros", "data"] + assert walk_metrics["data_forced"] == 0.0 + assert walk_metrics["rollout_step"] == 0.0 + assert forced_metrics["data_forced"] == 1.0 + # A forced call sits at a drawn rung, not on the walk's grid position. + assert "rollout_step" not in forced_metrics + + +def test_half_present_latent_pair_fails_loud() -> None: + method = _make_method(slots=1, sample_type="ode", student=_ForcingStudent(), data_forcing=True) + batch = _latent_batch() + batch["audio_latent"] = torch.zeros(1, 0) + with pytest.raises(ValueError, match="exactly one of"): + method.single_train_step(batch, iteration=0) + + +# ---------------------------------------------------------------------- +# (3) Forced-input math: uniform grid-rung draw, forward-noised real latents +# ---------------------------------------------------------------------- + + +def test_forced_input_is_real_latents_noised_at_a_grid_rung() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1, student=_ForcingStudent(), data_forcing=True) + _stub_losses(method) + student = method.student + # Pre-seed the slot: the forced step itself is under test, not the + # one-time stagger. + method._carry_slot_seeded[0] = True + + batch = _latent_batch(seed=3) + packed_real = torch.cat( + (batch["vae_latent"].reshape(1, -1), batch["audio_latent"].reshape(1, -1)), + dim=1, + ) + method.single_train_step(batch, iteration=0) + + forced = student.add_noise_calls[0] + assert forced["timestep"] in [float(t) for t in _GRID] + torch.testing.assert_close(forced["clean"], packed_real) + sigma = forced["timestep"] / 1000.0 + torch.testing.assert_close(forced["noisy"], (1.0 - sigma) * packed_real + sigma * forced["noise"]) + # The student trains exactly on that noised-real state at that rung. + main = student.predict_calls[-1] + assert main["timestep"] == forced["timestep"] + assert main["grad_enabled"] is True + assert main["attn_kind"] == "vsa" + vis = student.last_batch.dmd_latent_vis_dict + torch.testing.assert_close(vis["generator_timestep"], torch.tensor([forced["timestep"]])) + + +def test_forced_rung_draw_covers_the_whole_grid_and_never_zero() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1, student=_ForcingStudent(), data_forcing=True) + _stub_losses(method) + method._carry_slot_seeded[0] = True + + for call in range(64): + method.single_train_step(_latent_batch(seed=call), iteration=call) + # _stub_losses replaces the loss paths, so the only add_noise per call is + # the forced one; FastGen's sample_from_t_list never yields t = 0. + drawn = {call["timestep"] for call in method.student.add_noise_calls} + assert drawn == {float(t) for t in _GRID} + assert 0.0 not in drawn + + +# ---------------------------------------------------------------------- +# (4) The walk pauses on forced batches and resumes untouched +# ---------------------------------------------------------------------- + + +def test_forced_batches_pause_the_walk_and_text_batches_resume_it() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1, student=_ForcingStudent(), data_forcing=True) + _stub_losses(method) + + # Start the walk with a text-only batch (rung 0 trained, carry at rung 1). + method.single_train_step(_text_only_batch(), iteration=0) + carried = method._carry_slots[0] + assert carried is not None and carried["rung"] == 1 + state_before = carried["state"].clone() + + # Two forced calls: the slot's carry object and state are untouched. + method.single_train_step(_latent_batch(seed=5), iteration=1) + method.single_train_step(_latent_batch(seed=6), iteration=2) + assert method._carry_slots[0] is carried + torch.testing.assert_close(carried["state"], state_before) + + # The next text-only batch resumes at rung 1 and advances to rung 2. + _, _, metrics = method.single_train_step(_text_only_batch(), iteration=3) + assert metrics["rollout_step"] == 1.0 + assert method._carry_slots[0] is not None + assert method._carry_slots[0]["rung"] == 2 + + +def test_forced_batch_at_walk_boundary_leaves_the_boundary_state() -> None: + method = _make_method(slots=1, grid=[999, 500], sample_type="ode", interval=1, student=_ForcingStudent(), + data_forcing=True) + _stub_losses(method) + + # Walk the 2-rung grid to completion: offset 0, rungs 0 then 1, then clear. + method.single_train_step(_text_only_batch(), iteration=0) + method.single_train_step(_text_only_batch(), iteration=1) + assert method._carry_slots[0] is None + + # A forced batch at the boundary trains data-forced and does not restart + # the walk; the following text-only batch starts fresh at rung 0. + _, _, forced_metrics = method.single_train_step(_latent_batch(), iteration=2) + assert forced_metrics["data_forced"] == 1.0 + assert method._carry_slots[0] is None + _, _, metrics = method.single_train_step(_text_only_batch(), iteration=3) + assert metrics["rollout_step"] == 0.0 + + +# ---------------------------------------------------------------------- +# (5) Uniform seeding: a forced first call still runs the stagger pre-walk +# ---------------------------------------------------------------------- + + +def test_forced_first_call_runs_stagger_prewalk_with_uniform_forward_count() -> None: + forced_student = _ForcingStudent() + forced_method = _make_method(slots=1, sample_type="ode", interval=1, student=forced_student, data_forcing=True, + rank=1, world=2) + _stub_losses(forced_method) + text_student = _ForcingStudent() + text_method = _make_method(slots=1, sample_type="ode", interval=1, student=text_student, data_forcing=True, rank=1, + world=2) + _stub_losses(text_method) + + latent = _latent_batch(seed=7) + forced_method.single_train_step(latent, iteration=0) + text_method.single_train_step(_text_only_batch(), iteration=0) + + # Both branches pay the identical FSDP forward count on the slot's + # first-ever call: the whole-grid pre-walk plus this call's forward. + assert len(forced_student.predict_calls) == len(_GRID) + assert len(text_student.predict_calls) == len(_GRID) + + # The seeded walk waits at the stream's stagger rung with this batch's + # conditioning adopted; the forced call did not consume the walk. + carried = forced_method._carry_slots[0] + assert carried is not None + assert carried["rung"] == (1 * 1 + 0) % len(_GRID) + assert forced_method._carry_slot_seeded[0] is True + torch.testing.assert_close(carried["raw_batch"]["vae_latent"], latent["vae_latent"]) + + # The next text-only batch resumes the seeded walk at that rung with the + # adopted conditioning, not its own. + _, _, metrics = forced_method.single_train_step(_text_only_batch(), iteration=1) + assert metrics["rollout_step"] == float(carried["rung"]) + adopted = forced_student.prepare_calls[-1] + torch.testing.assert_close(adopted["text_embedding"], latent["text_embedding"]) + + +# ---------------------------------------------------------------------- +# (6) Knob parsing and validation +# ---------------------------------------------------------------------- + + +def test_data_forcing_requires_rollout_carry() -> None: + with pytest.raises(ValueError, match="rollout_carry: true"): + _make_method(carry=False, sample_type=None, student=_ForcingStudent(), data_forcing=True) + + +def test_data_forcing_rejects_non_bool() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {"rollout_data_forcing": "yes"}) + object.__setattr__(method, "_rollout_carry", True) + with pytest.raises(ValueError, match="must be a bool"): + method._parse_rollout_data_forcing() + + +def test_data_forcing_requires_explicit_legacy_mixed_regime_opt_in() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {"rollout_data_forcing": True}) + object.__setattr__(method, "_rollout_carry", True) + with pytest.raises(ValueError, match="not a FastGen recipe"): + method._parse_rollout_data_forcing() + + method.method_config["allow_mixed_rollout_regimes"] = True + assert method._parse_rollout_data_forcing() is True + + +def test_data_forcing_requires_t2va_schema() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "_rollout_mode", "simulate") + object.__setattr__(method, "_rollout_data_forcing", True) + object.__setattr__( + method, + "training_config", + SimpleNamespace(data=SimpleNamespace(preprocessed_data_type="text_only")), + ) + with pytest.raises(ValueError, match="t2va"): + method._validate_preprocessed_data_type() + object.__setattr__( + method, + "training_config", + SimpleNamespace(data=SimpleNamespace(preprocessed_data_type="t2va")), + ) + method._validate_preprocessed_data_type() + + +def test_batch_classifier_contract() -> None: + assert DMD2Method._batch_has_latents(_latent_batch()) is True + assert DMD2Method._batch_has_latents(_text_only_batch()) is False + # Rows without the latent keys at all (pure text_only schema) are text-only. + assert DMD2Method._batch_has_latents({"text_embedding": torch.ones(1, 4)}) is False + with pytest.raises(ValueError, match="exactly one of"): + DMD2Method._batch_has_latents({ + "vae_latent": torch.ones(1, 3), + "audio_latent": torch.zeros(1, 0), + }) + + +# ---------------------------------------------------------------------- +# Integration: forced call on the real H3 CPU trio, per-modality shifts +# ---------------------------------------------------------------------- + + +def _build_forcing_trio(monkeypatch: pytest.MonkeyPatch, *, interval: int) -> DMD2Method: + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import ( + _FIXTURE, + _make_model, + ) + from fastvideo.train.utils.config import load_run_config + + config = load_run_config(str(_FIXTURE)) + config.method["rollout_mode"] = "simulate" + config.method["generator_update_interval"] = interval + config.method["rollout_carry"] = True + config.method["rollout_carry_slots"] = 1 + config.method["rollout_sample_type"] = "ode" + config.method["rollout_data_forcing"] = True + config.method["allow_mixed_rollout_regimes"] = True + student = _make_model(monkeypatch, config.training, scale=1.0) + teacher = _make_model(monkeypatch, config.training, trainable=False, scale=0.5) + critic = _make_model(monkeypatch, config.training, scale=0.25) + student.init_preprocessors = lambda training_config: None + method = DMD2Method( + cfg=config, + role_models={ + "student": student, + "teacher": teacher, + "critic": critic, + }, + ) + method.cuda_generator = torch.Generator(device="cpu").manual_seed(0) + return method + + +def test_forced_noising_per_modality_shift_on_real_h3_trio(monkeypatch: pytest.MonkeyPatch) -> None: + """The forced input mixes each modality at its own shifted sigma.""" + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import _raw_batch + + method = _build_forcing_trio(monkeypatch, interval=5) + student = method.student + # Skip the one-time stagger so the first add_noise call is the forced one. + method._carry_slot_seeded[0] = True + + records: list[dict] = [] + original_add_noise = student.add_noise + + def spy_add_noise(clean, noise, timestep): + noisy = original_add_noise(clean, noise, timestep) + records.append({ + "clean": clean, + "noise": noise, + "timestep": timestep, + "noisy": noisy, + }) + return noisy + + monkeypatch.setattr(student, "add_noise", spy_add_noise) + + raw = _raw_batch(seed=1) + loss_map, outputs, metrics = method.single_train_step(raw, iteration=1) + + forced = records[0] + expected_packed = student.pack_latents( + raw["vae_latent"].permute(0, 2, 1, 3, 4).to(torch.bfloat16), + raw["audio_latent"].to(torch.bfloat16), + ) + torch.testing.assert_close(forced["clean"], expected_packed) + rung = int(forced["timestep"].reshape(-1)[0]) + assert rung in method.method_config["dmd_denoising_steps"] + + slices = dict(student.modality_slices()) + base = torch.tensor([rung / 1000.0], dtype=torch.float64) + for name, shift in (("video", 12.0), ("audio", 3.0)): + sigma = shift_noise_amount(base, shift) + expected = ((1.0 - sigma) * forced["clean"][:, slices[name]].to(torch.float64) + + sigma * forced["noise"][:, slices[name]].to(torch.float64)).to(torch.bfloat16) + torch.testing.assert_close(forced["noisy"][:, slices[name]], expected) + + # Critic phase: the forced generation feeds the critic loss; the paused + # slot stays at the boundary (never seeded a walk beyond the skip above). + assert metrics["data_forced"] == 1.0 + assert metrics["update_student"] == 0.0 + assert loss_map["fake_score_loss"].item() > 0.0 + assert method._carry_slots[0] is None + + method.backward(loss_map, outputs) + assert method.critic.transformer.scale.grad is not None + assert torch.isfinite(method.critic.transformer.scale.grad) + + +def test_forced_student_step_and_walk_adoption_on_real_h3_trio(monkeypatch: pytest.MonkeyPatch) -> None: + """A forced first call seeds the walk; the walk later reuses its prompt.""" + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import _raw_batch + + method = _build_forcing_trio(monkeypatch, interval=1) + student = method.student + + latent_raw = _raw_batch(seed=1) + loss_map, outputs, metrics = method.single_train_step(latent_raw, iteration=0) + + assert metrics["data_forced"] == 1.0 + assert metrics["update_student"] == 1.0 + assert torch.isfinite(loss_map["total_loss"]) + assert loss_map["generator_loss"].item() > 0.0 + assert "generator_loss_video" in metrics and "generator_loss_audio" in metrics + # The stagger seeded the walk at offset (0*1+0)%3 = 0 without training it. + carried = method._carry_slots[0] + assert carried is not None and carried["rung"] == 0 + method.backward(loss_map, outputs) + assert student.transformer.scale.grad is not None + assert torch.isfinite(student.transformer.scale.grad) + + # A text-only batch resumes the seeded walk under the adopted prompt. + text_raw = { + "vae_latent": torch.zeros(1, 0), + "audio_latent": torch.zeros(1, 0), + "text_embedding": _raw_batch(seed=9)["text_embedding"], + "text_attention_mask": torch.tensor([[1, 1, 0, 0]], dtype=torch.float32), + } + _, _, walk_metrics = method.single_train_step(text_raw, iteration=1) + assert walk_metrics["data_forced"] == 0.0 + assert walk_metrics["rollout_step"] == 0.0 + adopted_text = latent_raw["text_embedding"][:, :2].to(torch.bfloat16) + torch.testing.assert_close(student.transformer.last_encoder_hidden_states, adopted_text) diff --git a/fastvideo/tests/train/methods/test_dmd2_fake_score_loss_space.py b/fastvideo/tests/train/methods/test_dmd2_fake_score_loss_space.py new file mode 100644 index 0000000000..a2731b3403 --- /dev/null +++ b/fastvideo/tests/train/methods/test_dmd2_fake_score_loss_space.py @@ -0,0 +1,136 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Critic regression-space coverage for DMD2.""" +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method + +_SPLIT = 24 +_TOTAL = 32 +_SIGMA = {"video": 0.8, "audio": 0.5} +_TIMESTEP = torch.tensor([500]) + + +class _Student: + + def modality_slices(self): + return (("video", slice(0, _SPLIT)), ("audio", slice(_SPLIT, _TOTAL))) + + def add_noise(self, clean, noise, timestep): + out = clean.clone() + for name, sl in self.modality_slices(): + sigma = _SIGMA[name] + out[:, sl] = (1.0 - sigma) * clean[:, sl] + sigma * noise[:, sl] + return out + + +class _Critic: + + def __init__(self): + self.scale = torch.nn.Parameter(torch.tensor(0.1)) + self.predict_noise_calls = 0 + self.predict_x0_calls = 0 + + def predict_noise(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + self.predict_noise_calls += 1 + return self.scale * torch.ones_like(noisy) + + def predict_x0(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + self.predict_x0_calls += 1 + return self.scale * torch.ones_like(noisy) + + +def _method(space: str | None, critic: _Critic) -> DMD2Method: + method = object.__new__(DMD2Method) + config = {} if space is None else {"fake_score_loss_space": space} + object.__setattr__(method, "method_config", config) + object.__setattr__(method, "student", _Student()) + object.__setattr__(method, "critic", critic) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(0)) + object.__setattr__(method, "_cfg_uncond", None) + object.__setattr__(method, "_fake_score_loss_space", method._parse_fake_score_loss_space()) + object.__setattr__(method, "_sample_score_timestep", lambda device: _TIMESTEP) + return method + + +def _run(space: str | None, critic: _Critic) -> torch.Tensor: + method = _method(space, critic) + batch = SimpleNamespace(timesteps=None, attn_metadata=None, fake_score_latent_vis_dict=None) + gen = torch.randn(1, _TOTAL, generator=torch.Generator().manual_seed(7)) + loss, _, _, metrics = method._critic_flow_matching_loss(batch, generator_pred_x0=gen) + assert set(metrics) == {"fake_score_loss_video", "fake_score_loss_audio"} + return loss + + +def _expected(space: str | dict[str, str]) -> float: + gen = torch.randn(1, _TOTAL, generator=torch.Generator().manual_seed(7)) + if space == "x0": + pred_x0 = 0.1 * torch.ones_like(gen) + return sum(torch.mean((pred_x0[:, sl] - gen[:, sl])**2).item() + for _, sl in _Student().modality_slices()) + + noise = torch.randn(gen.shape, generator=torch.Generator().manual_seed(0)) + target = noise - gen + pred = 0.1 * torch.ones_like(gen) + expected = 0.0 + for name, sl in _Student().modality_slices(): + mse = torch.mean((pred[:, sl] - target[:, sl])**2).item() + space_m = space.get(name, "velocity") if isinstance(space, dict) else space + weight = _SIGMA[name]**2 if space_m == "x0" else 1.0 + expected += weight * mse + return expected + + +def test_default_is_velocity_space() -> None: + loss = _run(None, _Critic()) + assert loss.item() == pytest.approx(_expected("velocity"), rel=1e-5) + + +def test_x0_space_uses_direct_x0_regression() -> None: + critic = _Critic() + loss = _run("x0", critic) + assert loss.item() == pytest.approx(_expected("x0"), rel=1e-4) + assert critic.predict_x0_calls == 1 + assert critic.predict_noise_calls == 0 + + +def test_x0_space_keeps_critic_gradient() -> None: + critic = _Critic() + loss = _run("x0", critic) + loss.backward() + assert critic.scale.grad is not None + assert critic.scale.grad.abs().item() > 0.0 + + +def test_invalid_loss_space_rejected() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {"fake_score_loss_space": "eps"}) + with pytest.raises(ValueError, match="velocity, x0"): + method._parse_fake_score_loss_space() + + +def test_per_modality_space_mapping() -> None: + """Legacy mixed spaces retain their one-forward compatibility path.""" + spec = {"video": "x0", "audio": "velocity"} + loss = _run(spec, _Critic()) + assert loss.item() == pytest.approx(_expected(spec), rel=1e-4) + + +def test_mapping_unknown_modality_falls_back_to_default() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {"fake_score_loss_space": {"video": "x0"}}) + mapping = method._parse_fake_score_loss_space() + object.__setattr__(method, "_fake_score_loss_space", mapping) + assert method._fake_score_space_for("video") == "x0" + assert method._fake_score_space_for("audio") == "velocity" + + +def test_mapping_invalid_value_rejected() -> None: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {"fake_score_loss_space": {"audio": "eps"}}) + with pytest.raises(ValueError, match="velocity, x0"): + method._parse_fake_score_loss_space() diff --git a/fastvideo/tests/train/methods/test_dmd2_fastgen_parity.py b/fastvideo/tests/train/methods/test_dmd2_fastgen_parity.py new file mode 100644 index 0000000000..8099c1460e --- /dev/null +++ b/fastvideo/tests/train/methods/test_dmd2_fastgen_parity.py @@ -0,0 +1,229 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU regression tests for DMD2 behavior matched to the FastGen H3 recipe. + +The references below intentionally spell out the small pieces of FastGen math +instead of importing a second checkout. FastGen's rectified-flow schedule uses +float64 arithmetic over ``max_t=0.999`` and casts latent results back only once. +""" +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method +from fastvideo.train.models.minimax_h3.minimax_h3_dmd import MiniMaxH3DMDModel + +_FASTGEN_MAX_T = 0.999 +_TIMESTEP_SCALE = 1000.0 + + +def _fastgen_time_shift(value: torch.Tensor, shift: float) -> torch.Tensor: + """FastGen ``time_shift`` on the H3 schedule's finite time domain.""" + value = value.to(torch.float64) + return value * shift * _FASTGEN_MAX_T / (value * (shift - 1.0) + _FASTGEN_MAX_T) + + +def _h3_adapter() -> MiniMaxH3DMDModel: + model = MiniMaxH3DMDModel.__new__(MiniMaxH3DMDModel) + model.training_config = SimpleNamespace(data=SimpleNamespace( + num_latent_t=2, + num_frames=5, + num_height=64, + num_width=64, + )) + return model + + +def _score_method() -> DMD2Method: + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", { + "min_timestep_ratio": 0.001, + "max_timestep_ratio": 0.999, + "score_timestep_shift": 2.4, + "score_timestep_warp_max": _FASTGEN_MAX_T, + "score_timestep_continuous": True, + }) + object.__setattr__(method, "student", SimpleNamespace( + num_train_timesteps=1000, + shift_and_clamp_timestep=lambda timestep: timestep, + )) + lo, hi = method._parse_score_timestep_bounds() + object.__setattr__(method, "_score_min_timestep", lo) + object.__setattr__(method, "_score_max_timestep", hi) + object.__setattr__(method, "_score_timestep_shift", 2.4) + object.__setattr__(method, "_score_timestep_warp_max", _FASTGEN_MAX_T) + object.__setattr__(method, "_score_timestep_continuous", True) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(0)) + return method + + +@pytest.mark.parametrize("unit_draw", [0.0, 0.37, 1.0]) +def test_shifted_score_sampler_draws_uniform_coordinate_before_inverse_warp( + monkeypatch: pytest.MonkeyPatch, + unit_draw: float, +) -> None: + """The random draw is uniform on FastGen's pre-shift ``[.001, .999]``. + + H3 represents the score time on an unshifted base clock. Applying the + inverse 2.4 warp there makes its video shift compose to 5 and its audio + shift compose to 1.25, exactly matching FastGen's video-clock schedule. + """ + + def fixed_rand(size, *, device=None, dtype=None, generator=None): + del generator + assert tuple(size) == (1, ) + return torch.full((1, ), unit_draw, device=device, dtype=dtype) + + monkeypatch.setattr(torch, "rand", fixed_rand) + method = _score_method() + sampled = method._sample_score_timestep(torch.device("cpu")) + + pre_shift = torch.tensor( + [0.001 + unit_draw * (0.999 - 0.001)], + dtype=torch.float64, + ) + expected_base = _fastgen_time_shift(pre_shift, 1.0 / 2.4) * _TIMESTEP_SCALE + assert sampled.dtype == torch.float64 + torch.testing.assert_close(sampled, expected_base, rtol=0.0, atol=1e-10) + + sigma_video, sigma_audio = _h3_adapter()._noise_amounts(sampled) + torch.testing.assert_close( + sigma_video, + _fastgen_time_shift(pre_shift, 5.0), + rtol=0.0, + atol=1e-12, + ) + torch.testing.assert_close( + sigma_audio, + _fastgen_time_shift(pre_shift, 1.25), + rtol=0.0, + atol=1e-12, + ) + + +def _packed_bfloat16_pair(model: MiniMaxH3DMDModel) -> tuple[torch.Tensor, torch.Tensor]: + slices = dict(model.modality_slices()) + total = slices["audio"].stop + clean = torch.linspace(-2.75, 3.125, total, dtype=torch.float32).reshape(1, -1).to(torch.bfloat16) + noise = torch.linspace(1.875, -3.5, total, dtype=torch.float32).reshape(1, -1).to(torch.bfloat16) + return clean, noise + + +def _reference_sigmas(timestep: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + base = (timestep.reshape(-1)[:1].to(torch.float64) / _TIMESTEP_SCALE).clamp(0.0, _FASTGEN_MAX_T) + return _fastgen_time_shift(base, 12.0), _fastgen_time_shift(base, 3.0) + + +def _reference_mix(clean: torch.Tensor, noise: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor: + return ((1.0 - sigma) * clean.to(torch.float64) + sigma * noise.to(torch.float64)).to(clean.dtype) + + +def _reference_unmix(noisy: torch.Tensor, clean: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor: + return ((noisy.to(torch.float64) - (1.0 - sigma) * clean.to(torch.float64)) / sigma.clamp_min(1e-6)).to( + noisy.dtype) + + +def test_h3_add_noise_matches_fastgen_fp64_then_cast_for_bfloat16() -> None: + model = _h3_adapter() + clean, noise = _packed_bfloat16_pair(model) + timestep = torch.tensor([413.375], dtype=torch.float64) + sigma_video, sigma_audio = _reference_sigmas(timestep) + clean_video, clean_audio = model.unpack_latents(clean) + noise_video, noise_audio = model.unpack_latents(noise) + expected = model.pack_latents( + _reference_mix(clean_video, noise_video, sigma_video), + _reference_mix(clean_audio, noise_audio, sigma_audio), + ) + + actual = model.add_noise(clean, noise, timestep) + + assert actual.dtype == torch.bfloat16 + torch.testing.assert_close(actual, expected, rtol=0.0, atol=0.0) + + +def test_h3_extract_eps_matches_fastgen_fp64_then_cast_at_low_audio_sigma() -> None: + model = _h3_adapter() + clean, noisy = _packed_bfloat16_pair(model) + # This is the base-clock value reached from FastGen's lower score bound; + # audio sigma is about 0.00125, where early BF16 rounding is amplified. + timestep = _fastgen_time_shift(torch.tensor([0.001], dtype=torch.float64), 1.0 / 2.4) * _TIMESTEP_SCALE + sigma_video, sigma_audio = _reference_sigmas(timestep) + noisy_video, noisy_audio = model.unpack_latents(noisy) + clean_video, clean_audio = model.unpack_latents(clean) + expected = model.pack_latents( + _reference_unmix(noisy_video, clean_video, sigma_video), + _reference_unmix(noisy_audio, clean_audio, sigma_audio), + ) + + actual = model.extract_eps(noisy, clean, timestep) + + assert actual.dtype == torch.bfloat16 + torch.testing.assert_close(actual, expected, rtol=0.0, atol=0.0) + + +def test_h3_predict_x0_helper_matches_fastgen_fp64_then_cast() -> None: + model = _h3_adapter() + noisy, pred_noise = _packed_bfloat16_pair(model) + sigma = torch.tensor([0.3141592653589793], dtype=torch.float64) + expected = (noisy.to(torch.float64) - sigma * pred_noise.to(torch.float64)).to(torch.bfloat16) + + actual = model._to_x0(noisy, pred_noise, sigma) + + assert actual.dtype == torch.bfloat16 + torch.testing.assert_close(actual, expected, rtol=0.0, atol=0.0) + + +class _ScoreStudent: + + def modality_slices(self): + return (("video", slice(0, 3)), ("audio", slice(3, 5))) + + def add_noise(self, clean, noise, timestep): + del timestep + return 0.75 * clean + 0.25 * noise + + +class _DirectX0Critic: + + def __init__(self) -> None: + self.value = torch.nn.Parameter(torch.tensor(0.25)) + self.predict_x0_calls = 0 + + def predict_x0(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + del timestep, batch, conditional, cfg_uncond, attn_kind + self.predict_x0_calls += 1 + return self.value * torch.ones_like(noisy) + + def predict_noise(self, *args, **kwargs): + del args, kwargs + raise AssertionError("global x0 regression must call critic.predict_x0 directly") + + +def test_critic_x0_objective_is_direct_predict_x0_mse_per_modality() -> None: + method = object.__new__(DMD2Method) + critic = _DirectX0Critic() + object.__setattr__(method, "method_config", {"fake_score_loss_space": "x0"}) + object.__setattr__(method, "student", _ScoreStudent()) + object.__setattr__(method, "critic", critic) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(11)) + object.__setattr__(method, "_cfg_uncond", None) + object.__setattr__(method, "_fake_score_loss_space", {"__default__": "x0"}) + object.__setattr__(method, "_sample_score_timestep", lambda device: torch.tensor([500.0], device=device)) + batch = SimpleNamespace(timesteps=None, attn_metadata=None, fake_score_latent_vis_dict=None) + generated_x0 = torch.tensor([[1.0, -2.0, 4.0, 8.0, -16.0]]) + + loss, _, _, metrics = method._critic_flow_matching_loss(batch, generator_pred_x0=generated_x0) + + pred_x0 = torch.full_like(generated_x0, 0.25) + expected_video = torch.mean((pred_x0[:, :3] - generated_x0[:, :3])**2) + expected_audio = torch.mean((pred_x0[:, 3:] - generated_x0[:, 3:])**2) + assert critic.predict_x0_calls == 1 + torch.testing.assert_close(metrics["fake_score_loss_video"], expected_video) + torch.testing.assert_close(metrics["fake_score_loss_audio"], expected_audio) + torch.testing.assert_close(loss, expected_video + expected_audio) + + loss.backward() + assert critic.value.grad is not None + assert torch.isfinite(critic.value.grad) diff --git a/fastvideo/tests/train/methods/test_dmd2_rollout_carry.py b/fastvideo/tests/train/methods/test_dmd2_rollout_carry.py new file mode 100644 index 0000000000..c7d95ea1dd --- /dev/null +++ b/fastvideo/tests/train/methods/test_dmd2_rollout_carry.py @@ -0,0 +1,656 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Carried backward-simulation rollout for DMD2 (FastGen port). + +CPU contract tests for ``rollout_carry``: slot round-robin, rung +progression and clearing, first-ever staggered starts with a uniform +pre-walk, ODE renoise with per-modality shifts, carried conditioning, +and knob validation. ``rollout_carry: false`` (the default) keeps the +existing rollout path byte-identical; the pre-existing DMD2 suites +(``test_minimax_h3_dmd2.py`` et al.) cover that path unchanged. +""" +from __future__ import annotations + +from types import SimpleNamespace + +import pytest +import torch + +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method +from fastvideo.train.models.minimax_h3 import MiniMaxH3DMDModel +from fastvideo.train.models.minimax_h3.minimax_h3 import shift_noise_amount + +_GRID = [999, 749, 500, 250] +_LATENT_SHAPE = (1, 8) + + +class _CarryStudent: + """Rectified-flow fake on one packed tensor (sigma = t / 1000).""" + + device = torch.device("cpu") + + def __init__(self) -> None: + self.prepare_calls: list[dict] = [] + self.predict_calls: list[dict] = [] + self.last_batch: SimpleNamespace | None = None + self.last_pred: torch.Tensor | None = None + + def prepare_batch(self, raw_batch, *, generator, latents_source): + assert latents_source == "zeros" + self.prepare_calls.append(raw_batch) + batch = SimpleNamespace( + latents=torch.zeros(_LATENT_SHAPE), + timesteps=torch.tensor([0.0]), + attn_metadata=None, + attn_metadata_vsa="vsa-metadata", + dmd_latent_vis_dict={}, + fake_score_latent_vis_dict={}, + ) + self.last_batch = batch + return batch + + @staticmethod + def _sigma(timestep: torch.Tensor) -> float: + return float(timestep.reshape(-1)[0]) / 1000.0 + + def predict_x0(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + self.predict_calls.append({ + "timestep": float(timestep.reshape(-1)[0]), + "grad_enabled": torch.is_grad_enabled(), + "attn_kind": attn_kind, + }) + pred = noisy * 0.5 + self.last_pred = pred + return pred + + def add_noise(self, clean, noise, timestep): + sigma = self._sigma(timestep) + return (1.0 - sigma) * clean + sigma * noise + + def extract_eps(self, noisy, clean, timestep): + sigma = self._sigma(timestep) + return (noisy - (1.0 - sigma) * clean) / sigma + + +def _make_method( + *, + slots: int = 1, + grid: list[int] | None = None, + interval: int = 1, + sample_type: str | None = "ode", + carry: bool = True, + rank: int = 0, + world: int = 1, + grad_accum: int | None = None, + rollout_mode: str = "simulate", + student: object | None = None, + data_forcing: bool | None = None, + native_shape_bucketing: bool = False, +) -> DMD2Method: + method = object.__new__(DMD2Method) + config: dict = { + "rollout_mode": rollout_mode, + "rollout_carry": carry, + "dmd_denoising_steps": list(grid or _GRID), + "generator_update_interval": interval, + } + if carry: + config["rollout_carry_slots"] = slots + if sample_type is not None: + config["rollout_sample_type"] = sample_type + if data_forcing is not None: + config["rollout_data_forcing"] = data_forcing + if data_forcing: + config["allow_mixed_rollout_regimes"] = True + object.__setattr__(method, "method_config", config) + object.__setattr__(method, "student", student if student is not None else _CarryStudent()) + object.__setattr__( + method, + "training_config", + SimpleNamespace( + loop=SimpleNamespace(gradient_accumulation_steps=(slots if grad_accum is None else grad_accum)), + distributed=SimpleNamespace(sp_size=1), + data=SimpleNamespace(native_shape_bucketing=native_shape_bucketing), + ), + ) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(0)) + object.__setattr__(method, "_cfg_uncond", None) + object.__setattr__(method, "_denoising_step_list", None) + object.__setattr__(method, "_rollout_mode", method._parse_rollout_mode()) + object.__setattr__(method, "_rollout_carry_rank_world", lambda: (rank, world)) + knobs = method._parse_rollout_carry() + object.__setattr__(method, "_rollout_carry", knobs[0]) + object.__setattr__(method, "_rollout_carry_slot_count", knobs[1]) + object.__setattr__(method, "_rollout_sample_type", knobs[2]) + method._init_rollout_carry_state() + object.__setattr__( + method, + "_rollout_data_forcing", + method._parse_rollout_data_forcing(), + ) + return method + + +def _stub_losses(method: DMD2Method) -> list[torch.Tensor]: + """Replace the loss paths with recorders; carry mechanics stay real.""" + critic_preds: list[torch.Tensor] = [] + object.__setattr__(method, "_dmd_loss", lambda pred, batch: (torch.zeros(()), {})) + + def _critic(batch, *, generator_pred_x0=None): + critic_preds.append(generator_pred_x0) + return torch.zeros(()), "critic-ctx", {}, {} + + object.__setattr__(method, "_critic_flow_matching_loss", _critic) + return critic_preds + + +def _replay_ode_walk(state: torch.Tensor, grid: list[int], hops: int) -> torch.Tensor: + """Reference walk under the fake student's flow: pred = 0.5 * state.""" + for rung in range(hops): + sigma = grid[rung] / 1000.0 + pred = state * 0.5 + eps = (state - (1.0 - sigma) * pred) / sigma + sigma_next = grid[rung + 1] / 1000.0 + state = (1.0 - sigma_next) * pred + sigma_next * eps + return state + + +# ---------------------------------------------------------------------- +# (1) Slot round-robin and rung progression 0 -> 1 -> 2 -> 3 -> clear +# ---------------------------------------------------------------------- + + +def test_slot_round_robin_and_rung_progression_with_two_slots() -> None: + method = _make_method(slots=2, sample_type="sde", interval=1) + _stub_losses(method) + student = method.student + + rungs = [] + forwards_per_call = [] + for call in range(10): + before = len(student.predict_calls) + _, _, metrics = method.single_train_step({"text_embedding": torch.ones(1, 4)}, iteration=call) + rungs.append(metrics["rollout_step"]) + forwards_per_call.append(len(student.predict_calls) - before) + + # Offsets: slot 0 -> (0*2+0)%4 = 0, slot 1 -> (0*2+1)%4 = 1. Interleaved + # round-robin walks: slot0 = 0,1,2,3,clear,0 and slot1 = 1,2,3,clear,0,1. + assert rungs == [0.0, 1.0, 1.0, 2.0, 2.0, 3.0, 3.0, 0.0, 0.0, 1.0] + # First-ever fill of each slot pre-walks the whole grid (len-1 forwards) + # exactly once; every other call, including post-clear restarts, pays + # exactly one generation forward. + assert forwards_per_call == [4, 4, 1, 1, 1, 1, 1, 1, 1, 1] + # Slot 1 cleared after rung 3 (call 5), restarted at rung 0 (call 7), + # advanced to rung 1 (call 9) and carries rung 2; slot 0 cleared at + # call 6, restarted at call 8, and carries rung 1. + assert method._carry_slots[0] is not None and method._carry_slots[0]["rung"] == 1 + assert method._carry_slots[1] is not None and method._carry_slots[1]["rung"] == 2 + + +def test_carry_slots_cleared_after_last_rung() -> None: + method = _make_method(slots=1, grid=[999, 500], sample_type="ode", interval=1) + _stub_losses(method) + student = method.student + + counts, rungs = [], [] + for call in range(3): + before = len(student.predict_calls) + _, _, metrics = method.single_train_step({"x": torch.ones(1)}, iteration=call) + rungs.append(metrics["rollout_step"]) + counts.append(len(student.predict_calls) - before) + + # offset (0*1+0)%2 = 0: pre-walk (1 hop) then rung 0; rung 1 finishes the + # trajectory; the restart begins at rung 0 with NO stagger pre-walk. + assert rungs == [0.0, 1.0, 0.0] + assert counts == [2, 1, 1] + assert method._carry_slots[0] is not None + assert method._carry_slots[0]["rung"] == 1 + + +# ---------------------------------------------------------------------- +# (2) Staggered starts across (rank, slot) streams +# ---------------------------------------------------------------------- + + +@pytest.mark.parametrize( + ("rank", "slot", "expected_offset"), + [(0, 0, 0), (0, 1, 1), (1, 0, 2), (1, 1, 3), (2, 0, 0), (2, 1, 1)], +) +def test_stagger_offsets_follow_rank_slot_formula(rank: int, slot: int, expected_offset: int) -> None: + method = _make_method(slots=2, sample_type="ode", rank=rank, world=3) + student = method.student + step_list = method._get_denoising_step_list(torch.device("cpu")) + state = torch.randn(_LATENT_SHAPE, generator=torch.Generator().manual_seed(3)) + batch = SimpleNamespace(dmd_latent_vis_dict={}) + + snapshot, rung = method._staggered_start(state, batch, step_list, slot) + + assert rung == expected_offset == (rank * 2 + slot) % len(_GRID) + # Every rank walks the whole grid regardless of its offset (uniform FSDP + # collective count) and only keeps the snapshot at its own rung. + assert len(student.predict_calls) == len(_GRID) - 1 + assert all(not call["grad_enabled"] for call in student.predict_calls) + assert all(call["attn_kind"] == "vsa" for call in student.predict_calls) + torch.testing.assert_close(snapshot, _replay_ode_walk(state, _GRID, expected_offset)) + + +@pytest.mark.parametrize("rank", [0, 1, 7]) +@pytest.mark.parametrize("slot", [0, 1]) +def test_native_shape_stagger_is_rank_synchronous(rank: int, slot: int) -> None: + method = _make_method( + slots=2, + sample_type="ode", + rank=rank, + world=8, + native_shape_bucketing=True, + ) + step_list = method._get_denoising_step_list(torch.device("cpu")) + state = torch.randn(_LATENT_SHAPE, generator=torch.Generator().manual_seed(3)) + batch = SimpleNamespace(dmd_latent_vis_dict={}) + + _, rung = method._staggered_start(state, batch, step_list, slot) + + assert rung == slot + + +# ---------------------------------------------------------------------- +# (3) ODE renoise analytic identity with per-modality shifts (real H3 math) +# ---------------------------------------------------------------------- + + +def _h3_adapter() -> MiniMaxH3DMDModel: + model = MiniMaxH3DMDModel.__new__(MiniMaxH3DMDModel) + model.training_config = SimpleNamespace(data=SimpleNamespace( + num_latent_t=2, + num_frames=5, + num_height=64, + num_width=64, + )) + return model + + +def test_ode_renoise_analytic_identity_per_modality() -> None: + """From x_t = (1-s_m)x0 + s_m*eps and pred = x0, the advanced state is + (1-s'_m)x0 + s'_m*eps for both the video-shift and audio-shift slices.""" + model = _h3_adapter() + slices = dict(model.modality_slices()) + total = slices["audio"].stop + generator = torch.Generator().manual_seed(11) + x0 = torch.randn(1, total, generator=generator) + eps = torch.randn(1, total, generator=generator) + timestep = torch.tensor([500], dtype=torch.long) + next_timestep = torch.tensor([250], dtype=torch.long) + + x_t = model.add_noise(x0, eps, timestep) + implied = model.extract_eps(x_t, x0, timestep) + torch.testing.assert_close(implied, eps, rtol=1e-5, atol=1e-5) + + advanced = model.add_noise(x0, implied, next_timestep) + for name, shift in (("video", 12.0), ("audio", 3.0)): + sigma_next = float(shift_noise_amount(torch.tensor([0.25]), shift)) + expected = (1.0 - sigma_next) * x0[:, slices[name]] + sigma_next * eps[:, slices[name]] + torch.testing.assert_close(advanced[:, slices[name]], expected, rtol=1e-5, atol=1e-5) + + +def test_method_renoise_ode_uses_adapter_extract_eps() -> None: + model = _h3_adapter() + method = _make_method(slots=1, grid=[500, 250], sample_type="ode", student=model) + total = dict(model.modality_slices())["audio"].stop + generator = torch.Generator().manual_seed(13) + x0 = torch.randn(1, total, generator=generator) + eps = torch.randn(1, total, generator=generator) + timestep = torch.tensor([500], dtype=torch.long) + x_t = model.add_noise(x0, eps, timestep) + + step_list = method._get_denoising_step_list(torch.device("cpu")) + advanced = method._renoise(x_t, x0, timestep, 1, step_list) + expected = model.add_noise(x0, eps, torch.tensor([250], dtype=torch.long)) + torch.testing.assert_close(advanced, expected, rtol=1e-5, atol=1e-5) + + +def test_method_renoise_sde_draws_fresh_noise() -> None: + method = _make_method(slots=1, grid=[999, 500], sample_type="sde") + step_list = method._get_denoising_step_list(torch.device("cpu")) + state = torch.randn(_LATENT_SHAPE, generator=torch.Generator().manual_seed(5)) + pred = state * 0.5 + + advanced = method._renoise(state, pred, torch.tensor([999]), 1, step_list) + noise = torch.randn(_LATENT_SHAPE, generator=torch.Generator().manual_seed(0)) + torch.testing.assert_close(advanced, 0.5 * pred + 0.5 * noise) + + +# ---------------------------------------------------------------------- +# (4) Mid-walk calls reuse the carried conditioning +# ---------------------------------------------------------------------- + + +def test_mid_walk_reuses_carried_conditioning_and_ignores_fresh_batches() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1) + _stub_losses(method) + student = method.student + + batch_a = {"text_embedding": torch.full((1, 4), 1.0), "info_list": ["prompt-a"]} + batch_b = {"text_embedding": torch.full((1, 4), 2.0), "info_list": ["prompt-b"]} + batch_c = {"text_embedding": torch.full((1, 4), 3.0), "info_list": ["prompt-c"]} + + method.single_train_step(batch_a, iteration=0) + snapshot = student.prepare_calls[0] + # The adopted batch is a detached device snapshot, not the loader dict. + assert snapshot is not batch_a + torch.testing.assert_close(snapshot["text_embedding"], batch_a["text_embedding"]) + assert snapshot["info_list"] == ["prompt-a"] + + # Rungs 1..3: the fresh loader batches are ignored, the trajectory keeps + # the exact snapshot object it set out with. + for call in range(1, 4): + method.single_train_step(batch_b, iteration=call) + assert student.prepare_calls[call] is snapshot + torch.testing.assert_close(student.prepare_calls[call]["text_embedding"], batch_a["text_embedding"]) + + # Rung 3 finished the walk; the next call adopts the new conditioning. + method.single_train_step(batch_c, iteration=4) + torch.testing.assert_close(student.prepare_calls[4]["text_embedding"], batch_c["text_embedding"]) + assert student.prepare_calls[4]["info_list"] == ["prompt-c"] + + +# ---------------------------------------------------------------------- +# Phase wiring: one forward per call, critic consumes the carried pred +# ---------------------------------------------------------------------- + + +def test_student_phase_forward_has_grad_and_sets_ctx_and_vis() -> None: + method = _make_method(slots=1, sample_type="ode", interval=5) + _stub_losses(method) + student = method.student + + _, outputs, metrics = method.single_train_step({"x": torch.ones(1)}, iteration=5) + + assert metrics["update_student"] == 1.0 + assert metrics["rollout_step"] == 0.0 + main = student.predict_calls[-1] + assert main["grad_enabled"] is True + assert main["attn_kind"] == "vsa" + assert main["timestep"] == float(_GRID[0]) + student_ctx = outputs["_fv_backward"]["student_ctx"] + assert student_ctx[0] is student.last_batch.timesteps + assert student_ctx[1] == "vsa-metadata" + assert outputs["_fv_backward"]["critic_ctx"] is None + vis = student.last_batch.dmd_latent_vis_dict + torch.testing.assert_close(vis["generator_timestep"], torch.tensor([float(_GRID[0])])) + torch.testing.assert_close(vis["generator_pred_video"], student.last_pred) + assert "generator_timestep" in method.latent_vis + + +def test_critic_phase_forward_is_no_grad_and_feeds_carried_pred() -> None: + method = _make_method(slots=1, sample_type="ode", interval=5) + critic_preds = _stub_losses(method) + student = method.student + + _, outputs, metrics = method.single_train_step({"x": torch.ones(1)}, iteration=1) + + assert metrics["update_student"] == 0.0 + assert metrics["rollout_step"] == 0.0 + main = student.predict_calls[-1] + assert main["grad_enabled"] is False + assert main["attn_kind"] == "vsa" + # The critic is fit on the same simulated state the student trains on: + # it receives this call's prediction rather than rolling its own. + assert len(critic_preds) == 1 + assert critic_preds[0] is student.last_pred + assert outputs["_fv_backward"]["critic_ctx"] == "critic-ctx" + assert outputs["_fv_backward"]["student_ctx"] is None + # Both phases advance the trajectory. + assert method._carry_slots[0] is not None + assert method._carry_slots[0]["rung"] == 1 + assert "generator_timestep" in student.last_batch.dmd_latent_vis_dict + + +def test_carried_advance_state_is_detached_and_matches_ode_math() -> None: + method = _make_method(slots=1, sample_type="ode", interval=1) + _stub_losses(method) + + method.single_train_step({"x": torch.ones(1)}, iteration=0) + carried = method._carry_slots[0] + assert carried is not None + assert carried["rung"] == 1 + assert not carried["state"].requires_grad + # offset 0: the main forward ran on the fresh-noise state drawn from the + # seeded generator; replay one ODE hop with the fake student's flow. + state0 = torch.randn(_LATENT_SHAPE, generator=torch.Generator().manual_seed(0)) + torch.testing.assert_close(carried["state"], _replay_ode_walk(state0, _GRID, 1)) + + +# ---------------------------------------------------------------------- +# (6) Knob parsing and validation; default off preserves existing behavior +# ---------------------------------------------------------------------- + + +def _parse_only( + config: dict, + *, + grad_accum: int = 1, + streams: tuple[int, int] = (0, 1), + native_shape_bucketing: bool = False, +): + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", dict(config)) + object.__setattr__(method, "student", _CarryStudent()) + object.__setattr__( + method, + "training_config", + SimpleNamespace( + loop=SimpleNamespace(gradient_accumulation_steps=grad_accum), + distributed=SimpleNamespace(sp_size=1), + data=SimpleNamespace(native_shape_bucketing=native_shape_bucketing), + ), + ) + object.__setattr__(method, "_rollout_mode", method._parse_rollout_mode()) + object.__setattr__(method, "_rollout_carry_rank_world", lambda: streams) + return method._parse_rollout_carry() + + +def test_rollout_carry_defaults_off() -> None: + knobs = _parse_only({ + "rollout_mode": "simulate", + "dmd_denoising_steps": _GRID, + }) + assert knobs == (False, 0, "sde") + + +def test_rollout_carry_defaults_slots_to_grad_accum() -> None: + knobs = _parse_only( + { + "rollout_mode": "simulate", + "rollout_carry": True, + "rollout_sample_type": "ode", + "dmd_denoising_steps": _GRID, + "generator_update_interval": 5, + }, + grad_accum=3, + ) + assert knobs == (True, 3, "ode") + + +def test_rollout_carry_requires_simulate_mode() -> None: + with pytest.raises(ValueError, match="rollout_mode: simulate"): + _parse_only({ + "rollout_mode": "data_latent", + "rollout_carry": True, + "dmd_denoising_steps": _GRID, + }) + + +def test_rollout_carry_slots_must_match_grad_accum() -> None: + with pytest.raises(ValueError, match="gradient_accumulation_steps"): + _parse_only( + { + "rollout_mode": "simulate", + "rollout_carry": True, + "rollout_carry_slots": 4, + "dmd_denoising_steps": _GRID, + }, + grad_accum=2, + ) + + +def test_carry_knobs_require_rollout_carry_enabled() -> None: + with pytest.raises(ValueError, match="rollout_carry: true"): + _parse_only({ + "rollout_mode": "simulate", + "rollout_carry_slots": 2, + "dmd_denoising_steps": _GRID, + }) + with pytest.raises(ValueError, match="rollout_carry: true"): + _parse_only({ + "rollout_mode": "simulate", + "rollout_sample_type": "ode", + "dmd_denoising_steps": _GRID, + }) + + +def test_rollout_sample_type_rejects_unknown_values() -> None: + with pytest.raises(ValueError, match="ode, sde"): + _parse_only({ + "rollout_mode": "simulate", + "rollout_carry": True, + "rollout_sample_type": "euler", + "dmd_denoising_steps": _GRID, + }) + + +def test_ode_requires_student_extract_eps() -> None: + method = object.__new__(DMD2Method) + object.__setattr__( + method, + "method_config", + { + "rollout_mode": "simulate", + "rollout_carry": True, + "rollout_sample_type": "ode", + "dmd_denoising_steps": _GRID, + }, + ) + object.__setattr__(method, "student", SimpleNamespace()) + object.__setattr__( + method, + "training_config", + SimpleNamespace( + loop=SimpleNamespace(gradient_accumulation_steps=1), + distributed=SimpleNamespace(sp_size=1), + ), + ) + object.__setattr__(method, "_rollout_mode", "simulate") + with pytest.raises(ValueError, match="extract_eps"): + method._parse_rollout_carry() + + +def test_coverage_guard_rejects_uncovered_rung_phases() -> None: + # gcd(4, 2) = 2 phase classes: one stream cannot cover both. + with pytest.raises(ValueError, match="cannot cover"): + DMD2Method._validate_rollout_carry_coverage(streams=1, grid_len=4, interval=2) + DMD2Method._validate_rollout_carry_coverage(streams=2, grid_len=4, interval=2) + # The H3 recipe (4-rung grid, interval 5) has gcd 1 and always passes. + DMD2Method._validate_rollout_carry_coverage(streams=1, grid_len=4, interval=5) + + +def test_native_shape_coverage_does_not_count_rank_staggering() -> None: + config = { + "rollout_mode": "simulate", + "rollout_carry": True, + "rollout_carry_slots": 1, + "rollout_sample_type": "ode", + "dmd_denoising_steps": _GRID, + "generator_update_interval": 2, + } + assert _parse_only(config, streams=(0, 2)) == (True, 1, "ode") + with pytest.raises(ValueError, match="cannot cover"): + _parse_only(config, streams=(0, 2), native_shape_bucketing=True) + + +# ---------------------------------------------------------------------- +# Integration: full carried steps on the real H3 CPU trio +# ---------------------------------------------------------------------- + + +def _build_carry_trio(monkeypatch: pytest.MonkeyPatch, *, interval: int) -> DMD2Method: + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import ( + _FIXTURE, + _make_model, + ) + from fastvideo.train.utils.config import load_run_config + + config = load_run_config(str(_FIXTURE)) + config.method["rollout_mode"] = "simulate" + config.method["generator_update_interval"] = interval + config.method["rollout_carry"] = True + config.method["rollout_carry_slots"] = 1 + config.method["rollout_sample_type"] = "ode" + student = _make_model(monkeypatch, config.training, scale=1.0) + teacher = _make_model(monkeypatch, config.training, trainable=False, scale=0.5) + critic = _make_model(monkeypatch, config.training, scale=0.25) + student.init_preprocessors = lambda training_config: None + method = DMD2Method( + cfg=config, + role_models={ + "student": student, + "teacher": teacher, + "critic": critic, + }, + ) + method.cuda_generator = torch.Generator(device="cpu").manual_seed(0) + return method + + +def test_full_carried_student_step_on_real_h3_trio(monkeypatch: pytest.MonkeyPatch) -> None: + """End-to-end carried step: real prepare_batch, losses, and ODE advance.""" + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import _raw_batch + + method = _build_carry_trio(monkeypatch, interval=1) + student = method.student + + loss_map, outputs, metrics = method.single_train_step(_raw_batch(seed=1), iteration=0) + + assert metrics["update_student"] == 1.0 + assert metrics["rollout_step"] == 0.0 # (rank 0 * 1 + slot 0) % 3 + assert torch.isfinite(loss_map["total_loss"]) + assert loss_map["generator_loss"].item() > 0.0 + assert "generator_loss_video" in metrics and "generator_loss_audio" in metrics + carried = method._carry_slots[0] + assert carried is not None and carried["rung"] == 1 + assert carried["state"].shape == student.prepare_batch( + carried["raw_batch"], + generator=torch.Generator().manual_seed(0), + latents_source="zeros", + ).latents.shape + assert not carried["state"].requires_grad + + method.backward(loss_map, outputs) + assert student.transformer.scale.grad is not None + assert torch.isfinite(student.transformer.scale.grad) + + # Mid-walk call: a new prompt arrives but the student must still be + # conditioned on the trajectory's adopted text. + adopted = _raw_batch(seed=1) + adopted_text = adopted["text_embedding"][:, :2].to(torch.bfloat16) + method.single_train_step(_raw_batch(seed=9), iteration=1) + torch.testing.assert_close(student.transformer.last_encoder_hidden_states, adopted_text) + + +def test_full_carried_critic_step_on_real_h3_trio(monkeypatch: pytest.MonkeyPatch) -> None: + from fastvideo.tests.train.methods.test_minimax_h3_dmd2 import _raw_batch + + method = _build_carry_trio(monkeypatch, interval=5) + critic = method.critic + + loss_map, outputs, metrics = method.single_train_step(_raw_batch(seed=1), iteration=1) + + assert metrics["update_student"] == 0.0 + assert metrics["rollout_step"] == 0.0 + assert loss_map["generator_loss"].item() == 0.0 + assert loss_map["fake_score_loss"].item() > 0.0 + assert "fake_score_loss_video" in metrics and "fake_score_loss_audio" in metrics + assert method._carry_slots[0] is not None and method._carry_slots[0]["rung"] == 1 + + method.backward(loss_map, outputs) + assert critic.transformer.scale.grad is not None + assert torch.isfinite(critic.transformer.scale.grad) + assert method.student.transformer.scale.grad is None diff --git a/fastvideo/tests/train/methods/test_dmd2_timestep_bounds.py b/fastvideo/tests/train/methods/test_dmd2_timestep_bounds.py index 8f100f3874..ba9f30af0a 100644 --- a/fastvideo/tests/train/methods/test_dmd2_timestep_bounds.py +++ b/fastvideo/tests/train/methods/test_dmd2_timestep_bounds.py @@ -4,6 +4,7 @@ from types import SimpleNamespace import pytest +import torch from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method @@ -18,7 +19,7 @@ def _method_with_ratios(min_ratio, max_ratio) -> DMD2Method: return method -def test_dmd2_score_timestep_bounds_match_legacy_recipe() -> None: +def test_dmd2_score_timestep_bounds_apply_ratios() -> None: method = _method_with_ratios(0.02, 0.98) assert method._parse_score_timestep_bounds() == (20, 980) @@ -39,3 +40,64 @@ def test_dmd2_score_timestep_bounds_reject_invalid_ranges( method = _method_with_ratios(min_ratio, max_ratio) with pytest.raises(ValueError, match="0 <= min <= max <= 1"): method._parse_score_timestep_bounds() + + +@pytest.mark.parametrize("warp_max", [0.0, -0.1, 1.1]) +def test_dmd2_score_timestep_warp_max_rejects_invalid_endpoint(warp_max: float) -> None: + method = _method_with_ratios(0.001, 0.999) + lo, hi = method._parse_score_timestep_bounds() + object.__setattr__(method, "_score_min_timestep", lo) + object.__setattr__(method, "_score_max_timestep", hi) + method.method_config["score_timestep_warp_max"] = warp_max + with pytest.raises(ValueError, match="0 < max <= 1"): + method._parse_score_timestep_warp_max() + + +def test_dmd2_score_timestep_warp_max_must_cover_upper_bound() -> None: + method = _method_with_ratios(0.001, 1.0) + lo, hi = method._parse_score_timestep_bounds() + object.__setattr__(method, "_score_min_timestep", lo) + object.__setattr__(method, "_score_max_timestep", hi) + method.method_config["score_timestep_warp_max"] = 0.999 + with pytest.raises(ValueError, match="must not exceed"): + method._parse_score_timestep_warp_max() + + +def test_dmd2_score_timestep_continuous_requires_bool() -> None: + method = _method_with_ratios(0.001, 0.999) + method.method_config["score_timestep_continuous"] = 1 + with pytest.raises(ValueError, match="must be a bool"): + method._parse_score_timestep_continuous() + + +def _sampler(min_ratio: float, max_ratio: float, shift: float) -> DMD2Method: + method = _method_with_ratios(min_ratio, max_ratio) + lo, hi = method._parse_score_timestep_bounds() + object.__setattr__(method, "_score_min_timestep", lo) + object.__setattr__(method, "_score_max_timestep", hi) + object.__setattr__(method, "_score_timestep_shift", shift) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(0)) + method.student.shift_and_clamp_timestep = lambda t: t + return method + + +def test_uniform_sampler_draws_in_bounds_without_boundary_atoms() -> None: + """shift=1 samples uniformly inside the configured integer bounds.""" + method = _sampler(0.02, 0.98, shift=1.0) + device = torch.device("cpu") + draws = torch.cat([method._sample_score_timestep(device) for _ in range(4000)]) + + assert int(draws.min()) >= 20 + assert int(draws.max()) <= 980 + # Uniform over 961 values yields about four samples per endpoint. + assert int((draws == 20).sum()) < 20 + assert int((draws == 980).sum()) < 20 + + +def test_shifted_sampler_respects_bounds() -> None: + method = _sampler(0.005, 0.98, shift=12.0) + device = torch.device("cpu") + draws = torch.cat([method._sample_score_timestep(device) for _ in range(2000)]) + + assert int(draws.min()) >= 5 + assert int(draws.max()) <= 980 diff --git a/fastvideo/tests/train/methods/test_dmd2_vsd_normalizer.py b/fastvideo/tests/train/methods/test_dmd2_vsd_normalizer.py new file mode 100644 index 0000000000..e7f0102178 --- /dev/null +++ b/fastvideo/tests/train/methods/test_dmd2_vsd_normalizer.py @@ -0,0 +1,47 @@ +# SPDX-License-Identifier: Apache-2.0 +"""The VSD normalizer must stay finite when |gen - real| degenerates to 0.""" +from __future__ import annotations + +from types import SimpleNamespace + +import torch + +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method + + +class _EchoTeacher: + """Returns the generator's own prediction: |gen - real| == 0 exactly.""" + + def __init__(self, gen: torch.Tensor): + self._gen = gen + + def predict_x0(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + return self._gen.detach().clone() + + +class _OffsetCritic(_EchoTeacher): + + def predict_x0(self, noisy, timestep, batch, *, conditional, cfg_uncond, attn_kind): + return self._gen.detach().clone() + 1.0 + + +def test_degenerate_denominator_yields_finite_loss_and_grad() -> None: + gen = torch.randn(1, 16, generator=torch.Generator().manual_seed(3)).requires_grad_(True) + + method = object.__new__(DMD2Method) + object.__setattr__(method, "method_config", {}) + object.__setattr__(method, "cuda_generator", torch.Generator().manual_seed(0)) + object.__setattr__(method, "_cfg_uncond", None) + object.__setattr__(method, "student", SimpleNamespace(add_noise=lambda clean, noise, t: clean)) + object.__setattr__(method, "teacher", _EchoTeacher(gen)) + object.__setattr__(method, "critic", _OffsetCritic(gen)) + object.__setattr__(method, "_sample_score_timestep", lambda device: torch.tensor([500])) + + batch = SimpleNamespace(dmd_latent_vis_dict={}) + loss, metrics = method._dmd_loss(gen, batch) + + assert torch.isfinite(loss) + assert metrics == {} + loss.backward() + assert gen.grad is not None + assert torch.isfinite(gen.grad).all() diff --git a/fastvideo/tests/train/methods/test_minimax_h3_dmd2.py b/fastvideo/tests/train/methods/test_minimax_h3_dmd2.py new file mode 100644 index 0000000000..6ee19600bb --- /dev/null +++ b/fastvideo/tests/train/methods/test_minimax_h3_dmd2.py @@ -0,0 +1,1154 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU contract tests for MiniMax H3 DMD2 distillation. + +Covers the packed dual-modality adapter (MiniMaxH3DMDModel) and one full +DMD2Method.single_train_step on a tiny CPU trio: student rollout, critic +flow-matching loss, generator DMD loss, both backwards, both optimizers. +""" + +import math +from pathlib import Path +from types import SimpleNamespace + +import pytest +import torch +import yaml + +from fastvideo.attention.backends.video_sparse_attn_h3 import MiniMaxH3VSAMetadata +from fastvideo.forward_context import get_forward_context +from fastvideo.pipelines.basic.minimax_h3.packing import audio_latent_num_frames, video_latent_num_frames +from fastvideo.platforms import AttentionBackendEnum +from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method +from fastvideo.train.models.minimax_h3 import MiniMaxH3DMDModel, MiniMaxH3Model +from fastvideo.train.models.minimax_h3.minimax_h3 import shift_noise_amount +from fastvideo.train.utils.config import load_run_config + +_FIXTURE = Path(__file__).resolve().parent.parent / "fixtures" / "minimax_h3_dmd2_min.yaml" +_REPO_ROOT = Path(__file__).resolve().parents[4] +_EXPERIMENT_CONFIG = (_REPO_ROOT / "examples/train/configs/distribution_matching/minimax_h3/dmd2_sp1_fsdp40_nuva_v9_dataforce_vsa64.yaml") +_V10_EXPERIMENT_CONFIG = (_REPO_ROOT / "examples/train/configs/distribution_matching/minimax_h3/dmd2_sp1_fsdp64_v10_dataonly_mixed_vsa64.yaml") +_V10_MAXSHAPE_CONFIG = (_REPO_ROOT / "examples/train/configs/distribution_matching/minimax_h3/dmd2_sp1_fsdp64_v10_maxshape_gate_vsa64.yaml") +_V10_PREPARE_LAUNCHER = _REPO_ROOT / "examples/train/slurm/prepare_h3_dmd2_v10_slinky.sh" +_H3_SBATCH = _REPO_ROOT / "examples/train/slurm/dmd2_32xgb200.sbatch" +_V10_GATED_LAUNCHER = _REPO_ROOT / "scripts/train/run_h3_v10_gated.sh" +_V10_MAXSHAPE_RUNNER = _REPO_ROOT / "scripts/train/run_h3_v10_maxshape_gate.sh" +_V10_KERNEL_GATE = _REPO_ROOT / "scripts/train/gate_h3_v10_kernel.sh" +_V10_KERNEL_REBUILD = _REPO_ROOT / "scripts/train/rebuild_h3_v10_kernel.sh" +_V10_KERNEL_RECEIPT_HELPER = _REPO_ROOT / "scripts/train/h3_v10_kernel_receipt.py" + +# Fixture geometry: video latents [1, 24, 2, 4, 4] and audio latents +# [1, 2, 32, 8]; the packed adapter stores video-major [1, T, C, H, W]. +_VIDEO_SHAPE = (1, 2, 24, 4, 4) +_AUDIO_SHAPE = (1, 2, 32, 8) +_PACKED_NUMEL = math.prod(_VIDEO_SHAPE) + math.prod(_AUDIO_SHAPE) + + +class _TinyJointTransformer(torch.nn.Module): + """Scale packed H3 rows with one trainable parameter.""" + + patch_size = (1, 2, 2) + + def __init__(self, scale: float = 1.0) -> None: + super().__init__() + self.scale = torch.nn.Parameter(torch.tensor(scale)) + self.last_encoder_hidden_states: torch.Tensor | None = None + self.last_attn_metadata = None + + def forward(self, **kwargs): + self.last_encoder_hidden_states = kwargs["encoder_hidden_states"] + self.last_attn_metadata = get_forward_context().attn_metadata + return ( + kwargs["hidden_states"] * self.scale, + kwargs["audio_hidden_states"] * self.scale, + ) + + +def _make_model( + monkeypatch: pytest.MonkeyPatch, + training_config, + *, + trainable: bool = True, + scale: float = 1.0, +) -> MiniMaxH3DMDModel: + monkeypatch.setattr(MiniMaxH3Model, "device", property(lambda _self: torch.device("cpu"))) + model = MiniMaxH3DMDModel.__new__(MiniMaxH3DMDModel) + model._trainable = trainable + model.transformer = _TinyJointTransformer(scale) + model.training_config = training_config + model.sp_group = None + model.attention_backend = None + return model + + +def _tiny_training_config(): + return SimpleNamespace( + data=SimpleNamespace( + num_latent_t=2, + num_frames=5, + num_height=64, + num_width=64, + ), + distributed=SimpleNamespace(sp_size=1), + vsa_sparsity=0.0, + vsa_tile_size=256, + ) + + +def _raw_batch(seed: int = 1) -> dict[str, torch.Tensor]: + generator = torch.Generator().manual_seed(seed) + return { + "vae_latent": torch.randn(1, 24, 2, 4, 4, generator=generator), + "audio_latent": torch.randn(1, 2, 32, 8, generator=generator), + "text_embedding": torch.randn(1, 4, 5120, generator=generator), + "text_attention_mask": torch.tensor([[1, 1, 0, 0]], dtype=torch.float32), + } + + +def _native_raw_batch( + width: int, + height: int, + num_frames: int, + *, + seed: int = 1, +) -> dict: + generator = torch.Generator().manual_seed(seed) + return { + "vae_latent": torch.randn( + 1, + 24, + video_latent_num_frames(num_frames), + height // 16, + width // 16, + generator=generator, + dtype=torch.bfloat16, + ), + "audio_latent": torch.randn( + 1, + 2, + 32, + audio_latent_num_frames(num_frames), + generator=generator, + dtype=torch.bfloat16, + ), + "text_embedding": torch.randn(1, 4, 5120, generator=generator), + "text_attention_mask": torch.tensor([[1, 1, 0, 0]], dtype=torch.float32), + "_shape_bucket_id": f"bucket={width}x{height}-{num_frames}f", + "info_list": [{ + "width": width, + "height": height, + "num_frames": num_frames, + "fps": 24.0, + "audio_sample_rate": 32_000, + }], + } + + +def _build_method( + monkeypatch: pytest.MonkeyPatch, + *, + rollout_mode: str, + generator_update_interval: int = 1, +) -> DMD2Method: + config = load_run_config(str(_FIXTURE)) + config.method["rollout_mode"] = rollout_mode + config.method["generator_update_interval"] = generator_update_interval + # Distinct role scales keep the critic-vs-teacher DMD gradient non-zero. + student = _make_model(monkeypatch, config.training, scale=1.0) + teacher = _make_model(monkeypatch, config.training, trainable=False, scale=0.5) + critic = _make_model(monkeypatch, config.training, scale=0.25) + student.init_preprocessors = lambda training_config: None + method = DMD2Method( + cfg=config, + role_models={ + "student": student, + "teacher": teacher, + "critic": critic, + }, + ) + method.cuda_generator = torch.Generator(device="cpu").manual_seed(0) + return method + + +# ---------------------------------------------------------------------- +# Core gate: one full DMD2 train step on CPU +# ---------------------------------------------------------------------- + + +@pytest.mark.parametrize("rollout_mode", ["data_latent", "simulate"]) +def test_dmd2_student_iteration_updates_only_student( + monkeypatch: pytest.MonkeyPatch, + rollout_mode: str, +) -> None: + """Student iterations do not train or step the critic.""" + method = _build_method(monkeypatch, rollout_mode=rollout_mode) + student = method.student + teacher = method.teacher + critic = method.critic + + loss_map, outputs, metrics = method.single_train_step(_raw_batch(), iteration=0) + + assert metrics["update_student"] == 1.0 + for key in ("total_loss", "generator_loss", "fake_score_loss"): + assert torch.isfinite(loss_map[key]), key + assert loss_map["generator_loss"].item() > 0.0 + assert loss_map["fake_score_loss"].item() == 0.0 + torch.testing.assert_close(loss_map["total_loss"], loss_map["generator_loss"]) + assert "generator_pred_video" in method.latent_vis + assert "real_score_pred_video" in method.latent_vis + assert "faker_score_pred_video" in method.latent_vis + + method.backward(loss_map, outputs) + assert student.transformer.scale.grad is not None + assert torch.isfinite(student.transformer.scale.grad) + assert critic.transformer.scale.grad is None + assert teacher.transformer.scale.grad is None + assert method.get_optimizers(0) == [method._student_optimizer] + assert method.get_lr_schedulers(0) == [method._student_lr_scheduler] + assert method.get_grad_clip_targets(0) == {"student": student.transformer} + + student_before = student.transformer.scale.detach().clone() + critic_before = critic.transformer.scale.detach().clone() + method.optimizers_schedulers_step(0) + assert student.transformer.scale.detach() != student_before + torch.testing.assert_close(critic.transformer.scale.detach(), critic_before) + + +def test_dmd2_critic_iteration_updates_only_critic(monkeypatch: pytest.MonkeyPatch) -> None: + """Off-interval iterations do not train or step the student.""" + method = _build_method( + monkeypatch, + rollout_mode="data_latent", + generator_update_interval=5, + ) + + student = method.student + critic = method.critic + loss_map, outputs, metrics = method.single_train_step(_raw_batch(), iteration=1) + + assert metrics["update_student"] == 0.0 + assert loss_map["generator_loss"].item() == 0.0 + assert loss_map["fake_score_loss"].item() > 0.0 + torch.testing.assert_close(loss_map["total_loss"], loss_map["fake_score_loss"]) + method.backward(loss_map, outputs) + assert student.transformer.scale.grad is None + assert critic.transformer.scale.grad is not None + assert torch.isfinite(critic.transformer.scale.grad) + assert method.get_optimizers(1) == [method._critic_optimizer] + assert method.get_lr_schedulers(1) == [method._critic_lr_scheduler] + assert method.get_grad_clip_targets(1) == {"critic": critic.transformer} + + student_before = student.transformer.scale.detach().clone() + critic_before = critic.transformer.scale.detach().clone() + method.optimizers_schedulers_step(1) + torch.testing.assert_close(student.transformer.scale.detach(), student_before) + assert critic.transformer.scale.detach() != critic_before + + +def test_dmd2_five_step_cadence_and_resume_state(monkeypatch: pytest.MonkeyPatch) -> None: + method = _build_method( + monkeypatch, + rollout_mode="simulate", + generator_update_interval=5, + ) + + assert [method._should_update_student(i) for i in range(1, 6)] == [ + False, + False, + False, + False, + True, + ] + method.method_config["generator_update_interval"] = 0 + with pytest.raises(ValueError, match="must be positive"): + method._should_update_student(0) + + method.seed_optimizer_state_for_resume() + for optimizer in (method._student_optimizer, method._critic_optimizer): + assert optimizer.state + assert all("exp_avg" in state for state in optimizer.state.values()) + + method.method_config.pop("generator_update_interval") + assert [method._should_update_student(i) for i in range(1, 6)] == [ + False, + False, + False, + False, + True, + ] + + +# ---------------------------------------------------------------------- +# Packed dual-modality adapter units +# ---------------------------------------------------------------------- + + +def test_packed_adapter_roundtrip_and_prepare_batch(monkeypatch: pytest.MonkeyPatch) -> None: + """Verify pack/unpack inversion and packed clean latents in the batch.""" + model = _make_model(monkeypatch, _tiny_training_config()) + video = torch.randn(_VIDEO_SHAPE) + audio = torch.randn(_AUDIO_SHAPE) + + packed = model.pack_latents(video, audio) + assert packed.shape == (1, _PACKED_NUMEL) + video_out, audio_out = model.unpack_latents(packed) + torch.testing.assert_close(video_out, video) + torch.testing.assert_close(audio_out, audio) + + raw_batch = _raw_batch() + batch = model.prepare_batch( + raw_batch, + generator=torch.Generator().manual_seed(7), + ) + assert batch.latents.shape == (1, _PACKED_NUMEL) + video_clean, audio_clean = model.unpack_latents(batch.latents) + torch.testing.assert_close( + video_clean, + raw_batch["vae_latent"].permute(0, 2, 1, 3, 4).to(torch.bfloat16), + ) + torch.testing.assert_close(audio_clean, batch.audio_latents) + + +def test_native_layout_is_batch_local_across_successive_shapes(monkeypatch: pytest.MonkeyPatch) -> None: + """A later shape must not mutate how an earlier packed tensor is split.""" + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + + first = model.prepare_batch( + _native_raw_batch(64, 64, 5), + generator=torch.Generator().manual_seed(3), + ) + first_packed = first.latents.clone() + first_layout = first.minimax_h3_dmd_layout + second = model.prepare_batch( + _native_raw_batch(96, 64, 22), + generator=torch.Generator().manual_seed(4), + ) + + assert first.minimax_h3_dmd_layout is first_layout + assert first_layout != second.minimax_h3_dmd_layout + first_video, first_audio = model.unpack_latents(first_packed, layout=first_layout) + assert first_video.shape == (1, 2, 24, 4, 4) + assert first_audio.shape == (1, 2, 32, 8) + second_video, second_audio = model.unpack_latents(second.latents, layout=second.minimax_h3_dmd_layout) + assert second_video.shape == (1, 7, 24, 4, 6) + assert second_audio.shape == (1, 2, 32, 37) + + noise = torch.zeros_like(second.latents) + mixed = model.add_noise_for_batch(second.latents, noise, torch.tensor([500]), second) + assert mixed.shape == second.latents.shape + slices = dict(model.modality_slices_for_batch(second)) + assert slices["video"].stop == math.prod(second_video.shape) + assert slices["audio"].stop == second.latents.shape[1] + + prediction = model.predict_noise( + second.latents, + torch.tensor([500]), + second, + conditional=True, + ) + torch.testing.assert_close(prediction, -second.latents) + + +def test_dmd_losses_follow_successive_native_modality_slices(monkeypatch: pytest.MonkeyPatch) -> None: + """Student and critic phases both use the active batch's exact split.""" + method = _build_method( + monkeypatch, + rollout_mode="data_latent", + generator_update_interval=2, + ) + method.training_config.data.native_shape_bucketing = True + method.method_config["modality_loss_weights"] = {"video": 0.25, "audio": 2.0} + + student_losses, _, student_metrics = method.single_train_step( + _native_raw_batch(64, 64, 5), + iteration=2, + ) + critic_losses, _, critic_metrics = method.single_train_step( + _native_raw_batch(96, 64, 22), + iteration=3, + ) + + assert student_metrics["update_student"] == 1.0 + assert {"generator_loss_video", "generator_loss_audio"} <= student_metrics.keys() + assert torch.isfinite(student_losses["generator_loss"]) + assert critic_metrics["update_student"] == 0.0 + assert {"fake_score_loss_video", "fake_score_loss_audio"} <= critic_metrics.keys() + assert torch.isfinite(critic_losses["fake_score_loss"]) + + +@pytest.mark.parametrize( + ("width", "height", "num_frames"), + [ + (1344, 768, 124), + (768, 1344, 362), + (832, 480, 90), + (480, 832, 124), + ], +) +def test_native_validation_accepts_min_max_portrait_and_lowres( + monkeypatch: pytest.MonkeyPatch, + width: int, + height: int, + num_frames: int, +) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + raw = _native_raw_batch(width, height, num_frames) + + video, audio = model._resolve_clean_latents(raw, "data", torch.bfloat16, torch.device("cpu")) + + assert video.shape == ( + 1, + 24, + video_latent_num_frames(num_frames), + height // 16, + width // 16, + ) + assert audio.shape == (1, 2, 32, audio_latent_num_frames(num_frames)) + + +@pytest.mark.parametrize( + ("mutation", "message"), + [ + (lambda raw: raw["info_list"][0].update(width=96), "disagrees with row metadata"), + (lambda raw: raw["info_list"][0].update(fps=30.0), "24 fps clock"), + (lambda raw: raw["info_list"][0].update(audio_sample_rate=44_100), "32000 Hz"), + (lambda raw: raw.update(audio_latent=raw["audio_latent"][..., :-1]), "audio clock"), + (lambda raw: raw.update(vae_latent=raw["vae_latent"][:, :, :-1]), "vae_latent shape"), + ], +) +def test_native_validation_rejects_bucket_metadata_and_clock_mismatches( + monkeypatch: pytest.MonkeyPatch, + mutation, + message: str, +) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + raw = _native_raw_batch(64, 64, 5) + mutation(raw) + + with pytest.raises(ValueError, match=message): + model._resolve_clean_latents(raw, "data", torch.bfloat16, torch.device("cpu")) + + +def test_native_validation_requires_production_canvas_multiple(monkeypatch: pytest.MonkeyPatch) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + raw = _native_raw_batch(64, 64, 5) + raw["_shape_bucket_id"] = "bucket=80x64-5f" + raw["info_list"][0]["width"] = 80 + raw["vae_latent"] = torch.zeros(1, 24, 2, 4, 5) + + with pytest.raises(ValueError, match="canvas multiple 32"): + model._resolve_clean_latents(raw, "data", torch.bfloat16, torch.device("cpu")) + + +def test_legacy_fixed_data_path_still_truncates_to_config(monkeypatch: pytest.MonkeyPatch) -> None: + config = _tiny_training_config() + model = _make_model(monkeypatch, config) + raw = _raw_batch() + raw["vae_latent"] = torch.cat((raw["vae_latent"], raw["vae_latent"][:, :, :1]), dim=2) + raw["audio_latent"] = torch.cat((raw["audio_latent"], raw["audio_latent"][..., :2]), dim=-1) + + video, audio = model._resolve_clean_latents(raw, "data", torch.bfloat16, torch.device("cpu")) + + assert video.shape == (1, 24, 2, 4, 4) + assert audio.shape == (1, 2, 32, 8) + + +def test_simulate_zeros_follow_native_shape_bucket(monkeypatch: pytest.MonkeyPatch) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + + raw = {"_shape_bucket_id": "bucket=96x64-22f"} + video, audio = model._resolve_clean_latents(raw, "zeros", torch.bfloat16, torch.device("cpu")) + + assert video.shape == (1, 24, 7, 4, 6) + assert audio.shape == (1, 2, 32, 37) + + +def test_native_simulate_requires_shape_bucket(monkeypatch: pytest.MonkeyPatch) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + + with pytest.raises(ValueError, match="data-free batches require.*_shape_bucket_id"): + model._resolve_clean_latents({}, "zeros", torch.bfloat16, torch.device("cpu")) + + +@pytest.mark.parametrize( + ("bucket_id", "video_shape", "audio_frames"), + [ + ("bucket=1760x768-362f", (1, 24, 107, 48, 110), audio_latent_num_frames(362)), + ("bucket=768x1344-124f", (1, 24, 37, 84, 48), audio_latent_num_frames(124)), + ], +) +def test_native_simulate_zeros_cover_production_extremes( + monkeypatch: pytest.MonkeyPatch, + bucket_id: str, + video_shape: tuple[int, ...], + audio_frames: int, +) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + + video, audio = model._resolve_clean_latents( + {"_shape_bucket_id": bucket_id}, + "zeros", + torch.bfloat16, + torch.device("cpu"), + ) + + assert video.shape == video_shape + assert audio.shape == (1, 2, 32, audio_frames) + + +def test_native_simulate_rejects_non_aligned_canvas(monkeypatch: pytest.MonkeyPatch) -> None: + config = _tiny_training_config() + config.data.native_shape_bucketing = True + model = _make_model(monkeypatch, config) + + with pytest.raises(ValueError, match="canvas multiple 32"): + model._resolve_clean_latents( + {"_shape_bucket_id": "bucket=80x64-5f"}, + "zeros", + torch.bfloat16, + torch.device("cpu"), + ) + + +def test_packed_add_noise_applies_modality_shifts(monkeypatch: pytest.MonkeyPatch) -> None: + """One shared base timestep must map to two shifted noise amounts.""" + model = _make_model(monkeypatch, _tiny_training_config()) + clean = torch.ones(1, _PACKED_NUMEL) + noise = torch.zeros(1, _PACKED_NUMEL) + + torch.testing.assert_close( + model.add_noise(clean, noise, torch.tensor([0])), + clean, + ) + torch.testing.assert_close( + model.add_noise(clean, noise, torch.tensor([1000])), + noise, + ) + + mixed = model.add_noise(clean, noise, torch.tensor([500])) + video_mixed, audio_mixed = model.unpack_latents(mixed) + base = torch.tensor([0.5]) + torch.testing.assert_close( + video_mixed, + torch.full(_VIDEO_SHAPE, float(1.0 - shift_noise_amount(base, 12.0))), + ) + torch.testing.assert_close( + audio_mixed, + torch.full(_AUDIO_SHAPE, float(1.0 - shift_noise_amount(base, 3.0))), + ) + + +def test_packed_predict_noise_plumbs_timesteps_and_tolerates_vsa(monkeypatch: pytest.MonkeyPatch, ) -> None: + """Explicit method timesteps must rewrite both modality clean-times.""" + model = _make_model(monkeypatch, _tiny_training_config()) + batch = model.prepare_batch(_raw_batch(), generator=torch.Generator().manual_seed(7)) + noisy = torch.randn(1, _PACKED_NUMEL).to(torch.bfloat16) + timestep = torch.tensor([757], dtype=torch.long) + + # attn_kind="vsa" must silently mean dense (both metadata views are None). + prediction = model.predict_noise( + noisy, + timestep, + batch, + conditional=True, + attn_kind="vsa", + ) + + base = torch.tensor([0.757]) + torch.testing.assert_close( + batch.timesteps, + 1.0 - shift_noise_amount(base.double(), 12.0), + ) + torch.testing.assert_close( + batch.audio_timesteps, + 1.0 - shift_noise_amount(base.double(), 3.0), + ) + # The unit-scale transformer echoes packed rows, and the H3 wrapper + # negates them into noise-minus-clean form. + torch.testing.assert_close(prediction, -noisy) + + x0 = model.predict_x0(noisy, timestep, batch, conditional=True) + noisy_video, noisy_audio = model.unpack_latents(noisy) + sigma_video = shift_noise_amount(base.double(), 12.0) + sigma_audio = shift_noise_amount(base.double(), 3.0) + expected = model.pack_latents( + (noisy_video.double() + sigma_video * noisy_video.double()).to(torch.bfloat16), + (noisy_audio.double() + sigma_audio * noisy_audio.double()).to(torch.bfloat16), + ) + torch.testing.assert_close(x0, expected) + + +def test_uncond_forward_zeroes_text_and_guards_policies(monkeypatch: pytest.MonkeyPatch) -> None: + """Teacher-CFG unconditional forwards zero text; other policies fail fast.""" + model = _make_model(monkeypatch, _tiny_training_config()) + batch = model.prepare_batch(_raw_batch(), generator=torch.Generator().manual_seed(7)) + noisy = torch.randn(1, _PACKED_NUMEL).to(torch.bfloat16) + timestep = torch.tensor([500], dtype=torch.long) + + model.predict_noise( + noisy, + timestep, + batch, + conditional=False, + cfg_uncond={"text": "zero"}, + ) + assert torch.all(model.transformer.last_encoder_hidden_states == 0) + + model.predict_noise( + noisy, + timestep, + batch, + conditional=True, + cfg_uncond={"text": "zero"}, + ) + assert torch.any(model.transformer.last_encoder_hidden_states != 0) + + with pytest.raises(ValueError, match="cfg_uncond"): + model.predict_noise(noisy, timestep, batch, conditional=False) + with pytest.raises(ValueError, match="negative-prompt"): + model.set_requires_negative_conditioning(True) + model.set_requires_negative_conditioning(False) + + +# ---------------------------------------------------------------------- +# VSA-H3 wiring +# ---------------------------------------------------------------------- + + +def test_prepare_batch_builds_vsa_h3_metadata(monkeypatch: pytest.MonkeyPatch) -> None: + """The VSA-H3 role gets real packed-sequence metadata; dense view stays None.""" + tc = _tiny_training_config() + tc.vsa_sparsity = 0.35 + tc.vsa_tile_size = 256 + model = _make_model(monkeypatch, tc) + model.attention_backend = AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3 + + batch = model.prepare_batch(_raw_batch(), generator=torch.Generator().manual_seed(7)) + + meta = batch.attn_metadata_vsa + assert isinstance(meta, MiniMaxH3VSAMetadata) + assert batch.attn_metadata is None + assert meta.VSA_sparsity == pytest.approx(0.35) + # Packed layout: 2 text rows | 0 condition rows | 16 stereo audio rows | + # 8 video rows ([1, 24, 2, 4, 4] latents at patch (1, 2, 2)). + assert meta.total_seq_length == 26 + assert meta.num_prefix_tiles == 2 + assert meta.num_video_tiles == 1 + assert meta.variable_block_sizes.tolist() == [2, 16, 8] + assert int(meta.variable_block_sizes.sum()) == meta.total_seq_length + + +def test_predict_noise_routes_vsa_metadata_by_attn_kind(monkeypatch: pytest.MonkeyPatch) -> None: + """Student "vsa" forwards see the VSA metadata; "dense" forwards see None.""" + model = _make_model(monkeypatch, _tiny_training_config()) + model.attention_backend = AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3 + batch = model.prepare_batch(_raw_batch(), generator=torch.Generator().manual_seed(7)) + noisy = torch.randn(1, _PACKED_NUMEL).to(torch.bfloat16) + timestep = torch.tensor([757], dtype=torch.long) + + model.predict_noise(noisy, timestep, batch, conditional=True, attn_kind="vsa") + assert model.transformer.last_attn_metadata is batch.attn_metadata_vsa + assert isinstance(model.transformer.last_attn_metadata, MiniMaxH3VSAMetadata) + + model.predict_noise(noisy, timestep, batch, conditional=True, attn_kind="dense") + assert model.transformer.last_attn_metadata is None + + +def _ctor_training_config() -> SimpleNamespace: + """The minimum surface MiniMaxH3Model.__init__ reads from TrainingConfig.""" + return SimpleNamespace( + pipeline_config=SimpleNamespace(dit_config=SimpleNamespace(uniform_parameter_dtype=False)), + data=SimpleNamespace( + train_batch_size=1, + training_cfg_rate=0.0, + preprocessed_data_type="t2va", + ), + model=SimpleNamespace(enable_gradient_checkpointing_type=None), + ) + + +def test_per_role_attention_backend_override_resolves(monkeypatch: pytest.MonkeyPatch) -> None: + """Each role's backend reaches the loader; unsupported backends fail fast.""" + captured: dict[str, AttentionBackendEnum | None] = {} + + def _fake_load(**kwargs): + captured[kwargs["model_path"]] = kwargs["attention_backend"] + return _TinyJointTransformer() + + monkeypatch.setattr( + "fastvideo.train.models.minimax_h3.minimax_h3.load_module_from_path", + _fake_load, + ) + + student_config = _ctor_training_config() + student = MiniMaxH3DMDModel( + init_from="role/student", + training_config=student_config, + trainable=True, + attention_backend="VIDEO_SPARSE_ATTN_H3", + ) + teacher = MiniMaxH3DMDModel( + init_from="role/teacher", + training_config=_ctor_training_config(), + trainable=False, + attention_backend="FLASH_ATTN", + ) + + assert student.attention_backend is AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3 + assert student_config.pipeline_config.dit_config.uniform_parameter_dtype is False + assert teacher.attention_backend is AttentionBackendEnum.FLASH_ATTN + # load_module_from_path turns this request into the construction scope + # that binds the backend to the transformer's attention layers. + assert captured["role/student"] is AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3 + assert captured["role/teacher"] is AttentionBackendEnum.FLASH_ATTN + assert not any(p.requires_grad for p in teacher.transformer.parameters()) + + with pytest.raises(ValueError, match="supports the attention backends"): + MiniMaxH3DMDModel( + init_from="role/bad", + training_config=_ctor_training_config(), + attention_backend="VIDEO_SPARSE_ATTN", + ) + + +# ---------------------------------------------------------------------- +# Config contracts +# ---------------------------------------------------------------------- + + +def test_h3_dmd2_fixture_resolves_trio_contract() -> None: + """The fixture must wire the H3 DMD trio through the modular builder path.""" + config = load_run_config(str(_FIXTURE)) + + for role in ("student", "teacher", "critic"): + assert config.models[role]["_target_"] == ("fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel") + assert config.models["teacher"]["trainable"] is False + assert config.method["_target_"] == ("fastvideo.train.methods.distribution_matching.dmd2.DMD2Method") + assert config.training.data.preprocessed_data_type == "t2va" + + +def test_h3_dmd2_current_config_pins_recipe() -> None: + """The production config pins the intended H3 DMD2 knobs.""" + config = yaml.safe_load(_EXPERIMENT_CONFIG.read_text()) + method = config["method"] + training = config["training"] + + for role in ("student", "teacher", "critic"): + assert config["models"][role]["_target_"] == ("fastvideo.train.models.minimax_h3.MiniMaxH3DMDModel") + assert config["models"]["teacher"]["trainable"] is False + assert config["models"]["critic"]["trainable"] is True + assert method["_target_"] == ("fastvideo.train.methods.distribution_matching.dmd2.DMD2Method") + assert method["rollout_mode"] == "simulate" + assert method["rollout_carry"] is True + assert (method["rollout_carry_slots"] == training["loop"]["gradient_accumulation_steps"]) + # Global batch 128 = 32 DP x accum 4; the carry owns one stream per slot. + assert training["loop"]["gradient_accumulation_steps"] == 4 + assert method["rollout_sample_type"] == "ode" + # Historical v9 explicitly opts into its non-golden per-batch hybrid. + assert method["rollout_data_forcing"] is True + assert method["allow_mixed_rollout_regimes"] is True + assert method["generator_update_interval"] == 5 + assert method["real_score_guidance_scale"] == 1.0 + # FastGen h3_new grid: time_shift(linspace(0.999, 0, 5), 12) in base time. + assert method["dmd_denoising_steps"] == [999, 749, 500, 250] + assert "warp_denoising_step" not in method + # f_{1/2.4} == f_{5/12}: FastGen's shifted draw f_5(U) on the shift-12 clock. + assert method["score_timestep_shift"] == 2.4 + assert method["score_timestep_warp_max"] == 0.999 + assert method["score_timestep_continuous"] is True + assert method["min_timestep_ratio"] == 0.001 + assert method["max_timestep_ratio"] == 0.999 + assert method["fake_score_loss_space"] == "x0" + assert method["cfg_uncond"] == {"text": "zero"} + assert method["fake_score_learning_rate"] == training["optimizer"]["learning_rate"] + assert method["fake_score_betas"] == [0.9, 0.999] + assert method["fake_score_lr_scheduler"] == "constant" + assert training["optimizer"]["betas"] == [0.9, 0.999] + assert training["dit_precision"] == "fp32" + assert training["checkpoint"]["output_dir"].endswith("v9_dataforce_vsa64") + assert training["vsa"] == {"sparsity": 0.9, "tile_size": 64} + assert (config["models"]["student"]["attention_backend"] == "VIDEO_SPARSE_ATTN_H3") + assert config["models"]["teacher"]["attention_backend"] == "FLASH_ATTN" + assert config["models"]["critic"]["attention_backend"] == "FLASH_ATTN" + assert config["pipeline"]["dit_config"]["uniform_parameter_dtype"] is False + # Mixed loading is declared t2va (the superset schema); text-only roots + # yield empty latent columns and route to the carried walk. + assert training["data"]["preprocessed_data_type"] == "t2va" + data_paths = training["data"]["data_path"] + assert any("nuva_t2va" in str(path) for path in data_paths) + assert any("text_only" in str(path) for path in data_paths) + assert training["data"]["train_batch_size"] == 1 + assert training["data"]["training_cfg_rate"] == 0.0 + assert config["callbacks"]["grad_clip"]["max_grad_norm"] == 1.0 + # Regional compile of the dense roles; gated on the A/B verdict before + # launch (see the YAML's PENDING GATE note). + assert training["model"]["enable_torch_compile"] is True + # The compile A/B (vsa_gate/compile_ab/VERDICT.md) validated the flip + # with NO torch_compile_kwargs — the config must not add any. + assert "torch_compile_kwargs" not in training["model"] + + +def test_h3_dmd2_v10_config_pins_data_only_native_shape_recipe() -> None: + """V10 is a fresh 64-GPU, global-batch-64, all-real-latent lineage.""" + config = yaml.safe_load(_V10_EXPERIMENT_CONFIG.read_text()) + method = config["method"] + training = config["training"] + distributed = training["distributed"] + data = training["data"] + + assert method["rollout_mode"] == "data_latent" + for carry_key in ( + "rollout_carry", + "rollout_carry_slots", + "rollout_sample_type", + "rollout_data_forcing", + ): + assert carry_key not in method + assert method["dmd_denoising_steps"] == [999, 749, 500, 250] + assert method["fake_score_learning_rate"] == 2.0e-6 + assert training["optimizer"]["learning_rate"] == 2.0e-6 + + assert distributed == { + "num_gpus": 64, + "sp_size": 1, + "tp_size": 1, + "hsdp_replicate_dim": 1, + "hsdp_shard_dim": 64, + } + global_batch = (distributed["num_gpus"] // distributed["sp_size"] * data["train_batch_size"] * + training["loop"]["gradient_accumulation_steps"]) + assert global_batch == 64 + assert data["preprocessed_data_type"] == "t2va" + assert data["native_shape_bucketing"] is True + assert len(data["data_path"]) == 5 + assert all(path.startswith("/mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3/") + and path.endswith("/data") for path in data["data_path"]) + + checkpoint = training["checkpoint"] + assert checkpoint["output_dir"].endswith("v10_dataonly_mixed_vsa64") + assert training["loop"]["gradient_accumulation_steps"] == 1 + assert checkpoint["save_inference_checkpoint_on_validation"] is True + assert checkpoint["inference_checkpoint_role"] == "student" + assert checkpoint["inference_checkpoint_dtype"] == "bfloat16" + assert checkpoint["training_state_checkpointing_steps"] == 100 + assert checkpoint["require_complete_training_checkpoint"] is True + assert checkpoint["checkpointing_start_step"] == 100 + assert checkpoint["checkpoints_total_limit"] == 3 + assert training["tracker"]["run_name"] == "dmd2_sp1_fsdp64_v10_dataonly_mixed_vsa64" + assert training["model"]["enable_torch_compile"] is True + assert training["model"]["torch_compile_kwargs"] == { + "dynamic": True, + "recompile_limit": 32, + } + assert training["vsa"] == {"sparsity": 0.9, "tile_size": 64} + assert config["models"]["student"]["attention_backend"] == "VIDEO_SPARSE_ATTN_H3" + for role in ("teacher", "critic"): + assert config["models"][role]["attention_backend"] == "FLASH_ATTN" + + validation = config["callbacks"]["validation"] + assert validation["dataset_file"].endswith("/validation/heldout60.json") + assert validation["every_steps"] == 100 + assert validation["run_at_start"] is True + assert validation["sampling_steps"] == [4] + assert validation["use_record_dimensions"] is True + assert validation["max_record_num_frames"] == 345 + assert validation["use_validation_media_conditioning"] is False + + +def test_h3_dmd2_v10_prepare_launcher_pins_finalized_data_and_execution_clone() -> None: + """The non-submitting helper gates the dedicated clone and immutable dataset.""" + launcher = _V10_PREPARE_LAUNCHER.read_text() + + assert "/mnt/lustre/vlm-wlsaidhi/fastvideo/FastVideo-v10" in launcher + assert ('readonly CONFIG="${REPO}/examples/train/configs/distribution_matching/minimax_h3/' + 'dmd2_sp1_fsdp64_v10_dataonly_mixed_vsa64.yaml"') in launcher + assert 'readonly DATA_ROOT="/mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v3"' in launcher + assert 'readonly BASE_DATA_ROOT="/mnt/lustre/vlm-shared/h3_t2av_preprocessed/v10_mixed_native_v2"' in launcher + assert 'readonly VALIDATION_MANIFEST="${DATA_ROOT}/validation/heldout60.json"' in launcher + assert "readonly VALIDATION_MAX_RECORD_NUM_FRAMES=345" in launcher + assert ('readonly OUTPUT_DIR="/mnt/lustre/vlm-wlsaidhi/fastvideo/outputs/' + 'minimax_h3_dmd2_sp1_fsdp64_v10_dataonly_mixed_vsa64"') in launcher + assert "readonly NUM_NODES=16" in launcher + assert "readonly WORLD_SIZE=64" in launcher + assert "readonly HSDP_REPLICATE=1" in launcher + assert "readonly HSDP_SHARD=64" in launcher + assert "readonly GRADIENT_ACCUMULATION_STEPS=1" in launcher + assert "readonly GLOBAL_BATCH_SIZE=64" in launcher + assert 'PARTITION="${PARTITION:-hpc-rack-3}"' in launcher + assert "readonly EXPECTED_TRAINING_ROWS=60549" in launcher + assert "readonly EXPECTED_SHAPE_BUCKETS=87" in launcher + assert "readonly EXPECTED_SCHEDULED_ROWS=63424" in launcher + assert "readonly EXPECTED_PADDED_ROWS=2875" in launcher + assert "readonly EXPECTED_STEPS_PER_EPOCH=991" in launcher + assert "readonly EXPECTED_VALIDATION_ROWS=60" in launcher + assert "readonly EXPECTED_VALIDATION_DP_PADDING=4" in launcher + assert "readonly MIN_OUTPUT_FREE_BYTES=" in launcher + for removed_override in ("CONFIG", "DATA_ROOT", "VALIDATION_MANIFEST", "OUTPUT_DIR"): + assert f'${{{removed_override}:-' not in launcher + assert 'readonly REVIEWED_V10_COMMIT="7635a5295b027000a00f6d70789c5cb5886218c3"' in launcher + assert 'merge-base --is-ancestor "${REVIEWED_V10_COMMIT}" HEAD' in launcher + assert "actual_data_paths = training[\"data\"][\"data_path\"]" in launcher + assert 'actual_validation = validation["dataset_file"]' in launcher + assert "actual_validation_max_record_num_frames" in launcher + assert 'validation.get("sampling_steps") != [4]' in launcher + assert "actual_output = training[\"checkpoint\"][\"output_dir\"]" in launcher + assert "actual_topology != expected_topology" in launcher + assert "actual_global_batch_size != global_batch_size" in launcher + assert 'compile_kwargs.get("dynamic") is not True' in launcher + assert "kernel receipt source" in launcher + assert 'readonly MAXSHAPE_AUDIT_ROOT=' in launcher + assert 'audit_root.glob("job-*/RESULT.json")' in launcher + assert "fastvideo-h3-v10-maxshape-gate-v1" in launcher + assert "no successful final-commit 64-GPU 1760x768x362 capacity receipt" in launcher + assert "available_bytes < MIN_OUTPUT_FREE_BYTES" in launcher + assert "validated same-recipe pre-step-100 restart namespace" in launcher + assert 'checkpoint_dir_re = re.compile(r"^checkpoint-([0-9]+)$")' in launcher + assert 'dcp_metadata = checkpoint / "dcp" / ".metadata"' in launcher + assert 'marker != "complete\\n"' in launcher + assert 'expected_rng_names = {f"rng_state_rank{rank}.pt" for rank in range(world_size)}' in launcher + assert 'len(rng_paths) != world_size' in launcher + assert "incomplete training checkpoint is not newer than the latest strict fallback" in launcher + assert "required_inference_steps = set(range(0, latest_step + 1, 100))" in launcher + assert "missing unlimited-retention inference checkpoints" in launcher + assert "completed validation step {step} does not have all 64 DP-padded names" in launcher + assert 'saved_method.get("dmd_denoising_steps") != [999, 749, 500, 250]' in launcher + assert "belongs to a different data/heldout60 recipe" in launcher + assert "shard/tensor counts differ from its index" in launcher + assert "step-zero validation must preserve all 64 DP-padded four-forward names" in launcher + assert 'require_file "${DATA_ROOT}/READY.json"' in launcher + assert 'require_file "${DATA_ROOT}/DERIVATION_RECEIPT.json"' in launcher + assert 'require_file "${source_root}/READY.json"' in launcher + assert 'require_file "${source_root}/MANIFEST.json"' in launcher + assert 'require_file "${source_root}/MANIFEST_rows.jsonl"' in launcher + assert 'require_file "${source_root}/data/map_style_cache/file_info.pkl"' in launcher + assert "finalize_dataset.py" in launcher + assert "derive_filtered_dataset.py" in launcher + assert '"excluded_resolutions": ["576x576", "640x480", "832x480"]' in launcher + assert '"excluded_validation_rows": 4' in launcher + assert '"training_holdout_policy": "preserve_base_validation_conditioning_ids"' in launcher + assert '"validation_payload_path": "validation/heldout60.json"' in launcher + assert "minimax-h3-native-t2va-filtered-derivation-v2" in launcher + assert '"validation_summary_sha256": sha256(data_root / "validation" / "manifest.json")' in launcher + assert "validation DP padding does not repeat exactly the first four retained records" in launcher + assert "excluded resolution leaked into v3 cache" in launcher + assert "--verify-only" in launcher + assert 'require_file "${REPO}/scripts/train/run_h3_v10_gated.sh"' in launcher + assert "--nodes=%q" in launcher + assert "--no-requeue" in launcher + assert "--wrap" not in launcher + assert '"${REPO}/scripts/train/run_h3_v10_gated.sh" "${execution_commit}"' in launcher + assert 'git -C "${REPO}" status --porcelain' in launcher + assert "This helper never calls sbatch" in launcher + + sbatch = _H3_SBATCH.read_text() + assert "HSDP_REPLICATE * HSDP_SHARD != WORLD_SIZE" in sbatch + assert 'git -C "${REPO}" status --porcelain' in sbatch + assert "V10 SOURCE GATE FAILED: execution checkout is dirty" in sbatch + assert "export HOME=" not in sbatch + for runtime_variable in ( + "HF_HOME", + "XDG_CACHE_HOME", + "TORCH_HOME", + "TRITON_CACHE_DIR", + "TORCHINDUCTOR_CACHE_DIR", + "TORCH_EXTENSIONS_DIR", + "FLASHINFER_WORKSPACE_BASE", + "CUDA_CACHE_PATH", + "NUMBA_CACHE_DIR", + "WANDB_CONFIG_DIR", + "WANDB_CACHE_DIR", + "WANDB_DATA_DIR", + "NETRC", + ): + assert f"export {runtime_variable}=" in sbatch + + +def test_h3_dmd2_v10_maxshape_gate_matches_production_capacity_contract() -> None: + config = yaml.safe_load(_V10_MAXSHAPE_CONFIG.read_text()) + training = config["training"] + distributed = training["distributed"] + checkpoint = training["checkpoint"] + + assert distributed == { + "num_gpus": 64, + "sp_size": 1, + "tp_size": 1, + "hsdp_replicate_dim": 1, + "hsdp_shard_dim": 64, + } + assert training["data"]["data_path"].endswith("/v10_maxshape_64g/data") + assert training["data"]["native_shape_bucketing"] is True + assert training["data"]["train_batch_size"] == 1 + assert training["loop"] == {"max_train_steps": 2, "gradient_accumulation_steps": 1} + assert config["method"]["generator_update_interval"] == 2 + assert training["model"]["enable_torch_compile"] is True + assert training["model"]["torch_compile_kwargs"] == { + "dynamic": True, + "recompile_limit": 32, + } + assert checkpoint["resume_from_checkpoint"] == "latest" + assert checkpoint["save_inference_checkpoint_on_validation"] is False + assert checkpoint["training_state_checkpointing_steps"] == 0 + assert "validation" not in config.get("callbacks", {}) + + runner = _V10_MAXSHAPE_RUNNER.read_text() + assert '"${SLURM_JOB_NUM_NODES:-0}" != "16"' in runner + assert "/v10_mixed_native_v3/" in runner + assert "/v10_mixed_native_v2/" not in runner + assert '"${staged}" -ef "${source}"' in runner + assert "all_64_gpus_sampled" in runner + assert "critic_grad_finite_positive" in runner + assert "student_grad_finite_positive" in runner + assert "dense_teacher_and_critic_compiled" in runner + assert "vsa_grad_used_triton64" in runner + assert '"execution_commit_unchanged": observed_execution_commit == expected_execution_commit' in runner + assert '"execution_checkout_clean": execution_checkout_clean' in runner + assert '"execution_commit": expected_execution_commit' in runner + + sbatch = _H3_SBATCH.read_text() + assert 'TRAIN_LOG_ROOT="${H3_V10_TRAIN_LOG_ROOT:-${REPO}/examples/train/logs}"' in sbatch + + +def test_h3_dmd2_v10_gated_launcher_defers_requeue_until_all_gates_pass() -> None: + launcher = _V10_GATED_LAUNCHER.read_text() + + assert launcher.startswith("#!/bin/bash\n") + assert "#SBATCH --nodes=16" in launcher + assert "#SBATCH --gpus-per-node=4" in launcher + assert "#SBATCH --no-requeue" in launcher + assert "--wrap" not in launcher + assert '(( $# != 1 ))' in launcher + assert '[[ ! "$1" =~ ^[0-9a-f]{40}$ ]]' in launcher + assert 'readonly EXPECTED_V10_COMMIT="$1"' in launcher + assert 'requested commit ${EXPECTED_V10_COMMIT} != execution HEAD ${execution_commit}' in launcher + assert launcher.count("require_exact_execution_checkout") == 4 + assert 'MAXSHAPE_RECEIPT="${MAXSHAPE_ROOT}/audit/job-${SLURM_JOB_ID:' in launcher + assert 'if [[ ! -f "${MAXSHAPE_RECEIPT}" ]]' in launcher + assert '"job_id": job_id' in launcher + assert '"${MAXSHAPE_CONFIG}" "${EXPECTED_V10_COMMIT}" "${SLURM_JOB_ID}"' in launcher + assert 'bash "${REPO}/scripts/train/run_h3_v10_maxshape_gate.sh"' in launcher + preflight = 'bash "${REPO}/examples/train/slurm/prepare_h3_dmd2_v10_slinky.sh"' + enable_requeue = 'scontrol update JobId="${SLURM_JOB_ID}" Requeue=1' + production = 'exec bash "${REPO}/examples/train/slurm/dmd2_32xgb200.sbatch"' + assert launcher.index(preflight) < launcher.index(enable_requeue) < launcher.index(production) + assert 'scontrol update JobId="${SLURM_JOB_ID}" Requeue=0' in launcher + assert "export H3_V10_KERNEL_GATE=0" in launcher + assert "export GATE_TEST=0" in launcher + + +def test_h3_dmd2_v10_kernel_gate_pins_import_order_and_real_gpu_checks() -> None: + gated_launcher = _V10_GATED_LAUNCHER.read_text() + sbatch = _H3_SBATCH.read_text() + gate = _V10_KERNEL_GATE.read_text() + rebuild = _V10_KERNEL_REBUILD.read_text() + receipt_helper = _V10_KERNEL_RECEIPT_HELPER.read_text() + expected_pythonpath = "${KERNEL_PREFIX}:${FA4_OVERLAY}:${FA4_CUTLASS_PACKAGES}" + + assert f'export PYTHONPATH="{expected_pythonpath}"' in gated_launcher + assert "export H3_V10_KERNEL_GATE=1" in gated_launcher + assert 'if [ "${H3_V10_KERNEL_GATE}" = "1" ]' in sbatch + assert "scripts/train/gate_h3_v10_kernel.sh" in sbatch + assert "python' -m pytest" in sbatch + assert "test_vsa_triton_backward_scale.py" in gate + assert "test_forward_matches_reference[64]" in gate + assert "test_real_sm100a_no_grad_route_receipt" in gate + assert "timeout --signal=TERM --kill-after=30s 300s" in gate + assert "FASTVIDEO_KERNEL_V10_RECEIPT.json" in gate + assert 'if source_commit != execution_commit:' in gate + assert 'observed_wheel_sha256 != receipt.get("wheel_sha256")' in gate + assert 'observed_prefix_tree_sha256 != receipt.get("installed_prefix_tree_sha256")' in gate + assert '"installed_prefix_tree_sha256": installed_prefix_tree_sha256' in rebuild + assert 'UV="${UV:-${KERNEL_ROOT}/tools/uv}"' in rebuild + assert "/home/vlm-wlsaidhi/.local/bin/uv" not in rebuild + assert '"__pycache__" not in relative.parts' in receipt_helper + assert 'path.suffix != ".pyc"' in receipt_helper + assert "907f2100e" in rebuild and "56d4a6074" in rebuild + assert "TORCH_CUDA_ARCH_LIST=10.0a" in rebuild + + +def test_h3_dmd2_v10_kernel_prefix_receipt_hashes_only_stable_installed_files(tmp_path: Path) -> None: + from scripts.train.h3_v10_kernel_receipt import RECEIPT_FILENAME, installed_prefix_tree_sha256 + + prefix = tmp_path / "prefix" + package = prefix / "fastvideo_kernel" + package.mkdir(parents=True) + installed = package / "kernel.so" + installed.write_bytes(b"installed-kernel-v1") + (prefix / "metadata.txt").write_text("metadata-v1", encoding="utf-8") + + receipt = prefix / RECEIPT_FILENAME + receipt.write_text("receipt-v1", encoding="utf-8") + bytecode_dir = package / "__pycache__" + bytecode_dir.mkdir() + bytecode = bytecode_dir / "module.cpython-312.pyc" + bytecode.write_bytes(b"bytecode-v1") + stray_bytecode = package / "generated.pyc" + stray_bytecode.write_bytes(b"stray-v1") + + original = installed_prefix_tree_sha256(prefix) + receipt.write_text("receipt-v2", encoding="utf-8") + bytecode.write_bytes(b"bytecode-v2") + stray_bytecode.write_bytes(b"stray-v2") + assert installed_prefix_tree_sha256(prefix) == original + + installed.write_bytes(b"installed-kernel-v2") + assert installed_prefix_tree_sha256(prefix) != original + + +def test_validation_dmd_sigmas_match_training_noise_amounts() -> None: + """``pipeline_config.dmd_denoising_steps`` replays the trained jump points. + + The H3 denoising stage normalizes the method's integer steps to base time + and lets each scheduler apply its own shift; the resulting clean-times must + match ``1 - shift_noise_amount(base)`` — the exact noising the packed DMD + adapter applies during training rollouts — with one forward per step. + """ + from fastvideo.models.schedulers.scheduling_minimax_h3 import MiniMaxH3Scheduler + + steps = [1000, 667, 333] + base = torch.tensor([step / 1000.0 for step in steps] + [0.0], dtype=torch.float32) + video = MiniMaxH3Scheduler(shift=12.0) + audio = MiniMaxH3Scheduler(shift=3.0) + video.set_timesteps(sigmas=video.shift_sigmas(base)) + audio.set_timesteps(sigmas=audio.shift_sigmas(base)) + + assert video.num_inference_steps == len(steps) + assert audio.num_inference_steps == len(steps) + for index, step in enumerate(steps): + base_step = torch.tensor([step / 1000.0]) + assert video.timesteps[index].item() == pytest.approx(1.0 - shift_noise_amount(base_step, 12.0).item()) + assert audio.timesteps[index].item() == pytest.approx(1.0 - shift_noise_amount(base_step, 3.0).item()) + + +def test_validation_callback_injects_method_denoising_steps() -> None: + """The callback copies the trained step list onto the validation config.""" + from fastvideo.train.callbacks.validation import ValidationCallback + + callback = ValidationCallback.__new__(ValidationCallback) + callback.method = SimpleNamespace(method_config={"dmd_denoising_steps": [1000, 667, 333]}) + + config = SimpleNamespace(dmd_denoising_steps=None) + callback._inject_method_denoising_steps(config) + assert config.dmd_denoising_steps == [1000, 667, 333] + + explicit = SimpleNamespace(dmd_denoising_steps=[1000, 500]) + callback._inject_method_denoising_steps(explicit) + assert explicit.dmd_denoising_steps == [1000, 500] + + callback.method = SimpleNamespace(method_config={"dmd_denoising_steps": [1000, 757], "warp_denoising_step": True}) + warped = SimpleNamespace(dmd_denoising_steps=None) + callback._inject_method_denoising_steps(warped) + assert warped.dmd_denoising_steps is None diff --git a/fastvideo/tests/train/methods/test_minimax_h3_finetune.py b/fastvideo/tests/train/methods/test_minimax_h3_finetune.py index 0d1663e1a2..2ac6612546 100644 --- a/fastvideo/tests/train/methods/test_minimax_h3_finetune.py +++ b/fastvideo/tests/train/methods/test_minimax_h3_finetune.py @@ -47,7 +47,8 @@ class _IdentityJointTransformer: patch_size = (1, 2, 2) def __call__(self, **kwargs): - return kwargs["hidden_states"], kwargs["audio_hidden_states"] + self.autocast_enabled = torch.is_autocast_enabled("cpu") + return kwargs["hidden_states"].float(), kwargs["audio_hidden_states"].float() @pytest.mark.parametrize( @@ -144,6 +145,50 @@ def test_h3_uniform_parameter_dtype_uses_fsdp_dtype() -> None: assert parameter_dtype == torch.bfloat16 +def test_h3_default_parameter_dtype_keeps_compute_boundaries_fp32() -> None: + model = cast( + MiniMaxH3Transformer3DModel, + SimpleNamespace( + config=MiniMaxH3Config(uniform_parameter_dtype=False), + _keep_in_fp32_modules=MiniMaxH3Transformer3DModel._keep_in_fp32_modules, + ), + ) + + assert MiniMaxH3Transformer3DModel._get_parameter_dtype( + model, + "proj_in.weight", + torch.bfloat16, + ) == torch.float32 + assert MiniMaxH3Transformer3DModel._get_parameter_dtype( + model, + "transformer_blocks.0.attn.to_q.weight", + torch.bfloat16, + ) == torch.bfloat16 + + +def test_h3_folded_adaln_keeps_fp32_training_master() -> None: + """Rank-reduced AdaLN must not silently demote an FP32 training load.""" + model = cast( + MiniMaxH3Transformer3DModel, + SimpleNamespace( + config=MiniMaxH3Config(uniform_parameter_dtype=False), + adaln_rank=768, + _keep_in_fp32_modules=MiniMaxH3Transformer3DModel._keep_in_fp32_modules, + ), + ) + + assert MiniMaxH3Transformer3DModel._get_parameter_dtype( + model, + "transformer_blocks.0.adaln_proj.linear.weight", + torch.float32, + ) == torch.float32 + assert MiniMaxH3Transformer3DModel._get_parameter_dtype( + model, + "transformer_blocks.0.adaln_proj.linear.weight", + torch.bfloat16, + ) == torch.bfloat16 + + def test_h3_materializes_rotary_frequencies_on_loader_device() -> None: """Verify that checkpoint loading moves analytic rotary state to the model device.""" model = cast( @@ -209,8 +254,11 @@ def test_h3_plugin_prepares_and_restores_joint_latent_shapes(monkeypatch: pytest ) assert isinstance(prediction, tuple) + assert model.transformer.autocast_enabled is False assert prediction[0].shape == batch.latents.shape assert prediction[1].shape == batch.audio_latents.shape + assert prediction[0].dtype == batch.noisy_model_input.dtype + assert prediction[1].dtype == batch.audio_noisy_model_input.dtype torch.testing.assert_close( prediction[0], -batch.noisy_model_input.permute(0, 2, 1, 3, 4), diff --git a/fastvideo/tests/train/trainer/test_validation.py b/fastvideo/tests/train/trainer/test_validation.py index 5a08a36f8c..a3b1927f43 100644 --- a/fastvideo/tests/train/trainer/test_validation.py +++ b/fastvideo/tests/train/trainer/test_validation.py @@ -84,6 +84,27 @@ def optimizers_zero_grad(self, iteration: int) -> None: self.weight.grad = None +class _RecordingCheckpointManager: + + def __init__(self) -> None: + self.inference_events: list[tuple[int, bool]] = [] + self.training_events: list[int] = [] + self.final_events: list[int] = [] + + def maybe_resume(self, *, resume_from_checkpoint: str) -> None: + assert resume_from_checkpoint == "" + return None + + def maybe_save_inference(self, step: int, *, validation_scheduled: bool) -> None: + self.inference_events.append((step, validation_scheduled)) + + def maybe_save(self, step: int) -> None: + self.training_events.append(step) + + def save_final(self, step: int) -> None: + self.final_events.append(step) + + def test_trainer_runs_validation_callback_during_training(monkeypatch, ) -> None: tracker = _DummyTracker() group = SimpleNamespace(rank=0, local_rank=0, rank_in_group=0, world_size=1) @@ -119,6 +140,7 @@ def test_trainer_runs_validation_callback_during_training(monkeypatch, ) -> None callback_configs=callback_configs, ) method = _DummyMethod() + checkpoint_manager = _RecordingCheckpointManager() trainer.run( method, @@ -126,6 +148,7 @@ def test_trainer_runs_validation_callback_during_training(monkeypatch, ) -> None "sample": "x" }], max_steps=3, + checkpoint_manager=checkpoint_manager, ) validation = trainer.callbacks._callbacks["validation"] @@ -137,3 +160,11 @@ def test_trainer_runs_validation_callback_during_training(monkeypatch, ) -> None assert method.optimizer_steps == [1, 2, 3] assert [step for _, step in tracker.logs] == [1, 2, 3] assert tracker.finished is True + assert checkpoint_manager.inference_events == [ + (0, True), + (1, False), + (2, True), + (3, False), + ] + assert checkpoint_manager.training_events == [1, 2, 3] + assert checkpoint_manager.final_events == [3] diff --git a/fastvideo/tests/train/utils/test_checkpoint.py b/fastvideo/tests/train/utils/test_checkpoint.py index 1bbab529cb..ec9acee89b 100644 --- a/fastvideo/tests/train/utils/test_checkpoint.py +++ b/fastvideo/tests/train/utils/test_checkpoint.py @@ -4,21 +4,28 @@ Covers the pure-Python portions of the checkpoint manager: name parsing, resume-path resolution, metadata round-trip, rolling-delete cleanup, the ``_is_stateful`` predicate, and the ``maybe_save`` gating -logic. Code paths that touch DCP (``dcp.save`` / ``dcp.load``) and -CUDA RNG snapshots are intentionally not covered here — those need a -GPU runner and will be tested in later phases. +logic. The inference staging path is covered with mocked DCP I/O; real +distributed collectives and CUDA RNG snapshots require a GPU runner. """ from __future__ import annotations +import json from pathlib import Path +from types import SimpleNamespace from typing import Any import pytest +import torch +import fastvideo.train.utils.checkpoint as checkpoint_module +import fastvideo.train.utils.inference_checkpoint as inference_checkpoint +from fastvideo.train.methods.base import TrainingMethod from fastvideo.train.utils.checkpoint import ( CheckpointConfig, CheckpointManager, + _FullModelState, _find_latest_checkpoint, + _is_complete_training_checkpoint, _is_stateful, _parse_step_from_dir, _resolve_resume_checkpoint, @@ -34,29 +41,75 @@ def _make_checkpoint_dir( step: int, *, with_dcp: bool = True, + with_metadata: bool = True, ) -> Path: - """Create a fake ``checkpoint-/dcp`` directory tree.""" + """Create a fake ``checkpoint-/dcp`` directory tree. + + ``dcp/.metadata`` is dcp.save's completion marker (written last); + ``_find_latest_checkpoint`` requires it, so a complete fake checkpoint + must include it. ``with_metadata=False`` fakes a crashed mid-write save. + """ ckpt_dir = output_dir / f"checkpoint-{step}" ckpt_dir.mkdir(parents=True, exist_ok=True) if with_dcp: (ckpt_dir / "dcp").mkdir(exist_ok=True) + if with_metadata: + (ckpt_dir / "dcp" / ".metadata").touch() return ckpt_dir +def _publish_fake_training_checkpoint( + checkpoint_dir: Path, + *, + step: int, + world_size: int = 1, +) -> None: + (checkpoint_dir / "metadata.json").write_text( + json.dumps({ + "step": step, + "config": { + "training": { + "distributed": { + "num_gpus": world_size, + }, + }, + }, + }), + encoding="utf-8", + ) + for rank in range(world_size): + (checkpoint_dir / f"rng_state_rank{rank}.pt").write_bytes(b"rng") + (checkpoint_dir / ".complete").write_text("complete\n", encoding="utf-8") + + def _make_manager( tmp_path: Path, *, save_steps: int = 0, + save_inference_on_validation: bool = False, keep_last: int = 0, + start_step: int = 0, raw_config: dict[str, Any] | None = None, + require_complete_training_checkpoint: bool = False, ) -> CheckpointManager: """Build a minimal ``CheckpointManager`` for tests that don't touch DCP.""" return CheckpointManager( method=None, dataloader=None, output_dir=str(tmp_path), - config=CheckpointConfig(save_steps=save_steps, keep_last=keep_last), - raw_config=raw_config, + config=CheckpointConfig( + save_steps=save_steps, + keep_last=keep_last, + start_step=start_step, + save_inference_on_validation=save_inference_on_validation, + require_complete_training_checkpoint=require_complete_training_checkpoint, + ), + raw_config=( + raw_config + if raw_config is not None + else ({"training": {"distributed": {"num_gpus": 1}}} + if require_complete_training_checkpoint else None) + ), ) @@ -98,6 +151,36 @@ def test_is_stateful_false_when_missing_load_state_dict() -> None: assert _is_stateful(_MissingLoad()) is False +def test_full_model_state_keeps_frozen_inference_parameters() -> None: + module = torch.nn.Linear(3, 2) + module.requires_grad_(False) + + state = _FullModelState(module).state_dict() + + assert set(state) == {"weight", "bias"} + + +def test_inference_checkpoint_role_is_explicit() -> None: + student = SimpleNamespace( + transformer=torch.nn.Linear(2, 2), + _init_from="base/student", + ) + fake_method = SimpleNamespace(_role_models={"student": student}) + + modules = TrainingMethod.inference_checkpoint_modules(fake_method, "student") + base_path = TrainingMethod.inference_checkpoint_base_model_path(fake_method, "student") + + assert modules == {"transformer": student.transformer} + assert base_path == "base/student" + + +def test_unknown_inference_checkpoint_role_raises() -> None: + fake_method = SimpleNamespace(_role_models={}) + + with pytest.raises(ValueError, match="known roles"): + TrainingMethod.inference_checkpoint_modules(fake_method, "ema") + + # --------------------------------------------------------------------------- # B. _parse_step_from_dir # --------------------------------------------------------------------------- @@ -138,6 +221,45 @@ def test_find_latest_returns_largest_step(tmp_path: Path) -> None: assert latest.name == "checkpoint-200" +def test_find_latest_default_preserves_legacy_dcp_completion_contract(tmp_path: Path) -> None: + legacy = _make_checkpoint_dir(tmp_path, 10) + + assert not (legacy / ".complete").exists() + assert _find_latest_checkpoint(tmp_path) == legacy + + +def test_find_latest_strict_skips_unpublished_newer_checkpoint(tmp_path: Path) -> None: + older = _make_checkpoint_dir(tmp_path, 5) + _publish_fake_training_checkpoint(older, step=5, world_size=2) + newer = _make_checkpoint_dir(tmp_path, 10) + _publish_fake_training_checkpoint(newer, step=10, world_size=2) + (newer / ".complete").unlink() + + latest = _find_latest_checkpoint(tmp_path, require_complete_marker=True) + + assert latest == older + + +@pytest.mark.parametrize("defect", ["marker", "metadata_step", "missing_rng", "extra_rng"]) +def test_strict_training_checkpoint_requires_complete_publication( + tmp_path: Path, + defect: str, +) -> None: + checkpoint = _make_checkpoint_dir(tmp_path, 10) + _publish_fake_training_checkpoint(checkpoint, step=10, world_size=2) + if defect == "marker": + (checkpoint / ".complete").write_text("incomplete\n", encoding="utf-8") + elif defect == "metadata_step": + _publish_fake_training_checkpoint(checkpoint, step=9, world_size=2) + elif defect == "missing_rng": + (checkpoint / "rng_state_rank1.pt").unlink() + else: + (checkpoint / "rng_state_rank2.pt").write_bytes(b"stale") + + assert not _is_complete_training_checkpoint(checkpoint, require_complete_marker=True) + assert _find_latest_checkpoint(tmp_path, require_complete_marker=True) is None + + def test_find_latest_skips_dirs_without_dcp_subdir(tmp_path: Path) -> None: # checkpoint-10 is "corrupted" — has no dcp/ subdir, must be skipped. _make_checkpoint_dir(tmp_path, 10, with_dcp=False) @@ -147,6 +269,16 @@ def test_find_latest_skips_dirs_without_dcp_subdir(tmp_path: Path) -> None: assert latest.name == "checkpoint-5" +def test_find_latest_skips_incomplete_dcp_save(tmp_path: Path) -> None: + # checkpoint-10 crashed mid-save — dcp/ exists but .metadata (written + # last by dcp.save) does not; resuming from it would fail at boot. + _make_checkpoint_dir(tmp_path, 10, with_metadata=False) + _make_checkpoint_dir(tmp_path, 5) + latest = _find_latest_checkpoint(tmp_path) + assert latest is not None + assert latest.name == "checkpoint-5" + + def test_find_latest_skips_non_checkpoint_dirs(tmp_path: Path) -> None: (tmp_path / "logs").mkdir() (tmp_path / "wandb").mkdir() @@ -176,6 +308,28 @@ def test_resolve_latest_returns_latest_checkpoint(tmp_path: Path) -> None: assert resolved.name == "checkpoint-30" +def test_resolve_explicit_strict_checkpoint_rejects_missing_marker(tmp_path: Path) -> None: + checkpoint = _make_checkpoint_dir(tmp_path, 42) + + with pytest.raises(ValueError, match="incomplete"): + _resolve_resume_checkpoint( + str(checkpoint), + output_dir=str(tmp_path), + require_complete_marker=True, + ) + + +def test_resolve_latest_strict_refuses_fresh_start_over_incomplete_state(tmp_path: Path) -> None: + _make_checkpoint_dir(tmp_path, 42) + + with pytest.raises(ValueError, match="refusing to start from scratch"): + _resolve_resume_checkpoint( + "latest", + output_dir=str(tmp_path), + require_complete_marker=True, + ) + + def test_resolve_explicit_checkpoint_dir(tmp_path: Path) -> None: ckpt = _make_checkpoint_dir(tmp_path, 42) resolved = _resolve_resume_checkpoint(str(ckpt), output_dir=str(tmp_path)) @@ -254,6 +408,79 @@ def test_write_metadata_includes_raw_config(tmp_path: Path) -> None: assert loaded["config"] == raw +def test_strict_manager_requires_world_size_in_saved_config(tmp_path: Path) -> None: + with pytest.raises(ValueError, match="num_gpus"): + CheckpointManager( + method=None, + dataloader=None, + output_dir=str(tmp_path), + config=CheckpointConfig( + save_steps=1, + keep_last=1, + require_complete_training_checkpoint=True, + ), + raw_config={}, + ) + + +def test_save_publishes_training_complete_after_rng_barrier( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + raw_config = {"training": {"distributed": {"num_gpus": 1}}} + method = SimpleNamespace(checkpoint_state=lambda: {}) + manager = CheckpointManager( + method=method, + dataloader=None, + output_dir=str(tmp_path), + config=CheckpointConfig( + save_steps=1, + keep_last=0, + require_complete_training_checkpoint=True, + ), + raw_config=raw_config, + ) + checkpoint = tmp_path / "checkpoint-10" + checkpoint.mkdir() + marker = checkpoint / ".complete" + marker.write_text("complete\n", encoding="utf-8") + events: list[str] = [] + + def fake_dcp_save(states: dict[str, Any], *, checkpoint_id: str) -> None: + assert states == {} + assert not marker.exists() + events.append("dcp") + dcp_dir = Path(checkpoint_id) + dcp_dir.mkdir(parents=True, exist_ok=True) + (dcp_dir / ".metadata").touch() + + def fake_rng_save(checkpoint_dir: Path) -> None: + assert not marker.exists() + events.append("rng") + (checkpoint_dir / "rng_state_rank0.pt").write_bytes(b"rng") + + def fake_barrier() -> None: + events.append("barrier") + + publish = checkpoint_module._publish_training_checkpoint_complete + + def record_publish(checkpoint_dir: Path) -> None: + assert events[-1] == "barrier" + events.append("publish") + publish(checkpoint_dir) + + monkeypatch.setattr(checkpoint_module.dcp, "save", fake_dcp_save) + monkeypatch.setattr(checkpoint_module, "_barrier", fake_barrier) + monkeypatch.setattr(checkpoint_module, "_publish_training_checkpoint_complete", record_publish) + monkeypatch.setattr(manager, "_save_rng_snapshot", fake_rng_save) + + manager.save(10) + + assert events == ["barrier", "dcp", "barrier", "rng", "barrier", "publish", "barrier"] + assert marker.read_text(encoding="utf-8") == "complete\n" + assert _is_complete_training_checkpoint(checkpoint, require_complete_marker=True) + + def test_load_metadata_raises_on_missing_file(tmp_path: Path) -> None: ckpt_dir = _make_checkpoint_dir(tmp_path, 7) # No metadata.json written. @@ -304,6 +531,45 @@ def test_cleanup_skips_non_checkpoint_dirs(tmp_path: Path) -> None: assert remaining == ["checkpoint-3", "logs", "wandb"] +def test_cleanup_never_removes_inference_checkpoints(tmp_path: Path) -> None: + mgr = _make_manager(tmp_path, keep_last=1) + for step in (1, 2, 3): + _make_checkpoint_dir(tmp_path, step) + inference_dir = tmp_path / "inference" / f"checkpoint-{step}" + inference_dir.mkdir(parents=True) + (inference_dir / ".complete").touch() + + mgr._cleanup_old_checkpoints() + + assert sorted(path.name for path in tmp_path.glob("checkpoint-*")) == ["checkpoint-3"] + assert sorted(path.name for path in (tmp_path / "inference").iterdir()) == [ + "checkpoint-1", + "checkpoint-2", + "checkpoint-3", + ] + + +def test_strict_cleanup_does_not_count_incomplete_newer_directories(tmp_path: Path) -> None: + mgr = _make_manager( + tmp_path, + keep_last=2, + require_complete_training_checkpoint=True, + ) + for step in (100, 200, 300): + checkpoint = _make_checkpoint_dir(tmp_path, step) + _publish_fake_training_checkpoint(checkpoint, step=step) + # A failed later save has DCP metadata but no post-RNG publication marker. + _make_checkpoint_dir(tmp_path, 400) + + mgr._cleanup_old_checkpoints() + + assert sorted(path.name for path in tmp_path.glob("checkpoint-*")) == [ + "checkpoint-200", + "checkpoint-300", + "checkpoint-400", + ] + + # --------------------------------------------------------------------------- # G. maybe_save gating logic # --------------------------------------------------------------------------- @@ -353,3 +619,146 @@ def test_maybe_save_triggers_on_each_interval(tmp_path: Path) -> None: for step in range(1, 41): mgr.maybe_save(step=step) assert calls == [10, 20, 30, 40] + + +def _record_both_save_calls( + mgr: CheckpointManager, +) -> tuple[list[int], list[int]]: + training_calls: list[int] = [] + inference_calls: list[int] = [] + + def fake_training_save(step: int) -> None: + training_calls.append(step) + mgr._last_saved_step = step + + def fake_inference_save(step: int) -> None: + inference_calls.append(step) + mgr._last_inference_saved_step = step + + mgr.save = fake_training_save # type: ignore[method-assign] + mgr.save_inference = fake_inference_save # type: ignore[method-assign] + return training_calls, inference_calls + + +def test_inference_save_tracks_validation_events_and_ignores_start_gate(tmp_path: Path) -> None: + mgr = _make_manager( + tmp_path, + save_steps=10, + save_inference_on_validation=True, + start_step=12, + ) + training_calls, inference_calls = _record_both_save_calls(mgr) + + mgr.maybe_save_inference(0, validation_scheduled=True) + mgr.maybe_save_inference(4, validation_scheduled=False) + mgr.maybe_save_inference(10, validation_scheduled=True) + mgr.maybe_save_inference(10, validation_scheduled=True) + for step in range(1, 21): + mgr.maybe_save(step) + + assert training_calls == [20] + assert inference_calls == [0, 10] + + +def test_disabled_validation_inference_checkpointing_is_no_op(tmp_path: Path) -> None: + mgr = _make_manager( + tmp_path, + save_inference_on_validation=False, + ) + _, inference_calls = _record_both_save_calls(mgr) + mgr.maybe_save_inference(0, validation_scheduled=True) + + assert inference_calls == [] + + +def test_save_final_dedupes_training_checkpoint(tmp_path: Path) -> None: + mgr = _make_manager( + tmp_path, + save_steps=10, + save_inference_on_validation=True, + ) + training_calls, inference_calls = _record_both_save_calls(mgr) + + mgr.maybe_save_inference(20, validation_scheduled=True) + mgr.maybe_save(20) + mgr.save_final(20) + + assert training_calls == [20] + assert inference_calls == [20] + + +def test_save_final_does_not_create_off_validation_inference_product(tmp_path: Path) -> None: + mgr = _make_manager( + tmp_path, + save_steps=10, + save_inference_on_validation=True, + ) + training_calls, inference_calls = _record_both_save_calls(mgr) + + mgr.save_final(17) + + assert training_calls == [17] + assert inference_calls == [] + + +def test_save_inference_stages_full_state_and_publishes_once( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + module = torch.nn.Linear(2, 2) + module.requires_grad_(False) + base_model = tmp_path / "base" + base_model.mkdir() + + class _Method: + cuda_generator = None + + def inference_checkpoint_modules(self, role: str) -> dict[str, torch.nn.Module]: + assert role == "student" + return {"transformer": module} + + def inference_checkpoint_base_model_path(self, role: str) -> str: + assert role == "student" + return str(base_model) + + manager = CheckpointManager( + method=_Method(), + dataloader=None, + output_dir=str(tmp_path / "run"), + config=CheckpointConfig( + save_steps=0, + keep_last=0, + save_inference_on_validation=True, + ), + ) + staged_states: list[dict[str, Any]] = [] + + def fake_dcp_save(states: dict[str, Any], *, checkpoint_id: str) -> None: + staged_states.append(states) + dcp_dir = Path(checkpoint_id) + dcp_dir.mkdir(parents=True, exist_ok=True) + (dcp_dir / ".metadata").touch() + + def fake_export(**kwargs: Any) -> Path: + target = Path(kwargs["output_dir"]) + target.mkdir(parents=True) + (target / ".complete").touch() + return target + + monkeypatch.setattr("fastvideo.train.utils.checkpoint.dcp.save", fake_dcp_save) + monkeypatch.setattr(inference_checkpoint, "export_inference_checkpoint", fake_export) + monkeypatch.setattr( + inference_checkpoint, + "validate_complete_inference_checkpoint", + lambda path, *, step: path if (path / ".complete").is_file() else None, + ) + + manager.save_inference(10) + manager.save_inference(10) + + assert len(staged_states) == 1 + state = staged_states[0] + assert set(state) == {"roles.student.transformer"} + assert isinstance(state["roles.student.transformer"], _FullModelState) + assert not (tmp_path / "run" / ".inference-staging").exists() + assert (tmp_path / "run" / "inference" / "checkpoint-10" / ".complete").is_file() diff --git a/fastvideo/tests/train/utils/test_config.py b/fastvideo/tests/train/utils/test_config.py index 8471eea8c9..25988dcd19 100644 --- a/fastvideo/tests/train/utils/test_config.py +++ b/fastvideo/tests/train/utils/test_config.py @@ -75,6 +75,10 @@ def test_minimal_yaml_applies_all_defaults(tmp_path: Path) -> None: assert t.loop.gradient_accumulation_steps == 1 assert t.checkpoint.output_dir == "" + assert t.checkpoint.save_inference_checkpoint_on_validation is False + assert t.checkpoint.inference_checkpoint_role == "student" + assert t.checkpoint.inference_checkpoint_dtype == "bfloat16" + assert t.checkpoint.require_complete_training_checkpoint is False assert t.checkpoint.checkpoints_total_limit == 0 assert t.tracker.trackers == [] @@ -84,9 +88,12 @@ def test_minimal_yaml_applies_all_defaults(tmp_path: Path) -> None: assert t.model.weighting_scheme == "uniform" assert t.model.precondition_outputs is False assert t.model.moba_config == {} + assert t.model.enable_torch_compile is False + assert t.model.torch_compile_kwargs == {} assert t.dit_precision == "fp32" assert t.vsa_sparsity == 0.0 + assert t.vsa_tile_size == 256 assert t.pipeline_config is None @@ -126,7 +133,11 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None: }, "checkpoint": { "output_dir": "/out", + "save_inference_checkpoint_on_validation": True, + "inference_checkpoint_role": "student", + "inference_checkpoint_dtype": "float16", "training_state_checkpointing_steps": 50, + "require_complete_training_checkpoint": True, "checkpoints_total_limit": 3, }, "tracker": { @@ -135,13 +146,18 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None: "run_name": "myrun", }, "vsa": { - "sparsity": 0.5 + "sparsity": 0.5, + "tile_size": 64, }, "model": { "weighting_scheme": "logit_normal", "logit_mean": 0.5, "logit_std": 1.5, "precondition_outputs": True, + "enable_torch_compile": True, + "torch_compile_kwargs": { + "dynamic": False, + }, }, "dit_precision": "bf16", } @@ -167,14 +183,21 @@ def test_full_yaml_populates_all_training_fields(tmp_path: Path) -> None: assert t.loop.gradient_accumulation_steps == 4 assert t.checkpoint.output_dir == "/out" + assert t.checkpoint.save_inference_checkpoint_on_validation is True + assert t.checkpoint.inference_checkpoint_role == "student" + assert t.checkpoint.inference_checkpoint_dtype == "float16" + assert t.checkpoint.require_complete_training_checkpoint is True assert t.checkpoint.checkpoints_total_limit == 3 assert t.tracker.trackers == ["wandb"] assert t.tracker.project_name == "myproj" assert t.vsa_sparsity == pytest.approx(0.5) + assert t.vsa_tile_size == 64 assert t.model.weighting_scheme == "logit_normal" assert t.model.precondition_outputs is True + assert t.model.enable_torch_compile is True + assert t.model.torch_compile_kwargs == {"dynamic": False} assert t.dit_precision == "bf16" @@ -190,6 +213,46 @@ def test_missing_models_raises(tmp_path: Path) -> None: load_run_config(_write_yaml(tmp_path, data)) +def test_invalid_vsa_tile_size_raises(tmp_path: Path) -> None: + data = _minimal_yaml() + data["training"] = {"vsa": {"tile_size": 128}} + with pytest.raises(ValueError, match="training.vsa.tile_size must be 64 or 256"): + load_run_config(_write_yaml(tmp_path, data)) + + +def test_validation_inference_checkpoint_flag_requires_bool(tmp_path: Path) -> None: + data = _minimal_yaml() + data["training"] = {"checkpoint": {"save_inference_checkpoint_on_validation": 1}} + with pytest.raises(ValueError, match="save_inference_checkpoint_on_validation"): + load_run_config(_write_yaml(tmp_path, data)) + + +def test_training_checkpoint_completion_flag_requires_bool(tmp_path: Path) -> None: + data = _minimal_yaml() + data["training"] = {"checkpoint": {"require_complete_training_checkpoint": 1}} + with pytest.raises(ValueError, match="require_complete_training_checkpoint"): + load_run_config(_write_yaml(tmp_path, data)) + + +def test_enabled_inference_checkpoint_requires_role(tmp_path: Path) -> None: + data = _minimal_yaml() + data["training"] = { + "checkpoint": { + "save_inference_checkpoint_on_validation": True, + "inference_checkpoint_role": "", + } + } + with pytest.raises(ValueError, match="inference_checkpoint_role"): + load_run_config(_write_yaml(tmp_path, data)) + + +def test_invalid_inference_checkpoint_dtype_raises(tmp_path: Path) -> None: + data = _minimal_yaml() + data["training"] = {"checkpoint": {"inference_checkpoint_dtype": "fp8"}} + with pytest.raises(ValueError, match="inference_checkpoint_dtype"): + load_run_config(_write_yaml(tmp_path, data)) + + def test_missing_method_raises(tmp_path: Path) -> None: data = _minimal_yaml() del data["method"] diff --git a/fastvideo/tests/train/utils/test_inference_checkpoint.py b/fastvideo/tests/train/utils/test_inference_checkpoint.py new file mode 100644 index 0000000000..518d0d833e --- /dev/null +++ b/fastvideo/tests/train/utils/test_inference_checkpoint.py @@ -0,0 +1,478 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU contracts for bounded-memory modular inference checkpoint export.""" + +from __future__ import annotations + +import json +from pathlib import Path + +import pytest +import torch +import torch.distributed.checkpoint as dcp +from safetensors.torch import load_file, save_file + +import fastvideo.train.utils.inference_checkpoint as inference_checkpoint +from fastvideo.train.utils.inference_checkpoint import ( + InferenceCheckpointExportError, + UnsupportedMergedReverseMappingError, + export_inference_checkpoint, + export_inference_checkpoint_from_dcp, + validate_complete_inference_checkpoint, +) + + +class _LiveModule(torch.nn.Module): + + def __init__(self, reverse_mapping: dict | None = None) -> None: + super().__init__() + self.master = torch.nn.Parameter( + torch.tensor([1.25, -2.5], dtype=torch.float32) + ) + self.reverse_param_names_mapping = reverse_mapping or {} + + +@pytest.fixture +def base_model_dir(tmp_path: Path) -> Path: + base = tmp_path / "base" + transformer = base / "transformer" + transformer.mkdir(parents=True) + (transformer / "config.json").write_text( + json.dumps({"_class_name": "FakeTransformer"}), + encoding="utf-8", + ) + (base / "modular_model_index.json").write_text( + json.dumps({"_class_name": "FakePipeline"}), + encoding="utf-8", + ) + vae = base / "vae" + vae.mkdir() + (vae / "config.json").write_text("{}", encoding="utf-8") + (base / ".cache").mkdir() + return base + + +def _write_base_transformer_weights(base: Path, tensors: dict[str, torch.Tensor]) -> None: + """Write one indexed base shard to define the exact export contract.""" + transformer = base / "transformer" + filename = "diffusion_pytorch_model-00001-of-00001.safetensors" + save_file(tensors, transformer / filename) + index = { + "metadata": { + "total_size": sum(tensor.numel() * tensor.element_size() for tensor in tensors.values()) + }, + "weight_map": {key: filename for key in tensors}, + } + (transformer / "diffusion_pytorch_model.safetensors.index.json").write_text( + json.dumps(index, indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + + +def _save_model_only_dcp( + root: Path, + tensors: dict[str, torch.Tensor], + *, + role: str = "student", + module_name: str = "transformer", +) -> Path: + checkpoint = root / "temporary-model-checkpoint" + dcp_dir = checkpoint / "dcp" + state = { + f"roles.{role}.{module_name}.{key}": value.clone() + for key, value in tensors.items() + } + dcp.save(state, checkpoint_id=str(dcp_dir)) + assert (dcp_dir / ".metadata").is_file() + return checkpoint + + +def _load_exported_tensors(transformer_dir: Path) -> dict[str, torch.Tensor]: + tensors: dict[str, torch.Tensor] = {} + for shard in sorted(transformer_dir.glob("*.safetensors")): + tensors.update(load_file(shard)) + return tensors + + +def test_export_casts_maps_shards_and_publishes_complete_layout( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights( + base_model_dir, + { + "disk.proj.weight": torch.empty(2, 2), + "disk.position_ids": torch.empty(2, dtype=torch.int64), + }, + ) + source = _save_model_only_dcp( + tmp_path, + { + "proj.weight": torch.arange(4, dtype=torch.float32).reshape(2, 2), + "block.attn.to_gate_compress.weight": torch.tensor( + [0.25, 0.5, 0.75, 1.0], dtype=torch.float32 + ), + "position_ids": torch.tensor([3, 7], dtype=torch.int64), + }, + ) + module = _LiveModule( + { + "proj.weight": ("disk.proj.weight", None, None), + "position_ids": ("disk.position_ids", None, None), + } + ) + master_before = module.master.detach().clone() + + result = export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=100, + module=module, + base_model_dir=base_model_dir, + dtype="bfloat16", + max_shard_size_bytes=16, + ) + + assert result == (tmp_path / "run" / "inference" / "checkpoint-100").resolve() + assert (result / ".complete").read_text(encoding="utf-8") == "complete\n" + assert (result / "metadata.json").is_file() + assert (result / "modular_model_index.json").is_symlink() + assert (result / "vae").is_symlink() + assert not (result / ".cache").exists() + assert (result / "transformer" / "config.json").is_file() + assert not (result / "transformer" / "config.json").is_symlink() + assert not (result / "transformer" / "base-weights.safetensors").exists() + + exported = _load_exported_tensors(result / "transformer") + assert set(exported) == { + "disk.proj.weight", + "disk.position_ids", + "block.attn.to_gate_compress.weight", + } + assert exported["disk.proj.weight"].dtype == torch.bfloat16 + assert exported["block.attn.to_gate_compress.weight"].dtype == torch.bfloat16 + assert exported["disk.position_ids"].dtype == torch.int64 + torch.testing.assert_close( + exported["disk.proj.weight"], + torch.arange(4, dtype=torch.bfloat16).reshape(2, 2), + ) + assert torch.equal(exported["disk.position_ids"], torch.tensor([3, 7])) + assert module.master.dtype == torch.float32 + assert torch.equal(module.master.detach(), master_before) + + metadata = json.loads((result / "metadata.json").read_text(encoding="utf-8")) + assert metadata["kind"] == "inference" + assert metadata["step"] == 100 + assert metadata["role"] == "student" + assert metadata["module"] == "transformer" + assert metadata["dtype"] == "bfloat16" + assert metadata["tensor_count"] == 3 + assert metadata["total_size"] == 32 + assert metadata["shard_count"] == 3 + assert all(size <= 16 for size in metadata["shard_sizes"]) + + index = json.loads( + ( + result + / "transformer" + / "diffusion_pytorch_model.safetensors.index.json" + ).read_text(encoding="utf-8") + ) + assert index["metadata"]["total_size"] == 32 + assert set(index["weight_map"]) == set(exported) + assert len(set(index["weight_map"].values())) == 3 + + +def test_export_is_atomic_on_failure_and_retry_succeeds( + tmp_path: Path, + base_model_dir: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _write_base_transformer_weights( + base_model_dir, + { + "a": torch.empty(4), + "b": torch.empty(4), + }, + ) + source = _save_model_only_dcp( + tmp_path, + { + "a": torch.ones(4, dtype=torch.float32), + "b": torch.ones(4, dtype=torch.float32), + }, + ) + module = _LiveModule() + real_save_file = inference_checkpoint.save_file + calls = 0 + + def fail_after_first_write(tensors, filename, metadata=None): + nonlocal calls + calls += 1 + real_save_file(tensors, filename, metadata=metadata) + if calls == 1: + raise OSError("injected shard write failure") + + monkeypatch.setattr(inference_checkpoint, "save_file", fail_after_first_write) + kwargs = { + "dcp_checkpoint": source, + "output_dir": tmp_path / "run", + "step": 7, + "module": module, + "base_model_dir": base_model_dir, + "dtype": torch.bfloat16, + "max_shard_size_bytes": 8, + } + with pytest.raises(OSError, match="injected"): + export_inference_checkpoint_from_dcp(**kwargs) + + inference_root = tmp_path / "run" / "inference" + assert not (inference_root / "checkpoint-7").exists() + assert list(inference_root.glob(".checkpoint-7.tmp-*")) == [] + + monkeypatch.setattr(inference_checkpoint, "save_file", real_save_file) + result = export_inference_checkpoint_from_dcp(**kwargs) + assert (result / ".complete").is_file() + + # Completed outputs are immutable and retries are idempotent, even after + # the temporary DCP has been retired by the caller. + for child in (source / "dcp").iterdir(): + child.unlink() + (source / "dcp").rmdir() + source.rmdir() + assert export_inference_checkpoint_from_dcp(**kwargs) == result + + +def test_export_rejects_merged_reverse_mapping( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"q.weight": torch.empty(2, 2)}) + source = _save_model_only_dcp( + tmp_path, + {"fused_qkv.weight": torch.ones(6, 2)}, + ) + module = _LiveModule( + {"fused_qkv.weight": ("q.weight", 0, 3)} + ) + + with pytest.raises(UnsupportedMergedReverseMappingError, match="model-specific split"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=3, + module=module, + base_model_dir=base_model_dir, + max_shard_size_bytes=1024, + ) + assert not (tmp_path / "run" / "inference" / "checkpoint-3").exists() + + +def test_export_refuses_incomplete_source_or_destination( + tmp_path: Path, + base_model_dir: Path, +) -> None: + source = tmp_path / "incomplete-source" / "dcp" + source.mkdir(parents=True) + with pytest.raises(FileNotFoundError, match="missing .metadata"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=1, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + complete_source = _save_model_only_dcp( + tmp_path / "other", + {"weight": torch.ones(1)}, + ) + incomplete_destination = tmp_path / "run" / "inference" / "checkpoint-2" + incomplete_destination.mkdir(parents=True) + (incomplete_destination / "metadata.json").write_text("{}", encoding="utf-8") + with pytest.raises(InferenceCheckpointExportError, match="Refusing to overwrite incomplete"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=complete_source, + output_dir=tmp_path / "run", + step=2, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + +def test_export_rejects_unknown_unmapped_tensor( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"weight": torch.empty(2)}) + source = _save_model_only_dcp( + tmp_path, + { + "weight": torch.ones(2), + "surprise.weight": torch.ones(2), + }, + ) + + with pytest.raises(InferenceCheckpointExportError, match="unknown base key"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=4, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + +def test_export_rejects_missing_base_tensor( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights( + base_model_dir, + { + "weight": torch.empty(2), + "bias": torch.empty(2), + }, + ) + source = _save_model_only_dcp(tmp_path, {"weight": torch.ones(2)}) + + with pytest.raises(InferenceCheckpointExportError, match="missing 1 base transformer tensors"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=5, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + +def test_export_rejects_base_shape_mismatch( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"weight": torch.empty(3)}) + source = _save_model_only_dcp(tmp_path, {"weight": torch.ones(2)}) + + with pytest.raises(InferenceCheckpointExportError, match="shape .* != base transformer shape"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=6, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + +def test_export_rejects_base_index_header_mismatch( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"weight": torch.empty(2)}) + index_path = base_model_dir / "transformer" / "diffusion_pytorch_model.safetensors.index.json" + index = json.loads(index_path.read_text(encoding="utf-8")) + index["weight_map"]["ghost"] = next(iter(index["weight_map"].values())) + index_path.write_text(json.dumps(index), encoding="utf-8") + source = _save_model_only_dcp(tmp_path, {"weight": torch.ones(2)}) + + with pytest.raises(InferenceCheckpointExportError, match="index/header mismatch"): + export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "run", + step=7, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + + +def test_completed_checkpoint_validator_rejects_missing_and_corrupt_shards( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"weight": torch.empty(2)}) + source = _save_model_only_dcp(tmp_path, {"weight": torch.ones(2)}) + + missing = export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "missing-run", + step=8, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + missing_shard = next((missing / "transformer").glob("*.safetensors")) + missing_shard.unlink() + with pytest.raises(InferenceCheckpointExportError, match="shards differ from its index"): + validate_complete_inference_checkpoint(missing, step=8) + + corrupt = export_inference_checkpoint_from_dcp( + dcp_checkpoint=source, + output_dir=tmp_path / "corrupt-run", + step=9, + module=_LiveModule(), + base_model_dir=base_model_dir, + ) + corrupt_shard = next((corrupt / "transformer").glob("*.safetensors")) + corrupt_shard.write_bytes(b"corrupt") + with pytest.raises(InferenceCheckpointExportError, match="Cannot read inference checkpoint shard"): + validate_complete_inference_checkpoint(corrupt, step=9) + + +def test_checkpoint_manager_adapter_uses_exact_target( + tmp_path: Path, + base_model_dir: Path, +) -> None: + _write_base_transformer_weights(base_model_dir, {"weight": torch.empty(2)}) + source = _save_model_only_dcp(tmp_path, {"weight": torch.ones(2)}) + module = _LiveModule() + target = tmp_path / "run" / "inference" / "checkpoint-9" + + result = export_inference_checkpoint( + dcp_dir=source / "dcp", + output_dir=target, + base_model_path=base_model_dir, + role="student", + modules={"transformer": module}, + dtype="bfloat16", + step=9, + raw_config={"not": "serialized"}, + ) + + assert result == target.resolve() + metadata = json.loads((result / "metadata.json").read_text(encoding="utf-8")) + assert metadata["config"] == {"not": "serialized"} + assert "source_dcp" not in metadata + + for child in (source / "dcp").iterdir(): + child.unlink() + (source / "dcp").rmdir() + source.rmdir() + assert export_inference_checkpoint( + dcp_dir=source / "dcp", + output_dir=target, + base_model_path=base_model_dir, + role="student", + modules={"transformer": module}, + dtype="bfloat16", + step=9, + raw_config={"not": "serialized"}, + ) == result + + with pytest.raises(ValueError, match="CheckpointManager output_dir"): + export_inference_checkpoint( + dcp_dir=source / "dcp", + output_dir=tmp_path / "wrong-target", + base_model_path=base_model_dir, + role="student", + modules={"transformer": module}, + dtype="bfloat16", + step=9, + ) + + with pytest.raises(InferenceCheckpointExportError, match="exactly one module"): + export_inference_checkpoint( + dcp_dir=source / "dcp", + output_dir=tmp_path / "run" / "inference" / "checkpoint-10", + base_model_path=base_model_dir, + role="student", + modules={}, + dtype="bfloat16", + step=10, + ) diff --git a/fastvideo/tests/train/utils/test_inference_checkpoint_distributed.py b/fastvideo/tests/train/utils/test_inference_checkpoint_distributed.py new file mode 100644 index 0000000000..088e2ae32c --- /dev/null +++ b/fastvideo/tests/train/utils/test_inference_checkpoint_distributed.py @@ -0,0 +1,286 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Two-GPU FSDP2 gate for validation-triggered inference checkpoints. + +This test intentionally launches a real two-rank NCCL worker and must run on a +compute node. It covers frozen parameters, DCP staging, bounded rank-zero +export, strict reload, RNG preservation, and rank-zero export failure +propagation without leaving a long-running NCCL collective outstanding. +""" + +from __future__ import annotations + +import argparse +import json +import os +import random +import subprocess +import sys +from pathlib import Path +from typing import Any + +import numpy as np +import pytest +import torch +import torch.distributed as dist +from safetensors.torch import load_file, save_file +from torch.distributed import init_device_mesh +from torch.distributed.fsdp import fully_shard + +import fastvideo.train.utils.inference_checkpoint as inference_checkpoint +from fastvideo.train.utils.checkpoint import CheckpointConfig, CheckpointManager + +WORLD_SIZE = 2 + + +class _TinyTransformer(torch.nn.Module): + + def __init__(self, device: torch.device | str = "cpu") -> None: + super().__init__() + self.frozen = torch.nn.Parameter( + torch.tensor([1.25, -2.5, 3.75, -4.0], device=device), + requires_grad=False, + ) + self.trainable = torch.nn.Parameter( + torch.tensor([[0.5, 1.0], [1.5, 2.0]], device=device), + ) + self.reverse_param_names_mapping: dict[str, tuple[str, None, None]] = {} + + +class _Method: + + def __init__(self, transformer: torch.nn.Module, base_model_dir: Path, generator: torch.Generator) -> None: + self.transformer = transformer + self.base_model_dir = base_model_dir + self.cuda_generator = generator + + def inference_checkpoint_modules(self, role: str) -> dict[str, torch.nn.Module]: + if role != "student": + raise ValueError(role) + return {"transformer": self.transformer} + + def inference_checkpoint_base_model_path(self, role: str) -> str: + if role != "student": + raise ValueError(role) + return str(self.base_model_dir) + + +def _write_base_model(base_model_dir: Path) -> None: + transformer_dir = base_model_dir / "transformer" + transformer_dir.mkdir(parents=True) + (transformer_dir / "config.json").write_text( + json.dumps({"_class_name": "TinyTransformer"}), + encoding="utf-8", + ) + filename = "diffusion_pytorch_model-00001-of-00001.safetensors" + tensors = { + "frozen": torch.empty(4), + "trainable": torch.empty(2, 2), + } + save_file(tensors, transformer_dir / filename) + index = { + "metadata": { + "total_size": sum(tensor.numel() * tensor.element_size() for tensor in tensors.values()) + }, + "weight_map": {key: filename for key in tensors}, + } + (transformer_dir / "diffusion_pytorch_model.safetensors.index.json").write_text( + json.dumps(index, indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + (base_model_dir / "modular_model_index.json").write_text( + json.dumps({"_class_name": "TinyPipeline"}), + encoding="utf-8", + ) + + +def _load_exported_state(checkpoint_dir: Path) -> dict[str, torch.Tensor]: + state: dict[str, torch.Tensor] = {} + for shard in sorted((checkpoint_dir / "transformer").glob("*.safetensors")): + state.update(load_file(shard)) + return state + + +def _capture_rng(generator: torch.Generator) -> dict[str, Any]: + numpy_state = np.random.get_state() + return { + "torch": torch.get_rng_state().clone(), + "python": random.getstate(), + "numpy": ( + numpy_state[0], + numpy_state[1].copy(), + numpy_state[2], + numpy_state[3], + numpy_state[4], + ), + "cuda": torch.cuda.get_rng_state().clone(), + "generator": generator.get_state().clone(), + } + + +def _rng_equal(before: dict[str, Any], after: dict[str, Any]) -> bool: + before_numpy = before["numpy"] + after_numpy = after["numpy"] + return bool( + torch.equal(before["torch"], after["torch"]) + and before["python"] == after["python"] + and before_numpy[0] == after_numpy[0] + and np.array_equal(before_numpy[1], after_numpy[1]) + and before_numpy[2:] == after_numpy[2:] + and torch.equal(before["cuda"], after["cuda"]) + and torch.equal(before["generator"], after["generator"]) + ) + + +def _run_worker(work_dir: Path, result_path: Path) -> None: + dist.init_process_group("nccl") + rank = dist.get_rank() + local_rank = int(os.environ["LOCAL_RANK"]) + device = torch.device("cuda", local_rank) + torch.cuda.set_device(device) + try: + if dist.get_world_size() != WORLD_SIZE: + raise RuntimeError(f"Expected world size {WORLD_SIZE}, got {dist.get_world_size()}") + + base_model_dir = work_dir / "base" + output_dir = work_dir / "run" + if rank == 0: + _write_base_model(base_model_dir) + dist.barrier() + + torch.manual_seed(1000 + rank) + torch.cuda.manual_seed(2000 + rank) + random.seed(3000 + rank) + np.random.seed(4000 + rank) + generator = torch.Generator(device=device) + generator.manual_seed(5000 + rank) + + transformer = _TinyTransformer(device=device) + mesh = init_device_mesh("cuda", (WORLD_SIZE, )) + fully_shard(transformer, mesh=mesh) + method = _Method(transformer, base_model_dir, generator) + manager = CheckpointManager( + method=method, + dataloader=None, + output_dir=str(output_dir), + config=CheckpointConfig( + save_steps=0, + keep_last=0, + save_inference_on_validation=True, + inference_role="student", + inference_dtype="float32", + ), + ) + + success_rng_before = _capture_rng(generator) + manager.save_inference(1) + success_rng_after = _capture_rng(generator) + + checkpoint_dir = output_dir / "inference" / "checkpoint-1" + exported = _load_exported_state(checkpoint_dir) + reloaded = _TinyTransformer() + incompatible = reloaded.load_state_dict(exported, strict=True) + expected = _TinyTransformer().state_dict() + reload_equal = (not incompatible.missing_keys and not incompatible.unexpected_keys + and all(torch.equal(reloaded.state_dict()[key], value) for key, value in expected.items())) + frozen_present = "frozen" in exported and torch.equal(exported["frozen"], expected["frozen"]) + + real_export = inference_checkpoint.export_inference_checkpoint + if rank == 0: + + def _injected_failure(**_: Any) -> Path: + raise RuntimeError("injected rank-zero export failure") + + inference_checkpoint.export_inference_checkpoint = _injected_failure + + failure_rng_before = _capture_rng(generator) + failure: str | None = None + try: + manager.save_inference(2) + except RuntimeError as error: + failure = str(error) + finally: + if rank == 0: + inference_checkpoint.export_inference_checkpoint = real_export + failure_rng_after = _capture_rng(generator) + + local_result = { + "rank": rank, + "success_rng_equal": _rng_equal(success_rng_before, success_rng_after), + "failure_rng_equal": _rng_equal(failure_rng_before, failure_rng_after), + "reload_equal": reload_equal, + "frozen_present": frozen_present, + "failure": failure, + } + gathered: list[dict[str, Any] | None] = [None] * WORLD_SIZE + dist.all_gather_object(gathered, local_result) + if rank == 0: + result_path.write_text(json.dumps(gathered, indent=2, sort_keys=True) + "\n", encoding="utf-8") + dist.barrier() + finally: + dist.destroy_process_group() + + +def test_two_gpu_fsdp2_inference_checkpoint_gate(tmp_path: Path) -> None: + if not torch.cuda.is_available(): + pytest.skip("This test requires CUDA and must run on a compute node.") + if torch.cuda.device_count() < WORLD_SIZE: + pytest.skip(f"This test requires at least {WORLD_SIZE} CUDA devices.") + + result_path = tmp_path / "results.json" + command = [ + sys.executable, + "-m", + "torch.distributed.run", + "--standalone", + "--nproc_per_node", + str(WORLD_SIZE), + str(Path(__file__).resolve()), + "--worker", + "--work-dir", + str(tmp_path / "worker"), + "--result", + str(result_path), + ] + environment = os.environ.copy() + environment.setdefault("TORCHDYNAMO_DISABLE", "1") + try: + process = subprocess.run( + command, + capture_output=True, + text=True, + env=environment, + timeout=180, + ) + except subprocess.TimeoutExpired as error: + raise RuntimeError( + "Two-GPU inference checkpoint worker timed out after 180 seconds\n" + f"STDOUT:\n{error.stdout}\nSTDERR:\n{error.stderr}") from error + if process.returncode != 0: + raise RuntimeError( + f"Two-GPU inference checkpoint worker failed with code {process.returncode}\n" + f"STDOUT:\n{process.stdout}\nSTDERR:\n{process.stderr}") + + results = json.loads(result_path.read_text(encoding="utf-8")) + assert len(results) == WORLD_SIZE + assert all(result["success_rng_equal"] for result in results) + assert all(result["failure_rng_equal"] for result in results) + assert all(result["reload_equal"] for result in results) + assert all(result["frozen_present"] for result in results) + errors = [result["failure"] for result in results] + assert all(error is not None and "injected rank-zero export failure" in error for error in errors) + assert len(set(errors)) == 1 + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--worker", action="store_true") + parser.add_argument("--work-dir", type=Path) + parser.add_argument("--result", type=Path) + return parser.parse_args() + + +if __name__ == "__main__": + args = _parse_args() + if not args.worker or args.work_dir is None or args.result is None: + raise SystemExit("--worker, --work-dir, and --result are required") + _run_worker(args.work_dir, args.result) diff --git a/fastvideo/tests/train/utils/test_moduleloader_attention_backend.py b/fastvideo/tests/train/utils/test_moduleloader_attention_backend.py index a7e957c550..580049cedc 100644 --- a/fastvideo/tests/train/utils/test_moduleloader_attention_backend.py +++ b/fastvideo/tests/train/utils/test_moduleloader_attention_backend.py @@ -3,15 +3,22 @@ from __future__ import annotations +from types import SimpleNamespace + import torch import pytest -from fastvideo.attention.selector import _active_component_attention_backend_scope +from fastvideo.attention.selector import ( + _active_component_attention_backend_scope, + _component_attention_backend_scope, +) from fastvideo.configs.pipelines.base import PipelineConfig +from fastvideo.models.loader import component_loader from fastvideo.platforms import AttentionBackendEnum from fastvideo.train.utils import moduleloader from fastvideo.train.utils.training_config import ( DistributedConfig, + ModelTrainingConfig, TrainingConfig, ) @@ -37,7 +44,10 @@ def _fake_load_module(**kwargs): del kwargs scope = _active_component_attention_backend_scope() captured.append((scope.backend, scope.component) if scope else (None, None)) - return torch.nn.Linear(1, 1) + module = torch.nn.Linear(1, 1) + module.config = SimpleNamespace( # type: ignore[attr-defined] + _resolved_attention_backend=scope.backend if scope else None, ) + return module monkeypatch.setattr( moduleloader.PipelineComponentLoader, @@ -57,6 +67,107 @@ def _fake_load_module(**kwargs): assert _active_component_attention_backend_scope() is None +@pytest.mark.parametrize("role", ["student", "teacher", "critic"]) +def test_role_attention_backend_receipt_matches_request( + monkeypatch, + tmp_path, + role: str, +) -> None: + """Every DMD role returns a construction receipt for its explicit backend.""" + training_config = TrainingConfig( + distributed=DistributedConfig(hsdp_shard_dim=1), + pipeline_config=PipelineConfig(), + ) + requested = (AttentionBackendEnum.ATTN_QAT_TRAIN + if role == "student" else AttentionBackendEnum.FLASH_ATTN) + + monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path)) + monkeypatch.setattr( + moduleloader, + "verify_model_config_and_directory", + lambda path: {"transformer": ("diffusers", "FakeTransformer")}, + ) + + def _fake_load_module(**kwargs): + del kwargs + scope = _active_component_attention_backend_scope() + assert scope is not None + module = torch.nn.Linear(1, 1) + module.config = SimpleNamespace( # type: ignore[attr-defined] + _resolved_attention_backend=scope.backend, ) + return module + + monkeypatch.setattr(moduleloader.PipelineComponentLoader, "load_module", _fake_load_module) + + moduleloader.load_module_from_path( + model_path=f"fake/{role}", + module_type="transformer", + training_config=training_config, + disable_custom_init_weights=(role != "student"), + attention_backend=requested, + ) + + +def test_role_attention_backend_receipt_mismatch_fails(monkeypatch, tmp_path) -> None: + """A request that gets narrowed or lost may not silently start training.""" + training_config = TrainingConfig( + distributed=DistributedConfig(hsdp_shard_dim=1), + pipeline_config=PipelineConfig(), + ) + monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path)) + monkeypatch.setattr( + moduleloader, + "verify_model_config_and_directory", + lambda path: {"transformer": ("diffusers", "FakeTransformer")}, + ) + + def _fake_load_module(**kwargs): + del kwargs + module = torch.nn.Linear(1, 1) + module.config = SimpleNamespace( # type: ignore[attr-defined] + _resolved_attention_backend=AttentionBackendEnum.TORCH_SDPA, ) + return module + + monkeypatch.setattr(moduleloader.PipelineComponentLoader, "load_module", _fake_load_module) + + with pytest.raises(RuntimeError, match="requested attention backend FLASH_ATTN.*recorded TORCH_SDPA"): + moduleloader.load_module_from_path( + model_path="fake/teacher", + module_type="transformer", + training_config=training_config, + disable_custom_init_weights=True, + attention_backend="FLASH_ATTN", + ) + + +def test_teacher_critic_preserves_explicit_flash_attention_scope() -> None: + """disable_custom_init_weights must not erase a role-local dense backend.""" + args = SimpleNamespace( + _loading_teacher_critic_model=True, + attention_backend=None, + ) + + with _component_attention_backend_scope(AttentionBackendEnum.FLASH_ATTN, component="teacher"): + with component_loader._teacher_critic_attention_context(args): + scope = _active_component_attention_backend_scope() + assert scope is not None + assert scope.backend is AttentionBackendEnum.FLASH_ATTN + + +def test_teacher_critic_masks_generator_only_qat_attention_scope() -> None: + """The historical student-only QAT policy still narrows teacher/critic.""" + args = SimpleNamespace( + _loading_teacher_critic_model=True, + attention_backend=None, + ) + + with _component_attention_backend_scope(AttentionBackendEnum.ATTN_QAT_TRAIN, component="teacher"): + with component_loader._teacher_critic_attention_context(args): + scope = _active_component_attention_backend_scope() + assert scope is not None + assert scope.backend is None + + def test_load_transformer_restores_backend_when_loading_fails( monkeypatch, tmp_path, @@ -92,3 +203,75 @@ def _raise_during_load(**kwargs): attention_backend="ATTN_QAT_TRAIN", ) assert _active_component_attention_backend_scope() is None + + +def test_training_args_propagate_compile_settings() -> None: + training_config = TrainingConfig( + distributed=DistributedConfig(hsdp_shard_dim=1), + model=ModelTrainingConfig( + enable_torch_compile=True, + torch_compile_kwargs={"dynamic": False}, + ), + pipeline_config=PipelineConfig(), + ) + + args = moduleloader._make_training_args(training_config, model_path="fake/model") + + assert args.enable_torch_compile is True + assert args.torch_compile_kwargs == {"dynamic": False} + + +def test_load_transformer_forwards_pre_fsdp_transform(monkeypatch, tmp_path) -> None: + training_config = TrainingConfig( + distributed=DistributedConfig(hsdp_shard_dim=1), + pipeline_config=PipelineConfig(), + ) + + def transform(module): + return module + + captured = None + + monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path)) + monkeypatch.setattr( + moduleloader, + "verify_model_config_and_directory", + lambda path: {"transformer": ("diffusers", "FakeTransformer")}, + ) + + def _fake_load_module(**kwargs): + nonlocal captured + captured = getattr(kwargs["fastvideo_args"], "_pre_fsdp_transform", None) + return torch.nn.Linear(1, 1) + + monkeypatch.setattr(moduleloader.PipelineComponentLoader, "load_module", _fake_load_module) + + moduleloader.load_module_from_path( + model_path="fake/model", + module_type="transformer", + training_config=training_config, + pre_fsdp_transform=transform, + ) + + assert captured is transform + + +def test_pre_fsdp_transform_rejects_non_transformer(monkeypatch, tmp_path) -> None: + training_config = TrainingConfig( + distributed=DistributedConfig(hsdp_shard_dim=1), + pipeline_config=PipelineConfig(), + ) + monkeypatch.setattr(moduleloader, "maybe_download_model", lambda path: str(tmp_path)) + monkeypatch.setattr( + moduleloader, + "verify_model_config_and_directory", + lambda path: {"vae": ("diffusers", "FakeVAE")}, + ) + + with pytest.raises(ValueError, match="only be set when loading a transformer"): + moduleloader.load_module_from_path( + model_path="fake/model", + module_type="vae", + training_config=training_config, + pre_fsdp_transform=lambda module: module, + ) diff --git a/fastvideo/tests/train/utils/test_torch_compile.py b/fastvideo/tests/train/utils/test_torch_compile.py new file mode 100644 index 0000000000..5efa23e314 --- /dev/null +++ b/fastvideo/tests/train/utils/test_torch_compile.py @@ -0,0 +1,246 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Regional training compile policy regression tests.""" + +from __future__ import annotations + +import pytest +import torch + +from fastvideo.models.loader.fsdp_load import _compile_model_regions +from fastvideo.train.utils.activation_checkpoint import apply_activation_checkpointing + + +class _RepeatedModel(torch.nn.Module): + _compile_conditions = [ + lambda name, module: name.startswith("blocks.") and name.count(".") == 1 + ] + + def __init__(self) -> None: + super().__init__() + self.blocks = torch.nn.ModuleList([ + torch.nn.Linear(4, 4), + torch.nn.Linear(4, 4), + ]) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + for block in self.blocks: + value = block(value) + return value + + +def test_regional_compile_preserves_checkpoint_state_dict(monkeypatch) -> None: + model = apply_activation_checkpointing(_RepeatedModel()) + state_dict_keys = list(model.state_dict()) + calls: list[tuple[torch.nn.Module, dict]] = [] + + def _fake_compile(forward, **kwargs): + calls.append((forward.__self__, kwargs)) + return forward + + monkeypatch.setattr(torch, "compile", _fake_compile) + + assert _compile_model_regions(model, {}) == 2 + assert [target for target, _ in calls] == [ + block._checkpoint_wrapped_module for block in model.blocks + ] + assert [kwargs for _, kwargs in calls] == [ + {"fullgraph": True, "options": {"emulate_precision_casts": True}}, + {"fullgraph": True, "options": {"emulate_precision_casts": True}}, + ] + assert list(model.state_dict()) == state_dict_keys + + +def test_regional_compile_dispatches_grad_and_no_grad_calls(monkeypatch) -> None: + model = _RepeatedModel() + compiled_calls = 0 + + def _fake_compile(eager_forward, **kwargs): + del kwargs + + def _compiled(*args, **forward_kwargs): + nonlocal compiled_calls + compiled_calls += 1 + return eager_forward(*args, **forward_kwargs) + + return _compiled + + monkeypatch.setattr(torch, "compile", _fake_compile) + _compile_model_regions(model, {}) + + value = torch.randn(2, 4) + model(value) + assert compiled_calls == 2 + + with torch.no_grad(): + model(value) + assert compiled_calls == 4 + + +def test_regional_compile_forwards_supported_kwargs(monkeypatch) -> None: + model = _RepeatedModel() + calls: list[dict] = [] + global_recompile_limit = torch._dynamo.config.recompile_limit + + def _fake_compile(forward, **kwargs): + calls.append(kwargs) + return forward + + monkeypatch.setattr(torch, "compile", _fake_compile) + + _compile_model_regions(model, { + "dynamic": False, + "recompile_limit": 32, + "options": { + "emulate_precision_casts": False, + "max_autotune": True, + }, + }) + + assert calls == [ + { + "fullgraph": True, + "dynamic": False, + "recompile_limit": 32, + "options": { + "emulate_precision_casts": False, + "max_autotune": True, + }, + }, + { + "fullgraph": True, + "dynamic": False, + "recompile_limit": 32, + "options": { + "emulate_precision_casts": False, + "max_autotune": True, + }, + }, + ] + assert torch._dynamo.config.recompile_limit == global_recompile_limit + + +def test_regional_compile_rejects_partial_graph_mode() -> None: + with pytest.raises(ValueError, match="fullgraph=True"): + _compile_model_regions(_RepeatedModel(), {"fullgraph": False}) + + +def test_regional_compile_rejects_mode_kwarg() -> None: + """`mode` conflicts with the always-injected inductor options. + + torch.compile forbids mode+options together; the loader must fail with an + actionable message rather than letting torch blame an `options` key the + user never wrote (the CLI help's own example uses `mode`). + """ + model = _RepeatedModel() + with pytest.raises(ValueError, match="mode"): + _compile_model_regions(model, {"mode": "reduce-overhead"}) + + +def test_checkpoint_wrapper_prefix_normalization() -> None: + """AC-wrapped blocks must not break name-keyed weight-loader lookups. + + checkpoint_wrapper strips its prefix from state_dict() keys via hooks but + NOT from named_parameters()/named_buffers(); checkpoint keys are clean. + Pre-fix, a loaded buffer inside a wrapped block missed the named_buffers + membership test and was silently converted into a trainable nn.Parameter + by load_state_dict(assign=True). + """ + from fastvideo.models.loader.fsdp_load import _strip_checkpoint_wrapper_prefix + + class _BufferBlock(torch.nn.Module): + + def __init__(self) -> None: + super().__init__() + self.lin = torch.nn.Linear(4, 4) + self.register_buffer("freq", torch.arange(4.0)) + + def forward(self, value: torch.Tensor) -> torch.Tensor: + return self.lin(value) + self.freq + + class _BufferModel(torch.nn.Module): + + def __init__(self) -> None: + super().__init__() + self.blocks = torch.nn.ModuleList([_BufferBlock()]) + + from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( + checkpoint_wrapper, ) + + model = _BufferModel() + model.blocks[0] = checkpoint_wrapper(model.blocks[0]) + + # state_dict is clean; raw named_buffers is prefixed. + assert "blocks.0.freq" in model.state_dict() + raw_buffer_names = {name for name, _ in model.named_buffers()} + assert "blocks.0.freq" not in raw_buffer_names + assert "blocks.0._checkpoint_wrapped_module.freq" in raw_buffer_names + + # The canonicalized views match checkpoint keys exactly. + clean_buffers = {_strip_checkpoint_wrapper_prefix(name) for name, _ in model.named_buffers()} + clean_params = {_strip_checkpoint_wrapper_prefix(name) for name, _ in model.named_parameters()} + assert clean_buffers == {"blocks.0.freq"} + assert clean_params == {"blocks.0.lin.weight", "blocks.0.lin.bias"} + + # End-to-end: membership keyed on canonical names keeps a loaded buffer a + # buffer under load_state_dict(assign=True) instead of promoting it to a + # trainable parameter. + loaded = { + "blocks.0.freq": torch.ones(4), + "blocks.0.lin.weight": torch.ones(4, 4), + "blocks.0.lin.bias": torch.ones(4), + } + sharded_sd = { + key: (value if key in clean_buffers else torch.nn.Parameter(value)) + for key, value in loaded.items() + } + model.load_state_dict(sharded_sd, assign=True) + assert any("freq" in name for name, _ in model.named_buffers()) + assert not any("freq" in name for name, _ in model.named_parameters()) + + +def test_regional_compile_unsupported_when_attention_compile_disabled(monkeypatch) -> None: + """The attention-eager escape hatch must degrade the role, not crash. + + FASTVIDEO_DISABLE_ATTENTION_COMPILE=1 wraps attention forwards in + torch.compiler.disable; under fullgraph regional compile that raises + `torch._dynamo.exc.Unsupported: Skip inlining torch.compiler.disable()'d + function` at the first training step (observed on the h3-compile-ab + b4_on_attneager leg, job 2610). The guard must reject regional compile + for the role so it falls back to eager with the standard warning. + """ + from fastvideo.models.loader.fsdp_load import _regional_compile_unsupported_reason + + monkeypatch.setenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", "1") + reason = _regional_compile_unsupported_reason({"config": None}) + assert reason is not None + assert "FASTVIDEO_DISABLE_ATTENTION_COMPILE" in reason + + monkeypatch.setenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", "0") + assert _regional_compile_unsupported_reason({"config": None}) is None + + +def test_regional_compile_unsupported_for_vsa_backends() -> None: + """A VSA-backed role must fall back to eager instead of hard-failing. + + The VSA backends (Triton block-sparse kernels behind sequence-parallel + all-to-alls plus a host-synced metadata guard) are not fullgraph-traceable. + enable_torch_compile is a per-run switch shared by every DMD role, so the + loader must skip the VSA student while dense FLASH_ATTN/SDPA roles compile. + """ + from fastvideo.models.loader.fsdp_load import _regional_compile_unsupported_reason + from fastvideo.platforms import AttentionBackendEnum + + class _Config: + pass + + for backend in (AttentionBackendEnum.VIDEO_SPARSE_ATTN, AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3): + config = _Config() + config._resolved_attention_backend = backend + reason = _regional_compile_unsupported_reason({"config": config}) + assert reason is not None + assert backend.name in reason + + sdpa_config = _Config() + sdpa_config._resolved_attention_backend = AttentionBackendEnum.TORCH_SDPA + assert _regional_compile_unsupported_reason({"config": sdpa_config}) is None + assert _regional_compile_unsupported_reason({"config": None}) is None diff --git a/fastvideo/tests/train/utils/test_tracking.py b/fastvideo/tests/train/utils/test_tracking.py new file mode 100644 index 0000000000..9f20538159 --- /dev/null +++ b/fastvideo/tests/train/utils/test_tracking.py @@ -0,0 +1,136 @@ +# SPDX-License-Identifier: Apache-2.0 +"""CPU tests for the modular trainer's wandb tracking guard. + +The Trainer logs each step's loss dict through ``build_tracker``'s tracker on +rank 0; these tests pin the graceful no-op path (wandb missing / no +credentials) and the logged loss-dict shape with a monkeypatched wandb. +""" + +import builtins +import sys +from types import SimpleNamespace + +from fastvideo.train.utils import tracking +from fastvideo.train.utils.training_config import ( + CheckpointConfig, + TrackerConfig, +) +from fastvideo.training.trackers import DummyTracker, WandbTracker + + +class _FakeRun: + + def __init__(self, **init_kwargs): + self.init_kwargs = init_kwargs + self.logged = [] + + def log(self, metrics, step=None): + self.logged.append((metrics, step)) + + def finish(self): + pass + + +def _fake_wandb(): + runs = [] + + def _init(**kwargs): + run = _FakeRun(**kwargs) + runs.append(run) + return run + + return SimpleNamespace(init=_init, runs=runs, api=SimpleNamespace(api_key="key")) + + +def test_wandb_usable_false_when_not_importable(monkeypatch) -> None: + real_import = builtins.__import__ + + def _blocked(name, *args, **kwargs): + if name == "wandb": + raise ImportError("no wandb") + return real_import(name, *args, **kwargs) + + monkeypatch.setattr(builtins, "__import__", _blocked) + monkeypatch.delitem(sys.modules, "wandb", raising=False) + assert tracking._wandb_usable() is False + + +def test_wandb_usable_false_without_credentials(monkeypatch) -> None: + monkeypatch.delenv("WANDB_API_KEY", raising=False) + monkeypatch.delenv("WANDB_MODE", raising=False) + monkeypatch.setitem( + sys.modules, + "wandb", + SimpleNamespace(api=SimpleNamespace(api_key=None)), + ) + assert tracking._wandb_usable() is False + + +def test_wandb_usable_with_env_key(monkeypatch) -> None: + monkeypatch.setenv("WANDB_API_KEY", "key") + monkeypatch.setitem(sys.modules, "wandb", SimpleNamespace()) + assert tracking._wandb_usable() is True + + +def test_build_tracker_noops_cleanly_without_wandb(monkeypatch, tmp_path) -> None: + """A wandb-project config must degrade to the dummy tracker, not crash.""" + monkeypatch.setattr(tracking, "get_world_group", lambda: SimpleNamespace(rank=0)) + monkeypatch.setattr(tracking, "_wandb_usable", lambda: False) + + tracker = tracking.build_tracker( + TrackerConfig(project_name="h3-dmd2-vsa"), + CheckpointConfig(output_dir=str(tmp_path)), + config={"method": {}}, + ) + + assert isinstance(tracker, DummyTracker) + tracker.log({"total_loss": 1.0}, 1) # must not raise + tracker.finish() + + +def test_build_tracker_logs_loss_dict_with_monkeypatched_wandb(monkeypatch, tmp_path) -> None: + fake = _fake_wandb() + monkeypatch.setitem(sys.modules, "wandb", fake) + monkeypatch.setenv("WANDB_API_KEY", "key") + monkeypatch.setattr(tracking, "get_world_group", lambda: SimpleNamespace(rank=0)) + + run_config = {"method": {"dmd_denoising_steps": [1000, 757, 522]}} + tracker = tracking.build_tracker( + TrackerConfig(project_name="h3-dmd2-vsa", run_name="dmd2_vsa0_overfit"), + CheckpointConfig(output_dir=str(tmp_path)), + config=run_config, + ) + + assert isinstance(tracker, WandbTracker) + (run, ) = fake.runs + assert run.init_kwargs["project"] == "h3-dmd2-vsa" + assert run.init_kwargs["name"] == "dmd2_vsa0_overfit" + assert run.init_kwargs["config"] == run_config + + # The per-step dict the Trainer logs on rank 0 (DMD2 loss map + metrics). + metrics = { + "total_loss": 0.5, + "generator_loss": 0.25, + "fake_score_loss": 0.25, + "update_student": 1.0, + "step_time_sec": 0.1, + "vsa_sparsity": 0.0, + } + tracker.log(metrics, 7) + assert run.logged == [(metrics, 7)] + + +def test_build_tracker_nonzero_rank_never_inits_wandb(monkeypatch, tmp_path) -> None: + fake = _fake_wandb() + monkeypatch.setitem(sys.modules, "wandb", fake) + monkeypatch.setenv("WANDB_API_KEY", "key") + monkeypatch.setattr(tracking, "get_world_group", lambda: SimpleNamespace(rank=1)) + + tracker = tracking.build_tracker( + TrackerConfig(project_name="h3-dmd2-vsa"), + CheckpointConfig(output_dir=str(tmp_path)), + config=None, + ) + + assert isinstance(tracker, DummyTracker) + assert fake.runs == [] diff --git a/fastvideo/train/attn_qat/README.md b/fastvideo/train/attn_qat/README.md index 55edd75a1f..343aae1b88 100644 --- a/fastvideo/train/attn_qat/README.md +++ b/fastvideo/train/attn_qat/README.md @@ -71,7 +71,7 @@ The migration preserves these training semantics: |---|---| | Student fake-quantized attention | `models.student.attention_backend: ATTN_QAT_TRAIN` | | Teacher/critic full-precision attention | Role-local `FLASH_ATTN` | -| Generator update every five critic steps | `method.generator_update_interval: 5` | +| Four critic-only steps, then one student-only step | `method.generator_update_interval: 5` | | Three-step rollout | `method.dmd_denoising_steps: [1000, 757, 522]` | | Score timestep range | `method.min_timestep_ratio: 0.02`, `max_timestep_ratio: 0.98` | | Legacy teacher guidance `cond + 2(cond-uncond)` | Standard CFG scale `3.0` | diff --git a/fastvideo/train/callbacks/callback.py b/fastvideo/train/callbacks/callback.py index ca057c151b..c4ca9e699f 100644 --- a/fastvideo/train/callbacks/callback.py +++ b/fastvideo/train/callbacks/callback.py @@ -1,8 +1,5 @@ # SPDX-License-Identifier: Apache-2.0 -"""Callback base class and CallbackDict manager. - -Adapted from FastGen's callback pattern to FastVideo's types. -""" +"""Callback base class and CallbackDict manager.""" from __future__ import annotations @@ -24,6 +21,7 @@ "grad_clip": "fastvideo.train.callbacks.grad_clip.GradNormClipCallback", "validation": "fastvideo.train.callbacks.validation.ValidationCallback", "ema": "fastvideo.train.callbacks.ema.EMACallback", + "latent_vis": "fastvideo.train.callbacks.latent_vis.LatentVisCallback", } @@ -73,6 +71,11 @@ def on_validation_begin( ) -> None: pass + def will_run_validation(self, iteration: int = 0) -> bool: + """Return whether this callback will validate at ``iteration``.""" + del iteration + return False + def on_validation_end( self, method: TrainingMethod, @@ -179,3 +182,7 @@ def _dispatch(*args: Any, **kwargs: Any) -> None: fn(*args, **kwargs) return _dispatch + + def will_run_validation(self, iteration: int = 0) -> bool: + """Return whether any configured callback schedules validation now.""" + return any(cb.will_run_validation(iteration) for cb in self._callbacks.values()) diff --git a/fastvideo/train/callbacks/ema.py b/fastvideo/train/callbacks/ema.py index 2bd9f01f09..9ad7579256 100644 --- a/fastvideo/train/callbacks/ema.py +++ b/fastvideo/train/callbacks/ema.py @@ -92,6 +92,9 @@ def on_training_step_end( if self.student_ema is None: return + student_optimizer = getattr(method, "_student_optimizer", None) + if student_optimizer is not None and student_optimizer not in method.get_optimizers(iteration): + return if iteration < self._start_iter: return if not self._ema_started: diff --git a/fastvideo/train/callbacks/latent_vis.py b/fastvideo/train/callbacks/latent_vis.py new file mode 100644 index 0000000000..4d72c00a80 --- /dev/null +++ b/fastvideo/train/callbacks/latent_vis.py @@ -0,0 +1,93 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Intermediate-latent visualization callback for distillation methods. + +Port of the legacy ``fastvideo/training/distillation_pipeline.py`` latent +logging to the modular trainer: every ``every_steps`` iterations, rank 0 +decodes the method's latest latent snapshots (``method.latent_vis`` — the +student's rollout prediction plus the real- and fake-score predictions on +generator-update steps) through the model's ``decode_vis_latents`` hook and +logs them to the tracker as videos. + +Both hooks are optional: methods that never populate ``latent_vis`` or models +without ``decode_vis_latents`` turn this callback into a no-op. +""" + +from __future__ import annotations + +from typing import Any, TYPE_CHECKING + +import torch + +from fastvideo.distributed import get_world_group +from fastvideo.logger import init_logger +from fastvideo.train.callbacks.callback import Callback +from fastvideo.training.trackers import DummyTracker + +if TYPE_CHECKING: + from fastvideo.train.methods.base import TrainingMethod + +logger = init_logger(__name__) + +_DEFAULT_KEYS = ( + "generator_pred_video", + "real_score_pred_video", + "faker_score_pred_video", +) + + +class LatentVisCallback(Callback): + """Decode and log intermediate training latents as tracker videos.""" + + def __init__( + self, + every_steps: int = 100, + keys: list[str] | None = None, + fps: int = 24, + ) -> None: + self.every_steps = int(every_steps) + self.keys = tuple(keys) if keys else _DEFAULT_KEYS + self.fps = int(fps) + self.tracker: Any = DummyTracker() + + def on_train_start( + self, + method: TrainingMethod, + iteration: int = 0, + ) -> None: + tracker = getattr(method, "tracker", None) + if tracker is not None: + self.tracker = tracker + + def on_training_step_end( + self, + method: TrainingMethod, + loss_dict: dict[str, Any], + iteration: int = 0, + ) -> None: + if self.every_steps <= 0 or iteration % self.every_steps != 0: + return + if get_world_group().rank != 0: + return + vis = getattr(method, "latent_vis", None) + if not vis: + return + decode = getattr(getattr(method, "student", None), "decode_vis_latents", None) + if decode is None: + return + + artifacts: dict[str, Any] = {} + latent_layout = vis.get("_fv_latent_layout") + for key in self.keys: + latent = vis.get(key) + if not isinstance(latent, torch.Tensor): + continue + try: + clip = decode(latent) if latent_layout is None else decode(latent, layout=latent_layout) + except Exception as exc: + logger.warning("Latent visualization decode failed for %r: %s", key, exc) + continue + art = self.tracker.video(clip, fps=self.fps, format="mp4") + if art is not None: + artifacts[f"latent_vis/{key}"] = art + if artifacts: + self.tracker.log_artifacts(artifacts, iteration) diff --git a/fastvideo/train/callbacks/validation.py b/fastvideo/train/callbacks/validation.py index e83775ed04..770737a3cf 100644 --- a/fastvideo/train/callbacks/validation.py +++ b/fastvideo/train/callbacks/validation.py @@ -12,6 +12,7 @@ import gc import json import os +import re import time from copy import deepcopy from dataclasses import dataclass, field @@ -57,6 +58,7 @@ class _ValidationStepResult: overlay_videos: list[list[np.ndarray]] = field(default_factory=list) overlay_captions: list[str] = field(default_factory=list) ref_videos: list[str | None] = field(default_factory=list) + metadata: list[dict[str, Any]] = field(default_factory=list) actions: list[dict[str, Any] | None] = field(default_factory=list) mouse_pitch_signs: list[int | None] = field(default_factory=list) @@ -136,6 +138,8 @@ def __init__( sampling_steps: list[int] | None = None, guidance_scale: float | None = None, num_frames: int | None = None, + use_record_dimensions: bool = False, + max_record_num_frames: int | None = None, num_videos_per_prompt: int = 1, use_validation_media_conditioning: bool = True, output_dir: str | None = None, @@ -160,6 +164,10 @@ def __init__( self.sampling_steps = ([int(s) for s in sampling_steps] if sampling_steps else [40]) self.guidance_scale = (float(guidance_scale) if guidance_scale is not None else None) self.num_frames = (int(num_frames) if num_frames is not None else None) + self.use_record_dimensions = self._coerce_bool(use_record_dimensions) + self.max_record_num_frames = (int(max_record_num_frames) if max_record_num_frames is not None else None) + if self.max_record_num_frames is not None and self.max_record_num_frames <= 0: + raise ValueError("callbacks.validation.max_record_num_frames must be positive") self.num_videos_per_prompt = int(num_videos_per_prompt) if self.num_videos_per_prompt <= 0: raise ValueError("callbacks.validation.num_videos_per_prompt must be positive") @@ -240,18 +248,46 @@ def _coerce_bool(value: Any) -> bool: # Callback hooks # ---------------------------------------------------------- - def _adopt_training_denoising_ladder( + @staticmethod + def _assert_attention_contract( + inference_args: Any, + tc: Any, + ) -> None: + """Fail if validation would sample off the training attention contract. + + Sparsity and tile geometry are separate knobs and both have to survive + the training-config to inference-args hop: v8 validated a + tile-64-trained student at the tile-256 default because only the + sparsity was propagated, and nothing downstream noticed. + """ + for train_attr, args_attr in ( + ("vsa_sparsity", "VSA_sparsity"), + ("vsa_tile_size", "VSA_tile_size"), + ): + trained = getattr(tc, train_attr, None) + sampled = getattr(inference_args, args_attr, None) + if trained is None or sampled is None: + continue + if sampled != trained: + raise ValueError(f"validation would sample at {args_attr}={sampled} while " + f"training runs {train_attr}={trained}. The attention " + "contract must match; fix the training-config to " + "inference-args propagation rather than the symptom.") + + def _adopt_training_sampling_contract( self, method: TrainingMethod, ) -> None: - """Keep validation on the timestep ladder the method trains against. - - ``sampling_steps`` only sets ``num_inference_steps``, which the - scheduler turns into an N-point sigma grid -- N-1 forwards on its own - spacing, not the trained ladder. Validation reaches the trained - operating point only when ``sampling_timesteps`` repeats - ``method.dmd_denoising_steps`` exactly, so inherit it by default and - refuse to run when the two disagree. + """Keep validation sampling on the operating point training teaches. + + A few-step method trains its student on an explicit timestep ladder. + Validation only reaches that ladder when ``sampling_timesteps`` is + configured: ``sampling_steps`` sets ``num_inference_steps``, and the + public scheduler turns N of those into an N-point sigma grid, i.e. + N-1 forwards on the scheduler's own spacing. Duplicating the ladder by + hand in the callback config is the footgun that silently validated v8 + at three forwards on the wrong grid for its whole run, so derive it + from the method instead, and refuse to run when the two disagree. """ method_config = getattr(method, "method_config", None) if not isinstance(method_config, dict): @@ -264,8 +300,8 @@ def _adopt_training_denoising_ladder( if self.sampling_timesteps is None: self.sampling_timesteps = trained logger.info( - "validation: inheriting the trained denoising ladder %s " - "(%d forwards)", + "validation: adopting the trained denoising ladder %s " + "(%d forwards) from the training method", trained, len(trained), ) @@ -274,8 +310,9 @@ def _adopt_training_denoising_ladder( raise ValueError("callbacks.validation.sampling_timesteps " f"{self.sampling_timesteps} disagrees with the trained ladder " f"{trained} (method.dmd_denoising_steps). Validation would " - "sample off the operating point the student is trained for. " - "Drop the override to inherit the ladder, or align the two.") + "sample off the operating point the student was distilled " + "for. Drop the callback override to inherit the ladder, or " + "align the two deliberately.") def on_train_start( self, @@ -284,7 +321,7 @@ def on_train_start( ) -> None: self.method = method tc = self.training_config - self._adopt_training_denoising_ladder(method) + self._adopt_training_sampling_contract(method) self.world_group = get_world_group() self.sp_group = get_sp_group() @@ -309,16 +346,20 @@ def on_validation_begin( iteration: int = 0, ) -> None: """Run the optional step-zero baseline and each scheduled validation event.""" - if self.every_steps <= 0: - return - # Step zero measures the checkpoint before the first optimizer update. - if iteration == 0 and not self.run_at_start: - return - if iteration % self.every_steps != 0: + if not self.will_run_validation(iteration): return self._run_validation(method, iteration) + def will_run_validation(self, iteration: int = 0) -> bool: + """Return whether this callback schedules validation at ``iteration``.""" + if self.every_steps <= 0: + return False + # Step zero measures the checkpoint before the first optimizer update. + if iteration == 0 and not self.run_at_start: + return False + return iteration % self.every_steps == 0 + # ---------------------------------------------------------- # Core validation logic # ---------------------------------------------------------- @@ -420,6 +461,17 @@ def _validation_memory_context( self._restore_inactive_role_modules(module_records) self._restore_optimizer_states(optimizer_tensor_records) self._empty_cuda_cache() + off_cuda = [ + f"{name} on {getattr(p, '_local_tensor', p).device}" + for name, p in validation_transformer.named_parameters() + if getattr(p, "_local_tensor", p).device.type != "cuda" + ] + if off_cuda: + logger.warning( + "Post-validation: %d validation-transformer params off-CUDA; first: %s", + len(off_cuda), + off_cuda[:5], + ) def _offload_optimizer_states_to_cpu( self, @@ -504,6 +556,20 @@ def _offload_inactive_role_modules_to_cpu( device = self._first_cuda_tensor_device(module) if device is None: continue + if self._is_fsdp_managed(module): + # `.to()` round-trips on fully_shard modules replace DTensor + # local storage behind FSDP2's bookkeeping; the corruption + # surfaces as device-mismatch errors on the first backward + # through the module a few steps after restore. Keep sharded + # roles resident — the optimizer-state offload above already + # returns the bulk of the memory. + logger.info( + "Keeping role %r transformer on %s during validation " + "(FSDP-managed modules do not survive .to() round-trips).", + role, + device, + ) + continue try: module.to("cpu") except Exception as exc: @@ -532,6 +598,17 @@ def _restore_inactive_role_modules( device, ) + @staticmethod + def _is_fsdp_managed(module: torch.nn.Module) -> bool: + try: + from torch.distributed.fsdp import FSDPModule + from torch.distributed.tensor import DTensor + except ImportError: + return False + if isinstance(module, FSDPModule): + return True + return any(isinstance(p, DTensor) for p in module.parameters(recurse=True)) + @staticmethod def _first_cuda_tensor_device(module: torch.nn.Module) -> torch.device | None: for tensor in list(module.parameters(recurse=True)) + list(module.buffers(recurse=True)): @@ -630,6 +707,10 @@ def _run_validation_inner( result.ref_videos, local_videos.indices, ) + local_metadata = self._select_by_indices( + result.metadata, + local_videos.indices, + ) local_actions = self._select_by_indices( result.actions, local_videos.indices, @@ -655,6 +736,8 @@ def _run_validation_inner( all_overlay_video_filenames = list(local_overlay_video_filenames) all_captions = list(local_captions) all_overlay_captions = list(local_overlay_captions) + all_ref_videos = list(local_ref_videos) + all_metadata = list(local_metadata) all_audio_video_count = local_videos.audio_video_count all_metric_stats = local_metric_stats for sp_idx in range(1, num_sp_groups): @@ -665,12 +748,16 @@ def _run_validation_inner( recv_c = (self.world_group.recv_object(src=src)) recv_ov = (self.world_group.recv_object(src=src)) recv_oc = (self.world_group.recv_object(src=src)) + recv_ref = (self.world_group.recv_object(src=src)) + recv_metadata = (self.world_group.recv_object(src=src)) recv_m = (self.world_group.recv_object(src=src)) recv_audio_video_count = (self.world_group.recv_object(src=src)) all_video_filenames.extend(recv_v) all_overlay_video_filenames.extend(recv_ov) all_captions.extend(recv_c) all_overlay_captions.extend(recv_oc) + all_ref_videos.extend(recv_ref) + all_metadata.extend(recv_metadata) all_audio_video_count += int(recv_audio_video_count) self._merge_metric_stats( all_metric_stats, @@ -681,14 +768,39 @@ def _run_validation_inner( all_metric_stats, step=step, ) + display_captions = [ + self._validation_artifact_caption(caption, metadata) + for caption, metadata in zip(all_captions, all_metadata, strict=True) + ] + reference_filenames: list[str] = [] + reference_captions: list[str] = [] + for caption, ref_video, metadata in zip( + all_captions, + all_ref_videos, + all_metadata, + strict=True, + ): + if ref_video is None or not os.path.isfile(ref_video): + continue + reference_filenames.append(ref_video) + reference_captions.append( + self._validation_artifact_caption( + caption, + metadata, + prefix="held-out reference", + use_reference_num_frames=True, + )) # Media and completion counts share one tracker event so # artifacts and verification data remain aligned. self._log_validation_video_artifacts( all_video_filenames, - all_captions, + display_captions, key=f"validation_videos_{num_inference_steps}_steps", step=step, fps=sp.fps, + reference_video_filenames=reference_filenames, + reference_captions=reference_captions, + reference_key=f"validation_references_{num_inference_steps}_steps", scalar_metrics={ f"validation/{num_inference_steps}_steps_video_count": float(len(all_video_filenames)), @@ -696,6 +808,12 @@ def _run_validation_inner( (time.perf_counter() - validation_started_at), f"validation/{num_inference_steps}_steps_audio_video_count": float(all_audio_video_count), + f"validation/{num_inference_steps}_steps_reference_video_count": + float(len(reference_filenames)), + **self._validation_metadata_scalar_metrics( + all_metadata, + num_inference_steps=num_inference_steps, + ), }, ) if all_overlay_video_filenames: @@ -724,6 +842,14 @@ def _run_validation_inner( local_overlay_captions, dst=0, ) + self.world_group.send_object( + local_ref_videos, + dst=0, + ) + self.world_group.send_object( + local_metadata, + dst=0, + ) self.world_group.send_object( local_metric_stats, dst=0, @@ -818,6 +944,9 @@ def _log_validation_video_artifacts( step: int, fps: int, scalar_metrics: dict[str, float] | None = None, + reference_video_filenames: list[str] | None = None, + reference_captions: list[str] | None = None, + reference_key: str | None = None, ) -> None: """Log validation media and its scalar verification data at one step.""" video_logs = [] @@ -833,15 +962,87 @@ def _log_validation_video_artifacts( ) if art is not None: video_logs.append(art) + artifacts: dict[str, Any] = {} if video_logs: - artifacts: dict[str, Any] = {key: video_logs} - if scalar_metrics: - artifacts.update(scalar_metrics) + artifacts[key] = video_logs + if ((reference_video_filenames is None) != (reference_captions is None) + or (reference_video_filenames is not None) != (reference_key is not None)): + raise ValueError("Validation reference filenames, captions, and key must be provided together.") + if reference_video_filenames is not None: + reference_logs = [] + for fname, cap in zip( + reference_video_filenames, + reference_captions or [], + strict=True, + ): + art = self.tracker.video( + fname, + caption=cap, + fps=fps, + ) + if art is not None: + reference_logs.append(art) + if reference_logs: + assert reference_key is not None + artifacts[reference_key] = reference_logs + if scalar_metrics: + artifacts.update(scalar_metrics) + if artifacts: self.tracker.log_artifacts( artifacts, step, ) + @staticmethod + def _validation_artifact_caption( + caption: str, + metadata: dict[str, Any], + *, + prefix: str = "generated", + use_reference_num_frames: bool = False, + ) -> str: + fields = [prefix] + source = metadata.get("source") + sample_id = metadata.get("sample_id") + if source: + fields.append(f"source={source}") + if sample_id: + fields.append(f"id={sample_id}") + width = metadata.get("width") + height = metadata.get("height") + num_frames = (metadata.get("reference_num_frames", metadata.get("num_frames")) + if use_reference_num_frames else metadata.get("num_frames")) + if width and height and num_frames: + fields.append(f"shape={width}x{height}x{num_frames}f") + return f"[{' | '.join(fields)}] {caption}" + + @staticmethod + def _validation_metric_segment(value: Any) -> str: + segment = re.sub(r"[^A-Za-z0-9_.-]+", "_", str(value).strip()) + return segment.strip("_") or "unknown" + + @classmethod + def _validation_metadata_scalar_metrics( + cls, + metadata: list[dict[str, Any]], + *, + num_inference_steps: int, + ) -> dict[str, float]: + metrics: dict[str, float] = {} + prefix = f"validation/{num_inference_steps}_steps" + for record in metadata: + source = cls._validation_metric_segment(record.get("source", "unknown")) + source_key = f"{prefix}/source/{source}_count" + metrics[source_key] = metrics.get(source_key, 0.0) + 1.0 + width = record.get("width") + height = record.get("height") + num_frames = record.get("num_frames") + if width and height and num_frames: + shape = f"{int(width)}x{int(height)}x{int(num_frames)}f" + shape_key = f"{prefix}/shape/{shape}_count" + metrics[shape_key] = metrics.get(shape_key, 0.0) + 1.0 + return metrics + # ---------------------------------------------------------- # Metric evaluation # ---------------------------------------------------------- @@ -1275,6 +1476,25 @@ def _sync_runtime_dit_arch_config( getattr(transformer, name), ) + def _inject_method_denoising_steps(self, validation_config: Any) -> None: + """Sample validation at the method's trained DMD jump points. + + Distillation students only ever denoise from ``dmd_denoising_steps``; + stages that honor ``pipeline_config.dmd_denoising_steps`` should visit + the same points instead of the scheduler's native grid. An explicit + value already on the pipeline config wins, and warped lists stay out + (their raw entries are grid indices, not timesteps). + """ + method_config = getattr(self.method, "method_config", None) + if not isinstance(method_config, dict): + return + steps = method_config.get("dmd_denoising_steps") + if not steps or bool(method_config.get("warp_denoising_step", False)): + return + if getattr(validation_config, "dmd_denoising_steps", "unset") is not None: + return + validation_config.dmd_denoising_steps = [int(step) for step in steps] + def _get_pipeline( self, *, @@ -1326,6 +1546,7 @@ def _get_pipeline( validation_config, loaded_config, ) + self._inject_method_denoising_steps(validation_config) self._pipeline.fastvideo_args.pipeline_config = validation_config arch_config = self._pipeline.fastvideo_args.pipeline_config.dit_config.arch_config logger.info( @@ -1352,8 +1573,9 @@ def _prepare_validation_batch( tc = self.training_config sampling_param.prompt = validation_batch["prompt"] - sampling_param.height = tc.data.num_height - sampling_param.width = tc.data.num_width + height, width, num_frames = self._validation_sampling_dimensions(validation_batch) + sampling_param.height = height + sampling_param.width = width sampling_param.num_inference_steps = int(num_inference_steps) sampling_param.data_type = "video" if self.guidance_scale is not None: @@ -1372,14 +1594,7 @@ def _prepare_validation_batch( if img_path is not None and (img_path.startswith("http") or os.path.isfile(img_path)): sampling_param.image_path = img_path - temporal_compression_factor = int( - tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio # type: ignore[union-attr] - ) - default_num_frames = ((tc.data.num_latent_t - 1) * temporal_compression_factor + 1) - if self.num_frames is not None: - sampling_param.num_frames = int(self.num_frames) - else: - sampling_param.num_frames = int(default_num_frames) + sampling_param.num_frames = num_frames latents_size = [ (sampling_param.num_frames - 1) // 4 + 1, @@ -1397,6 +1612,7 @@ def _prepare_validation_batch( tc, model_path=tc.model_path, ) + self._assert_attention_contract(inference_args, tc) batch = ForwardBatch( **shallow_asdict(sampling_param), @@ -1429,6 +1645,50 @@ def _prepare_validation_batch( return batch + def _validation_sampling_dimensions( + self, + validation_batch: dict[str, Any], + ) -> tuple[int, int, int]: + """Resolve output geometry, optionally from one validation record. + + Native-shape validation is explicit because cached ``SamplingParam`` + instances are shared across records. A complete record triplet wins; + partial metadata fails instead of combining dimensions from unrelated + shapes. ``max_record_num_frames`` optionally caps only the temporal + member of a complete record triplet; fixed/default geometry is never + changed. With the option off (the default), legacy callback/config + behavior is unchanged. + """ + tc = self.training_config + temporal_compression_factor = int( + tc.pipeline_config.vae_config.arch_config.temporal_compression_ratio # type: ignore[union-attr] + ) + default_num_frames = ((tc.data.num_latent_t - 1) * temporal_compression_factor + 1) + dimensions = { + "height": int(tc.data.num_height), + "width": int(tc.data.num_width), + "num_frames": (int(self.num_frames) if self.num_frames is not None else int(default_num_frames)), + } + + if self.use_record_dimensions: + present = {name: validation_batch.get(name) is not None for name in dimensions} + if any(present.values()) and not all(present.values()): + missing = [name for name, is_present in present.items() if not is_present] + raise ValueError("Native-shape validation records must provide width, height, and num_frames together; " + f"missing {missing} for prompt {validation_batch.get('prompt')!r}.") + if all(present.values()): + dimensions = {name: int(validation_batch[name]) for name in dimensions} + if self.max_record_num_frames is not None: + dimensions["num_frames"] = min(dimensions["num_frames"], self.max_record_num_frames) + + for name, value in dimensions.items(): + if value <= 0: + raise ValueError(f"Validation {name} must be positive, got {value}") + if dimensions["height"] % 8 or dimensions["width"] % 8: + raise ValueError("Validation width and height must be divisible by 8, got " + f"{dimensions['width']}x{dimensions['height']}") + return dimensions["height"], dimensions["width"], dimensions["num_frames"] + def _attach_action_conditions( self, batch: ForwardBatch, @@ -1494,6 +1754,7 @@ def _run_validation_for_steps( tc, model_path=tc.model_path, ) + self._assert_attention_contract(inference_args, tc) self._sync_runtime_dit_arch_config( inference_args.pipeline_config, transformer, @@ -1515,6 +1776,7 @@ def _run_validation_for_steps( captions: list[str] = [] overlay_captions: list[str] = [] ref_videos: list[str | None] = [] + metadata: list[dict[str, Any]] = [] actions: list[dict[str, Any] | None] = [] mouse_pitch_signs: list[int | None] = [] @@ -1526,7 +1788,11 @@ def _run_validation_for_steps( ) assert (batch.prompt is not None and isinstance(batch.prompt, str)) - ref_video = validation_batch.get("ref_video") + # Text-only validation may still carry the held-out raw video for + # side-by-side logging. ``ref_video`` avoids decoding it during + # dataset iteration; ``video_path`` remains a backward-compatible + # fallback for existing manifests. + ref_video = (validation_batch.get("ref_video") or validation_batch.get("video_path")) action = self._validation_actions(validation_batch) with torch.no_grad(): @@ -1571,6 +1837,17 @@ def _run_validation_for_steps( audio_waveforms.append(output_audio) audio_sample_rates.append(int(output_audio_sample_rate) if output_audio_sample_rate is not None else None) ref_videos.append(ref_video if isinstance(ref_video, str) else None) + record_metadata: dict[str, Any] = { + "source": validation_batch.get("source", "unknown"), + "sample_id": validation_batch.get("sample_id", validation_batch.get("id")), + "width": int(batch.width), + "height": int(batch.height), + "num_frames": int(batch.num_frames), + } + reference_num_frames = validation_batch.get("num_frames") + if reference_num_frames is not None and int(reference_num_frames) != int(batch.num_frames): + record_metadata["reference_num_frames"] = int(reference_num_frames) + metadata.append(record_metadata) actions.append(action) mouse_pitch_signs.append(self._validation_mouse_pitch_sign(validation_batch)) if self.overlay_actions: @@ -1590,6 +1867,7 @@ def _run_validation_for_steps( overlay_videos=overlay_videos, overlay_captions=overlay_captions, ref_videos=ref_videos, + metadata=metadata, actions=actions, mouse_pitch_signs=mouse_pitch_signs, ) diff --git a/fastvideo/train/entrypoint/dcp_to_diffusers.py b/fastvideo/train/entrypoint/dcp_to_diffusers.py index da9d8e0b56..5485119024 100644 --- a/fastvideo/train/entrypoint/dcp_to_diffusers.py +++ b/fastvideo/train/entrypoint/dcp_to_diffusers.py @@ -62,6 +62,7 @@ def _save_role_pretrained( output_dir: str, module_names: list[str] | None = None, overwrite: bool = False, + link_base: bool = False, model: Any, ) -> str: """Export a role's modules into a diffusers-style model dir. @@ -117,16 +118,51 @@ def _copy_or_link(src: str, dest: str) -> None: logger.info( "Creating pretrained export dir at %s " - "(base=%s)", + "(base=%s, link_base=%s)", dst, local_base, + link_base, ) - shutil.copytree( - local_base, - dst, - symlinks=False, - copy_function=_copy_or_link, - ) + if link_base: + # Space-lean export: symlink every base component except the + # module dirs we are about to rewrite (those get a real dir with + # the base's non-weight files, e.g. config.json — weights and + # index are produced fresh below). Cross-user hardlinks are + # blocked by fs.protected_hardlinks, symlinks are not. + rewritten = set(module_names or ["transformer"]) + dst.mkdir(parents=True, exist_ok=True) + for entry in sorted(local_base.iterdir()): + if entry.name == ".cache" or entry.name.startswith(".git"): + continue + target = dst / entry.name + if entry.is_dir() and entry.name in rewritten: + target.mkdir() + for f in sorted(entry.iterdir()): + if f.name.endswith(".safetensors") or f.name.endswith(".safetensors.index.json"): + continue + shutil.copy2(os.path.realpath(f), target / f.name) + else: + os.symlink(os.path.realpath(entry), target) + else: + try: + shutil.copytree( + local_base, + dst, + symlinks=False, + copy_function=_copy_or_link, + # HF hub bookkeeping under the base checkpoint (.cache/) + # may be unreadable when the base belongs to another + # user; it is not part of the model. + ignore=shutil.ignore_patterns(".cache", ".git*"), + ) + except shutil.Error as exc: + # copytree collects per-file failures and raises at the end; + # tolerate stragglers as long as the component manifest made + # it. + logger.warning("copytree finished with %d skipped entries (first: %s)", len(exc.args[0]), + exc.args[0][0] if exc.args[0] else "?") + if not ((dst / "modular_model_index.json").is_file() or (dst / "model_index.json").is_file()): + raise FileNotFoundError(f"Export dir {dst} is missing its model index after copy.") _barrier() @@ -160,6 +196,10 @@ def _copy_or_link(src: str, dest: str) -> None: if _rank() == 0: for path in module_dir.glob("*.safetensors"): path.unlink(missing_ok=True) + # A leftover shard index from the base would point at the shards + # deleted above and shadow the fresh single-file export. + for path in module_dir.glob("*.safetensors.index.json"): + path.unlink(missing_ok=True) # Convert internal parameter names back to HF format. # load_model_from_full_model_state_dict builds reverse_param_names_mapping @@ -233,6 +273,8 @@ def convert( role: str = "student", overwrite: bool = False, verify: bool = False, + weights_only: bool = False, + link_base: bool = False, ) -> str: """Load a DCP checkpoint and export as a diffusers model. @@ -243,6 +285,8 @@ def convert( from fastvideo.distributed import ( maybe_init_distributed_environment_and_model_parallel, ) from fastvideo.train.utils.builder import build_from_config + from fastvideo.train.utils.instantiate import instantiate + from fastvideo.training.checkpointing_utils import ModelWrapper from fastvideo.train.utils.checkpoint import ( CheckpointManager, _resolve_resume_checkpoint, @@ -291,11 +335,23 @@ def convert( tc.distributed.hsdp_replicate_dim = 1 tc.distributed.hsdp_shard_dim = 1 - # -- Build model (loads pretrained weights + FSDP) -- - _, method, _, _ = build_from_config(cfg) - - # -- Load DCP weights into the model -- - states = method.checkpoint_state() + if weights_only: + # A role-only export must not construct unrelated roles or initialize + # the training dataloader. Large DMD2 checkpoints otherwise load the + # full teacher and critic merely to restore one student transformer, + # which can OOM and also makes export depend on stale dataset paths. + if role not in cfg.models: + raise KeyError(f"Role {role!r} is not present in the checkpoint config") + model = instantiate(cfg.models[role], training_config=tc) + if model.transformer is None: + raise ValueError(f"Role {role!r} has no transformer to export") + states = {f"roles.{role}.transformer": ModelWrapper(model.transformer)} + else: + # Full-state export retains the legacy behavior for callers that need + # method-managed optimizer or multi-role state. + _, method, _, _ = build_from_config(cfg) + states = method.checkpoint_state() + model = method._role_models[role] logger.info( "Loading DCP checkpoint from %s", resolved, @@ -303,7 +359,6 @@ def convert( dcp.load(states, checkpoint_id=str(dcp_dir)) # -- Export to diffusers format -- - model = method._role_models[role] base_model_path = str(tc.model_path) if not base_model_path: raise ValueError("Cannot determine base_model_path from " @@ -321,6 +376,7 @@ def convert( base_model_path=base_model_path, output_dir=output_dir, overwrite=overwrite, + link_base=link_base, model=model, ) logger.info("Export complete: %s", result) @@ -443,6 +499,22 @@ def main() -> None: "the exported directory to catch key-mapping bugs " "immediately."), ) + parser.add_argument( + "--weights-only", + action="store_true", + help=("Load only roles.* module weights from the checkpoint, " + "skipping optimizer/scheduler states (halves GPU memory " + "and allows exporting via a shim config whose optimizer " + "differs from the checkpoint's)."), + ) + parser.add_argument( + "--link-base", + action="store_true", + help=("Symlink base-model components into the export dir instead " + "of copying them (only the exported module dirs are real). " + "Saves hundreds of GB per export; the export then depends on " + "the base model dir staying in place."), + ) args = parser.parse_args(sys.argv[1:]) convert( @@ -452,6 +524,8 @@ def main() -> None: role=args.role, overwrite=args.overwrite, verify=args.verify, + weights_only=args.weights_only, + link_base=args.link_base, ) diff --git a/fastvideo/train/entrypoint/train.py b/fastvideo/train/entrypoint/train.py index 3f71e2a721..9d3f45ecf0 100644 --- a/fastvideo/train/entrypoint/train.py +++ b/fastvideo/train/entrypoint/train.py @@ -58,11 +58,19 @@ def run_training_from_config( # Auto-set attention backend for model families that require a specific # backend at load time, unless the user already overrode it explicitly. - if tc.vsa_sparsity > 0.0: + if tc.vsa_sparsity > 0.0 and "minimax" not in model_path_lower: os.environ.setdefault( "FASTVIDEO_ATTENTION_BACKEND", "VIDEO_SPARSE_ATTN", ) + # H3: do NOT push VSA into the env fallback. Training roles take their + # backend from per-role model config (student VSA, teacher/critic dense); + # the validation/inference pipeline resolves from this env and currently + # faults under the VSA-H3 kernel (async CUDA error surfacing as + # ncclUnhandledCudaError at the next FSDP all-gather, job 2307). + # Until the inference-side VSA path is debugged, H3 validation runs its + # native dense backend — a sparse-trained student evaluated dense is a + # known contract mismatch, noted in h3_dmd.md. elif ("turbodiffusion" in model_path_lower or "turbowan" in model_path_lower): os.environ.setdefault( "FASTVIDEO_ATTENTION_BACKEND", @@ -97,6 +105,11 @@ def run_training_from_config( ckpt_config = CheckpointConfig( save_steps=int(tc.checkpoint.training_state_checkpointing_steps or 0), keep_last=int(tc.checkpoint.checkpoints_total_limit or 0), + start_step=int(tc.checkpoint.checkpointing_start_step or 0), + save_inference_on_validation=bool(tc.checkpoint.save_inference_checkpoint_on_validation), + inference_role=str(tc.checkpoint.inference_checkpoint_role or "student"), + inference_dtype=str(tc.checkpoint.inference_checkpoint_dtype or "bfloat16"), + require_complete_training_checkpoint=bool(tc.checkpoint.require_complete_training_checkpoint), ) checkpoint_manager = CheckpointManager( diff --git a/fastvideo/train/methods/base.py b/fastvideo/train/methods/base.py index 5168313c5f..cbb45faa5b 100644 --- a/fastvideo/train/methods/base.py +++ b/fastvideo/train/methods/base.py @@ -109,6 +109,17 @@ def _optimizer_dict(self) -> dict[str, Any]: def _lr_scheduler_dict(self) -> dict[str, Any]: ... + def apply_configured_lrs(self) -> None: + """Re-apply the config's learning rates to live optimizers/schedulers. + + Called after a checkpoint resume when + ``training.checkpoint.reset_lr_on_resume`` is set: DCP restores + param-group ``lr``/``initial_lr`` and scheduler ``base_lrs`` from the + checkpoint, which would otherwise silently override an LR change in + the YAML. Methods with per-role LRs override this. + """ + return None + def checkpoint_state(self) -> dict[str, Any]: """Return DCP-ready checkpoint state for all trainable roles. @@ -144,6 +155,41 @@ def checkpoint_state(self) -> dict[str, Any]: return states + def inference_checkpoint_modules( + self, + role: str = "student", + ) -> dict[str, torch.nn.Module]: + """Return the complete modules that form an inference checkpoint. + + This intentionally does not reuse :meth:`checkpoint_state`: resumable + state includes optimizers and every trainable role, while inference + exports select one explicit role and must retain frozen parameters too. + Methods with a different deployable role contract may override this + hook; EMA is never selected implicitly. + """ + model = self._role_models.get(role) + if model is None: + raise ValueError(f"Inference checkpoint role {role!r} is not available; " + f"known roles: {sorted(self._role_models)}") + transformer = getattr(model, "transformer", None) + if not isinstance(transformer, torch.nn.Module): + raise ValueError(f"Inference checkpoint role {role!r} has no transformer module") + return {"transformer": transformer} + + def inference_checkpoint_base_model_path( + self, + role: str = "student", + ) -> str: + """Return the immutable model directory supplying non-trained parts.""" + model = self._role_models.get(role) + if model is None: + raise ValueError(f"Inference checkpoint role {role!r} is not available; " + f"known roles: {sorted(self._role_models)}") + path = str(getattr(model, "_init_from", "") or "") + if not path: + raise ValueError(f"Inference checkpoint role {role!r} does not expose its init_from model path") + return path + def backward( self, loss_map: dict[str, torch.Tensor], @@ -182,7 +228,7 @@ def seed_optimizer_state_for_resume(self) -> None: DCP needs matching entries to load into; without them the saved optimizer state is silently dropped. """ - for opt in self.get_optimizers(0): + for opt in self._optimizer_dict.values(): for group in opt.param_groups: for p in group["params"]: if not p.requires_grad: @@ -275,6 +321,6 @@ def on_train_start(self) -> None: def _infer_attn_kind(self) -> Literal["dense", "vsa"]: """Derive metadata mode from the student's resolved backend.""" backend = (self.student.attention_backend_name or envs.FASTVIDEO_ATTENTION_BACKEND) - if backend == "VIDEO_SPARSE_ATTN": + if backend in ("VIDEO_SPARSE_ATTN", "VIDEO_SPARSE_ATTN_H3"): return "vsa" return "dense" diff --git a/fastvideo/train/methods/distribution_matching/dmd2.py b/fastvideo/train/methods/distribution_matching/dmd2.py index ee1890bcb9..925b80f28b 100644 --- a/fastvideo/train/methods/distribution_matching/dmd2.py +++ b/fastvideo/train/methods/distribution_matching/dmd2.py @@ -3,9 +3,13 @@ from __future__ import annotations +import json +import math +from pathlib import Path from typing import Any, Literal import torch +import torch.distributed as dist import torch.nn.functional as F from fastvideo.train.methods.base import TrainingMethod, LogScalar @@ -55,6 +59,13 @@ def __init__( raise ValueError("DMD2Method requires critic to be trainable") self._cfg_uncond = self._parse_cfg_uncond() self._rollout_mode = self._parse_rollout_mode() + ( + self._rollout_carry, + self._rollout_carry_slot_count, + self._rollout_sample_type, + ) = self._parse_rollout_carry() + self._init_rollout_carry_state() + self._rollout_data_forcing = self._parse_rollout_data_forcing() self._validate_preprocessed_data_type() self._configure_student_negative_conditioning() self._denoising_step_list: torch.Tensor | None = (None) @@ -62,6 +73,10 @@ def __init__( self._score_min_timestep, self._score_max_timestep, ) = self._parse_score_timestep_bounds() + self._score_timestep_shift = self._parse_score_timestep_shift() + self._score_timestep_warp_max = self._parse_score_timestep_warp_max() + self._score_timestep_continuous = self._parse_score_timestep_continuous() + self._fake_score_loss_space = self._parse_fake_score_loss_space() # Initialize preprocessors on student. self.student.init_preprocessors(self.training_config) @@ -92,6 +107,9 @@ def single_train_step( dict[str, Any], dict[str, LogScalar], ]: + if self._rollout_carry: + return self._carried_train_step(batch, iteration) + latents_source: Literal["data", "zeros"] = "data" if self._rollout_mode == "simulate": latents_source = "zeros" @@ -107,22 +125,29 @@ def single_train_step( generator_loss = torch.zeros( (), device=training_batch.latents.device, - dtype=training_batch.latents.dtype, + dtype=torch.float32, ) student_ctx = None + generator_metrics: dict[str, LogScalar] = {} + fake_score_loss = torch.zeros_like(generator_loss) + critic_ctx = None + critic_outputs: dict[str, Any] = {} + critic_metrics: dict[str, LogScalar] = {} if update_student: generator_pred_x0 = self._student_rollout(training_batch, with_grad=True) student_ctx = ( training_batch.timesteps, training_batch.attn_metadata_vsa, ) - generator_loss = self._dmd_loss(generator_pred_x0, training_batch) - - ( - fake_score_loss, - critic_ctx, - critic_outputs, - ) = self._critic_flow_matching_loss(training_batch) + generator_loss, generator_metrics = self._dmd_loss(generator_pred_x0, training_batch) + training_batch.dmd_latent_vis_dict["generator_pred_video"] = generator_pred_x0.detach() + else: + ( + fake_score_loss, + critic_ctx, + critic_outputs, + critic_metrics, + ) = self._critic_flow_matching_loss(training_batch) total_loss = generator_loss + fake_score_loss loss_map = { @@ -137,7 +162,18 @@ def single_train_step( "student_ctx": student_ctx, "critic_ctx": critic_ctx, } - metrics: dict[str, LogScalar] = {"update_student": float(update_student)} + metrics: dict[str, LogScalar] = { + "update_student": float(update_student), + **generator_metrics, + **critic_metrics, + } + # Rank-local latent snapshots for LatentVisCallback. + self.latent_vis = { + **(training_batch.fake_score_latent_vis_dict or {}), + **(training_batch.dmd_latent_vis_dict or {}), + "_fv_latent_layout": + getattr(training_batch, "minimax_h3_dmd_layout", None), + } return loss_map, outputs, metrics # TrainingMethod override: backward @@ -168,6 +204,8 @@ def backward( student_ctx, grad_accum_rounds=grad_accum_rounds, ) + self._assert_finite_gradients("student", self.student) + return critic_ctx = backward_ctx.get("critic_ctx") if critic_ctx is None: @@ -177,39 +215,114 @@ def backward( critic_ctx, grad_accum_rounds=grad_accum_rounds, ) + self._assert_finite_gradients("critic", self.critic) + + @staticmethod + def _local_tensor(tensor: torch.Tensor) -> torch.Tensor: + return getattr(tensor, "_local_tensor", tensor) + + @classmethod + def _assert_finite_gradients(cls, role: str, model: ModelBase) -> None: + """Abort on numerical corruption instead of applying a partial update.""" + bad: list[str] = [] + for name, parameter in model.transformer.named_parameters(): + if parameter.grad is None: + continue + if not bool(torch.isfinite(cls._local_tensor(parameter.grad)).all()): + bad.append(name) + if len(bad) == 8: + break + if bad: + raise RuntimeError( + f"Nonfinite {role} gradients before clipping/Adam: {bad}" + ) # TrainingMethod override: get_optimizers def get_optimizers( self, iteration: int, ) -> list[torch.optim.Optimizer]: - optimizers: list[torch.optim.Optimizer] = [] - optimizers.append(self._critic_optimizer) if self._should_update_student(iteration): - optimizers.append(self._student_optimizer) - return optimizers + return [self._student_optimizer] + return [self._critic_optimizer] # TrainingMethod override: get_lr_schedulers def get_lr_schedulers( self, iteration: int, ) -> list[Any]: - schedulers: list[Any] = [] - schedulers.append(self._critic_lr_scheduler) if self._should_update_student(iteration): - schedulers.append(self._student_lr_scheduler) - return schedulers + return [self._student_lr_scheduler] + return [self._critic_lr_scheduler] # TrainingMethod override: get_grad_clip_targets def get_grad_clip_targets( self, iteration: int, ) -> dict[str, torch.nn.Module]: - targets: dict[str, torch.nn.Module] = {} if self._should_update_student(iteration): - targets["student"] = (self.student.transformer) - targets["critic"] = self.critic.transformer - return targets + return {"student": self.student.transformer} + return {"critic": self.critic.transformer} + + def optimizers_schedulers_step(self, iteration: int) -> None: + """Prove the first critic and student Adam updates are finite FP32.""" + role = "student" if self._should_update_student(iteration) else "critic" + model = self.student if role == "student" else self.critic + optimizer = self.get_optimizers(iteration)[0] + verified = getattr(self, "_verified_optimizer_roles", set()) + first = role not in verified + probes: list[tuple[torch.Tensor, torch.Tensor]] = [] + if first: + for parameter in model.transformer.parameters(): + if not parameter.requires_grad: + continue + local = self._local_tensor(parameter).detach().reshape(-1) + if local.dtype != torch.float32: + raise RuntimeError( + f"DMD2 {role} master weights must be FP32, got {local.dtype}" + ) + if local.numel() and len(probes) < 16: + probes.append((local, local[:4096].clone())) + + super().optimizers_schedulers_step(iteration) + + if first: + changed = sum( + int(torch.count_nonzero(current[:before.numel()] != before)) + for current, before in probes + ) + moments = [ + self._local_tensor(value) + for state in optimizer.state.values() + for key, value in state.items() + if key in {"exp_avg", "exp_avg_sq"} and torch.is_tensor(value) + ] + if ( + not probes + or changed == 0 + or not moments + or any(value.dtype != torch.float32 for value in moments) + or any(not bool(torch.isfinite(value).all()) for value in moments) + or any( + not bool(torch.isfinite(current[:before.numel()]).all()) + for current, before in probes + ) + ): + raise RuntimeError(f"No finite FP32 Adam update for {role}") + rank = dist.get_rank() if dist.is_initialized() else 0 + root = Path(self.training_config.checkpoint.output_dir) + root.mkdir(parents=True, exist_ok=True) + (root / f"dmd2_update_{role}_rank{rank}.json").write_text( + json.dumps({ + "role": role, + "iteration": iteration, + "changed_probe_elements": changed, + "passed": True, + }) + "\n", + encoding="utf-8", + ) + verified.add(role) + self._verified_optimizer_roles = verified def _parse_rollout_mode(self, ) -> Literal["simulate", "data_latent"]: """Parse how DMD2 obtains the latent point used for rollout. @@ -236,6 +349,224 @@ def _parse_rollout_mode(self, ) -> Literal["simulate", "data_latent"]: "{simulate, data_latent}, got " f"{raw!r}") + def _parse_rollout_carry(self) -> tuple[bool, int, Literal["ode", "sde"]]: + """Parse the carried backward-simulation knobs. + + ``rollout_carry: true`` walks the student's own sampling grid one + rung per ``single_train_step`` call, carrying the trajectory in + memory across calls (FastGen's backward simulation): exactly one + generation forward per call instead of one full-grid walk. Off + (default) keeps the existing full-rollout behavior unchanged. + + ``rollout_carry_slots`` is the number of independent trajectory + streams per rank; it must equal + ``training.loop.gradient_accumulation_steps`` because the trainer + calls ``single_train_step`` once per accumulation round and the + slots are selected round-robin over calls. Defaults to that value. + + ``rollout_sample_type`` picks how the walk re-noises onto the next + rung: ``sde`` draws fresh noise (the existing rollout behavior), + ``ode`` reuses the noise the current state implies per modality — + the deterministic step the FastGen H3 recipe uses. + """ + raw_carry = self.method_config.get("rollout_carry", None) + if raw_carry is None: + raw_carry = False + if not isinstance(raw_carry, bool): + raise ValueError("method.rollout_carry must be a bool, got " + f"{type(raw_carry).__name__}") + carry = bool(raw_carry) + + raw_sample_type = self.method_config.get("rollout_sample_type", None) + sample_type: Literal["ode", "sde"] = "sde" + if raw_sample_type is not None: + if not isinstance(raw_sample_type, str): + raise ValueError("method.rollout_sample_type must be a " + "string, got " + f"{type(raw_sample_type).__name__}") + normalized = raw_sample_type.strip().lower() + if normalized not in ("ode", "sde"): + raise ValueError("method.rollout_sample_type must be one of " + f"{{ode, sde}}, got {raw_sample_type!r}") + if not carry: + raise ValueError("method.rollout_sample_type requires " + "method.rollout_carry: true") + sample_type = normalized # type: ignore[assignment] + + slots_raw = get_optional_int( + self.method_config, + "rollout_carry_slots", + where="method.rollout_carry_slots", + ) + if not carry: + if slots_raw is not None: + raise ValueError("method.rollout_carry_slots requires " + "method.rollout_carry: true") + return False, 0, sample_type + + if self._rollout_mode != "simulate": + raise ValueError("method.rollout_carry: true requires " + "method.rollout_mode: simulate") + + grad_accum = max( + 1, + int(self.training_config.loop.gradient_accumulation_steps or 1), + ) + slots = grad_accum if slots_raw is None else int(slots_raw) + if slots <= 0: + raise ValueError("method.rollout_carry_slots must be positive, " + f"got {slots}") + if slots != grad_accum: + # The trainer calls single_train_step once per accumulation round + # without passing the round index; the round-robin slot selection + # only matches the trainer's cadence when the counts agree. + raise ValueError("method.rollout_carry_slots must equal " + "training.loop.gradient_accumulation_steps, got " + f"slots={slots} vs " + f"gradient_accumulation_steps={grad_accum}") + + if sample_type == "ode" and not callable(getattr(self.student, "extract_eps", None)): + raise ValueError("method.rollout_sample_type: ode requires the " + "student model to implement " + "extract_eps(noisy_latents, clean_latents, " + "timestep)") + + _, stagger_groups = self._rollout_carry_rank_world() + if bool(getattr(getattr(self.training_config, "data", None), "native_shape_bucketing", False)): + stagger_groups = 1 + self._validate_rollout_carry_coverage( + streams=stagger_groups * slots, + grid_len=self._rollout_grid_length(), + interval=self._generator_update_interval(), + ) + return True, slots, sample_type + + @staticmethod + def _validate_rollout_carry_coverage( + *, + streams: int, + grid_len: int, + interval: int, + ) -> None: + """Reject configurations that leave student rungs untrained. + + Student updates revisit a stream's rung modulo + ``gcd(len(dmd_denoising_steps), generator_update_interval)``, so + every residue class must be represented by a (rank, slot) stream; + the consecutive stagger offsets cover all classes exactly when + there are at least ``gcd`` streams. + """ + phase_classes = math.gcd(grid_len, interval) + if streams < phase_classes: + raise ValueError("Carried backward-simulation DMD2 cannot cover every student " + f"rung with len(dmd_denoising_steps)={grid_len}, " + f"generator_update_interval={interval}, and " + f"{streams} trajectory stream(s) (stagger groups x slots). " + "Student updates preserve the rung modulo " + f"gcd(grid, interval)={phase_classes}, but only {streams} " + "stream phase(s) are present. Use at least that many " + "rank-slot streams or choose a coprime update interval.") + + def _rollout_carry_rank_world(self) -> tuple[int, int]: + """Stagger rank and stream-group count for the carried rollout. + + Mirrors ``TrainingMethod.on_train_start``'s RNG grouping: ranks + inside one sequence-parallel group shard the same document and must + walk one shared trajectory, so they share a stagger rank (with + ``sp_size=1`` this is exactly the global rank). Falls back to a + single group when the distributed world is not initialized (CPU + tests, single-process runs). + """ + try: + from fastvideo.distributed import get_world_group + world_group = get_world_group() + global_rank = int(world_group.rank) + world_size = int(world_group.world_size) + except (AssertionError, ImportError, RuntimeError): + global_rank, world_size = 0, 1 + sp_size = max( + 1, + int(getattr(self.training_config.distributed, "sp_size", 1) or 1), + ) + return global_rank // sp_size, max(1, world_size // sp_size) + + def _parse_rollout_data_forcing(self) -> bool: + """Parse per-batch data forcing for the carried walk. + + ``rollout_data_forcing: true`` routes latent-bearing batches (t2va + parquet rows in a mixed ``data_path``) onto FastGen's data-driven + student inputs — the real packed latents forward-noised at a + uniformly drawn grid rung (``sample_from_t_list`` semantics) — + while text-only batches keep walking the carried backward + simulation. Off (default) keeps every batch on the walk, + byte-identical to the carry-only behavior. + """ + raw = self.method_config.get("rollout_data_forcing", None) + if raw is None: + return False + if not isinstance(raw, bool): + raise ValueError("method.rollout_data_forcing must be a bool, " + f"got {type(raw).__name__}") + if raw and not self._rollout_carry: + raise ValueError("method.rollout_data_forcing: true requires " + "method.rollout_carry: true; without the carry, " + "always-forced inputs are " + "method.rollout_mode: data_latent") + if raw: + allow_mixed = self.method_config.get( + "allow_mixed_rollout_regimes", + False, + ) + if not isinstance(allow_mixed, bool): + raise ValueError("method.allow_mixed_rollout_regimes must be a bool, " + f"got {type(allow_mixed).__name__}") + if not allow_mixed: + raise ValueError("method.rollout_data_forcing mixes carried and data-latent " + "rollout regimes per batch, which is not a FastGen recipe. " + "Choose one global regime with rollout_mode={simulate, " + "data_latent}; set allow_mixed_rollout_regimes: true only " + "to reproduce the legacy v9 experiment.") + return raw + + @staticmethod + def _batch_has_latents(batch: dict[str, Any]) -> bool: + """Classify a mixed-loading batch as latent-bearing or text-only. + + Under the t2va parquet schema the collate emits empty (numel-0) + tensors for latent columns a text-only row does not carry, so + presence means "key exists and non-empty". A row carrying exactly + one of the pair is corrupt data, not a batch type. + """ + video = batch.get("vae_latent") + audio = batch.get("audio_latent") + has_video = isinstance(video, torch.Tensor) and video.numel() > 0 + has_audio = isinstance(audio, torch.Tensor) and audio.numel() > 0 + if has_video != has_audio: + raise ValueError("Mixed-loading batch carries exactly one of " + "vae_latent/audio_latent non-empty; a t2va row " + "must carry both and a text_only row neither " + f"(vae_latent={'present' if has_video else 'empty/missing'}, " + f"audio_latent={'present' if has_audio else 'empty/missing'})") + return has_video + + def _rollout_grid_length(self) -> int: + raw = self.method_config.get("dmd_denoising_steps", None) + if not isinstance(raw, list) or not raw: + raise ValueError("method_config.dmd_denoising_steps must " + "be set for DMD2 distillation") + return len(raw) + + def _init_rollout_carry_state(self) -> None: + # Transient in-memory trajectory state: one independent slot per + # gradient-accumulation round, selected round-robin over calls. + # Intentionally never checkpointed (mirrors FastGen's CarryCallback): + # on resume every slot restarts from fresh noise — a brief warmup + # until the walk is mid-trajectory again. + slots = max(0, int(self._rollout_carry_slot_count)) + self._carry_call_count = 0 + self._carry_slots: list[dict[str, Any] | None] = [None] * slots + self._carry_slot_seeded: list[bool] = [False] * slots + def _validate_preprocessed_data_type(self) -> None: data_type = str(getattr( self.training_config.data, @@ -246,6 +577,13 @@ def _validate_preprocessed_data_type(self) -> None: raise ValueError("training.data.preprocessed_data_type='text_only' " "requires method.rollout_mode='simulate'; " "data_latent rollout requires vae_latent data.") + if self._rollout_data_forcing and data_type != "t2va": + raise ValueError("method.rollout_data_forcing: true requires " + "training.data.preprocessed_data_type='t2va': the " + "t2va parquet schema is the superset that reads " + "latent columns; text-only roots mixed into the " + "same data_path yield empty latent columns and " + "route to the carried walk.") def _uses_negative_prompt_conditioning(self) -> bool: if self._cfg_uncond is None: @@ -370,20 +708,87 @@ def _init_optimizers_and_schedulers(self) -> None: scheduler_name=critic_sched, ) - def _should_update_student( - self, - iteration: int, - ) -> bool: + @staticmethod + def _add_noise_for_batch( + model: ModelBase, + clean_latents: torch.Tensor, + noise: torch.Tensor, + timestep: torch.Tensor, + batch: Any, + ) -> torch.Tensor: + """Call the batch-aware hook while retaining lightweight test doubles.""" + hook = getattr(model, "add_noise_for_batch", None) + if hook is not None and batch is not None: + return hook(clean_latents, noise, timestep, batch) + return model.add_noise(clean_latents, noise, timestep) + + @staticmethod + def _extract_eps_for_batch( + model: ModelBase, + noisy_latents: torch.Tensor, + clean_latents: torch.Tensor, + timestep: torch.Tensor, + batch: Any, + ) -> torch.Tensor: + hook = getattr(model, "extract_eps_for_batch", None) + if hook is not None and batch is not None: + return hook(noisy_latents, clean_latents, timestep, batch) + extractor = getattr(model, "extract_eps", None) + if not callable(extractor): + raise TypeError(f"{type(model).__name__} does not implement extract_eps") + return extractor(noisy_latents, clean_latents, timestep) + + def _modality_slices(self, batch: Any) -> tuple[tuple[str, slice], ...] | None: + """Return slices used to normalize packed modalities independently.""" + batch_getter = getattr(self.student, "modality_slices_for_batch", None) + if batch_getter is not None: + slices = tuple(batch_getter(batch)) + return slices or None + getter = getattr(self.student, "modality_slices", None) + if getter is None: + return None + slices = tuple(getter()) + return slices or None + + def _modality_weight(self, name: str) -> float: + raw = self.method_config.get("modality_loss_weights", None) + if not isinstance(raw, dict): + return 1.0 + value = raw.get(name, 1.0) + return 1.0 if value is None else float(value) + + def apply_configured_lrs(self) -> None: + """Force student/critic LRs back to the configured values (post-resume).""" + student_lr = float(self.training_config.optimizer.learning_rate) + critic_lr = float(self.method_config.get("fake_score_learning_rate")) + for optimizer, scheduler, lr in ( + (self._student_optimizer, self._student_lr_scheduler, student_lr), + (self._critic_optimizer, self._critic_lr_scheduler, critic_lr), + ): + for group in optimizer.param_groups: + group["lr"] = lr + if "initial_lr" in group: + group["initial_lr"] = lr + if hasattr(scheduler, "base_lrs"): + scheduler.base_lrs = [lr] * len(scheduler.base_lrs) + + def _generator_update_interval(self) -> int: interval = get_optional_int( self.method_config, "generator_update_interval", where="method.generator_update_interval", ) if interval is None: - interval = 1 + interval = 5 if interval <= 0: - return True - return iteration % interval == 0 + raise ValueError("method.generator_update_interval must be positive") + return interval + + def _should_update_student( + self, + iteration: int, + ) -> bool: + return iteration % self._generator_update_interval() == 0 def _get_denoising_step_list( self, @@ -432,7 +837,7 @@ def _sample_rollout_timestep( return step_list[index] def _parse_score_timestep_bounds(self) -> tuple[int, int]: - """Resolve the score-model timestep window used by legacy DMD. + """Resolve the score-model timestep window. The student rollout schedule is controlled separately by ``dmd_denoising_steps``. These bounds apply only to the randomly @@ -461,15 +866,134 @@ def _parse_score_timestep_bounds(self) -> tuple[int, int]: int(max_ratio * num_timesteps), ) - def _sample_score_timestep(self, device: torch.device) -> torch.Tensor: - timestep = torch.randint( - 0, - int(self.student.num_train_timesteps), - [1], - device=device, - dtype=torch.long, - generator=self.cuda_generator, + def _parse_score_timestep_shift(self) -> float: + """Resolve the rational warp used by score-time sampling. + + Legacy integer sampling draws uniformly in the warped coordinate and + inverts it. Continuous FastGen parity draws the pre-warp coordinate + directly and applies this inverse warp before the model adapter adds + its modality-specific clock. + """ + shift = get_optional_float( + self.method_config, + "score_timestep_shift", + where="method.score_timestep_shift", + ) + shift = 1.0 if shift is None else float(shift) + if shift <= 0.0: + raise ValueError("method.score_timestep_shift must be > 0, " + f"got {shift}") + return shift + + def _parse_score_timestep_warp_max(self) -> float: + """Resolve the endpoint used by the continuous rational time warp.""" + warp_max = get_optional_float( + self.method_config, + "score_timestep_warp_max", + where="method.score_timestep_warp_max", ) + warp_max = 1.0 if warp_max is None else float(warp_max) + if not 0.0 < warp_max <= 1.0: + raise ValueError("method.score_timestep_warp_max must satisfy " + f"0 < max <= 1, got {warp_max}") + max_ratio = self._score_max_timestep / float(self.student.num_train_timesteps) + if max_ratio > warp_max: + raise ValueError("method.max_timestep_ratio must not exceed " + "method.score_timestep_warp_max, got " + f"{max_ratio} > {warp_max}") + return warp_max + + def _parse_score_timestep_continuous(self) -> bool: + """Select FastGen-style continuous score times instead of integer bins.""" + raw = self.method_config.get("score_timestep_continuous", False) + if not isinstance(raw, bool): + raise ValueError("method.score_timestep_continuous must be a bool, " + f"got {type(raw).__name__}") + return raw + + def _parse_fake_score_loss_space(self) -> dict[str, str]: + """Resolve the critic regression space, globally or per modality. + + ``velocity`` is plain velocity MSE. A global ``x0`` setting calls the + critic's x0 prediction directly, matching FastGen without estimating + sigma from rounded latents. Legacy mixed mappings retain the original + single-forward sigma-squared conversion for their x0 modalities. + """ + raw = self.method_config.get("fake_score_loss_space", None) + if raw is None: + return {"__default__": "velocity"} + if isinstance(raw, str): + mapping = {"__default__": raw} + elif isinstance(raw, dict): + mapping = {str(k).strip().lower(): str(v) for k, v in raw.items()} + mapping.setdefault("__default__", "velocity") + else: + raise ValueError("method.fake_score_loss_space must be a string " + "or a {modality: space} mapping, got " + f"{type(raw).__name__}") + normalized: dict[str, str] = {} + for key, value in mapping.items(): + space = str(value).strip().lower() + if space not in ("velocity", "x0"): + raise ValueError("method.fake_score_loss_space values must be " + f"one of {{velocity, x0}}, got {value!r} " + f"for {key!r}") + normalized[key] = space + return normalized + + def _fake_score_space_for(self, modality_name: str) -> str: + return self._fake_score_loss_space.get( + modality_name, + self._fake_score_loss_space["__default__"], + ) + + def _sample_score_timestep(self, device: torch.device) -> torch.Tensor: + shift = self._score_timestep_shift + num_timesteps = float(self.student.num_train_timesteps) + t_lo = self._score_min_timestep / num_timesteps + t_hi = self._score_max_timestep / num_timesteps + + if getattr(self, "_score_timestep_continuous", False): + # FastGen draws the pre-warp coordinate continuously in float64, + # then applies the rational shift. This method stores the inverse + # shift because model adapters (H3: video 12, audio 3) apply their + # own modality clocks afterwards. Bounds therefore belong to U, + # not to the inverse-warped base time. + u = torch.rand( + [1], + device=device, + dtype=torch.float64, + generator=self.cuda_generator, + ) * (t_hi - t_lo) + t_lo + inverse_shift = 1.0 / shift + warp_max = getattr(self, "_score_timestep_warp_max", 1.0) + t = (u * inverse_shift * warp_max / (u * (inverse_shift - 1.0) + warp_max)) + timestep = t * num_timesteps + timestep = self.student.shift_and_clamp_timestep(timestep) + return timestep.clamp(0.0, warp_max * num_timesteps) + + if shift == 1.0: + # Draw inside the bounds directly; drawing over the full range + # and clamping piles probability atoms onto both endpoints. + timestep = torch.randint( + self._score_min_timestep, + self._score_max_timestep + 1, + [1], + device=device, + dtype=torch.long, + generator=self.cuda_generator, + ) + else: + sigma_lo = shift * t_lo / (1.0 + (shift - 1.0) * t_lo) + sigma_hi = shift * t_hi / (1.0 + (shift - 1.0) * t_hi) + u = torch.rand( + [1], + device=device, + dtype=torch.float32, + generator=self.cuda_generator, + ) * (sigma_hi - sigma_lo) + sigma_lo + t = u / (shift - (shift - 1.0) * u) + timestep = (t * num_timesteps).round().to(torch.long) timestep = self.student.shift_and_clamp_timestep(timestep) return timestep.clamp( self._score_min_timestep, @@ -495,7 +1019,7 @@ def _student_rollout( dtype=dtype, generator=self.cuda_generator, ) - noisy_latents = self.student.add_noise(latents, noise, timestep) + noisy_latents = self._add_noise_for_batch(self.student, latents, noise, timestep, batch) pred_x0 = self.student.predict_x0( noisy_latents, timestep, @@ -561,11 +1085,13 @@ def _student_rollout( dtype=pred_clean.dtype, generator=self.cuda_generator, ) - current_noise_latents = (self.student.add_noise( + current_noise_latents = self._add_noise_for_batch( + self.student, pred_clean, noise, next_timestep_tensor, - )) + batch, + ) noise_latents.append(current_noise_latents.clone()) if noise_latent_index >= 0: @@ -598,12 +1124,477 @@ def _student_rollout( batch.dmd_latent_vis_dict["generator_timestep"] = target_timestep.float().detach() return pred_x0 + # ------------------------------------------------------------------ + # Carried backward simulation — the student's own trajectory, walked + # one rung per single_train_step call (port of FastGen's + # _backward_simulation / _staggered_start / _advance_carry). + # ------------------------------------------------------------------ + + def _carried_train_step( + self, + batch: dict[str, Any], + iteration: int, + ) -> tuple[ + dict[str, torch.Tensor], + dict[str, Any], + dict[str, LogScalar], + ]: + """One backward-simulation call: one generation forward, carried state. + + A multistep student is only ever correct on its own sampling + trajectory, and walking the full grid every call costs + ``len(dmd_denoising_steps)`` forwards. Instead the walk is spread + over consecutive calls: each call pays for exactly one student + forward at the carried rung, both phases (student and critic) + consume it — the critic is fit on the same simulated states the + student trains on — and both advance the trajectory. Each + grad-accum round owns an independent slot, selected round-robin + because the trainer does not pass the round index. + + An empty slot (first ever use, cleared after a finished trajectory, + or after a resume — the carry is transient and never checkpointed) + starts a fresh trajectory from noise and adopts the incoming loader + batch's conditioning; mid-walk calls ignore the fresh loader batch + and rebuild the training batch from the carried raw batch, since a + trajectory keeps the prompt it set out with. + """ + slot = self._carry_call_count % self._rollout_carry_slot_count + self._carry_call_count += 1 + + # Per-batch routing (off unless method.rollout_data_forcing): a batch + # that carries real latents trains on them at a noised grid rung and + # leaves this slot's walk untouched. + if self._rollout_data_forcing and self._batch_has_latents(batch): + return self._data_forced_train_step(batch, slot, iteration) + + carried = self._carry_slots[slot] + raw_batch = (self._carry_snapshot_raw_batch(batch) if carried is None else carried["raw_batch"]) + + training_batch = self.student.prepare_batch( + raw_batch, + generator=self.cuda_generator, + latents_source="zeros", + ) + latents = training_batch.latents + device = latents.device + step_list = self._get_denoising_step_list(device) + + if carried is None: + rung = 0 + state = torch.randn( + latents.shape, + device=device, + dtype=latents.dtype, + generator=self.cuda_generator, + ) + # Stagger only the first-ever fill of each slot; later fresh + # starts begin at rung 0 with no pre-walk and stay out of phase + # naturally. + if not self._carry_slot_seeded[slot]: + self._carry_slot_seeded[slot] = True + state, rung = self._staggered_start( + state, + training_batch, + step_list, + slot, + ) + else: + rung = int(carried["rung"]) + state = carried["state"] + + timestep = step_list[rung] * torch.ones( + 1, + device=device, + dtype=torch.long, + ) + + update_student = self._should_update_student(iteration) + + generator_loss = torch.zeros((), device=device, dtype=torch.float32) + fake_score_loss = torch.zeros_like(generator_loss) + student_ctx = None + critic_ctx = None + critic_outputs: dict[str, Any] = {} + generator_metrics: dict[str, LogScalar] = {} + critic_metrics: dict[str, LogScalar] = {} + if update_student: + generator_pred_x0 = self.student.predict_x0( + state, + timestep, + training_batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="vsa", + ) + student_ctx = ( + training_batch.timesteps, + training_batch.attn_metadata_vsa, + ) + generator_loss, generator_metrics = self._dmd_loss(generator_pred_x0, training_batch) + training_batch.dmd_latent_vis_dict["generator_pred_video"] = generator_pred_x0.detach() + else: + with torch.no_grad(): + generator_pred_x0 = self.student.predict_x0( + state, + timestep, + training_batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="vsa", + ) + ( + fake_score_loss, + critic_ctx, + critic_outputs, + critic_metrics, + ) = self._critic_flow_matching_loss( + training_batch, + generator_pred_x0=generator_pred_x0, + ) + training_batch.dmd_latent_vis_dict["generator_timestep"] = timestep.float().detach() + + # Advance after the loss path: both phases generated, so both hand + # the trajectory on. + self._advance_carry( + slot, + state, + generator_pred_x0, + timestep, + rung, + step_list, + training_batch, + raw_batch, + ) + + total_loss = generator_loss + fake_score_loss + loss_map = { + "total_loss": total_loss, + "generator_loss": generator_loss, + "fake_score_loss": fake_score_loss, + } + outputs: dict[str, Any] = dict(critic_outputs) + outputs["_fv_backward"] = { + "update_student": update_student, + "student_ctx": student_ctx, + "critic_ctx": critic_ctx, + } + metrics: dict[str, LogScalar] = { + "update_student": float(update_student), + "rollout_step": float(rung), + **generator_metrics, + **critic_metrics, + } + if self._rollout_data_forcing: + # The running mean of this metric is the realized latent-row + # fraction of the mix; emitted only when routing is enabled so + # carry-only runs keep their exact metric set. + metrics["data_forced"] = 0.0 + # Rank-local latent snapshots for LatentVisCallback. + self.latent_vis = { + **(training_batch.fake_score_latent_vis_dict or {}), + **(training_batch.dmd_latent_vis_dict or {}), + "_fv_latent_layout": + getattr(training_batch, "minimax_h3_dmd_layout", None), + } + return loss_map, outputs, metrics + + def _data_forced_train_step( + self, + batch: dict[str, Any], + slot: int, + iteration: int, + ) -> tuple[ + dict[str, torch.Tensor], + dict[str, Any], + dict[str, LogScalar], + ]: + """One data-forced call: train on real latents noised at a grid rung. + + FastGen's data-driven multistep student inputs (its + ``backward_simulation: false`` regime): ``t_student`` is drawn + uniformly over the student grid's rungs — ``sample_from_t_list`` + semantics, never t=0 — and the real packed latents are + forward-noised to that rung under each modality's shift, exactly + the uncarried ``rollout_mode: data_latent`` math. The slot's + carried walk pauses untouched and resumes on this stream's next + text-only batch: FastGen picks one regime per config, so pausing + is the minimal per-batch composition of its two modes. Both + phases consume the same forced generation, mirroring the carried + step's critic passthrough. + + The slot's one-time stagger pre-walk still runs on its first-ever + call even when that call is data-forced: the pre-walk's FSDP + collective count must stay uniform across ranks, and ranks whose + first batch is text-only run theirs on this same call. The seeded + walk adopts this batch's conditioning and waits at its stagger + rung. + """ + training_batch = self.student.prepare_batch( + batch, + generator=self.cuda_generator, + latents_source="data", + ) + latents = training_batch.latents + device = latents.device + if not self._carry_slot_seeded[slot]: + self._carry_slot_seeded[slot] = True + step_list = self._get_denoising_step_list(device) + state = torch.randn( + latents.shape, + device=device, + dtype=latents.dtype, + generator=self.cuda_generator, + ) + state, rung = self._staggered_start( + state, + training_batch, + step_list, + slot, + ) + self._carry_slots[slot] = { + "state": state.detach(), + "rung": rung, + "raw_batch": self._carry_snapshot_raw_batch(batch), + } + + forced_timestep = self._sample_rollout_timestep(device) + noise = torch.randn( + latents.shape, + device=device, + dtype=latents.dtype, + generator=self.cuda_generator, + ) + noisy_latents = self._add_noise_for_batch(self.student, latents, noise, forced_timestep, training_batch) + + update_student = self._should_update_student(iteration) + + generator_loss = torch.zeros((), device=device, dtype=torch.float32) + fake_score_loss = torch.zeros_like(generator_loss) + student_ctx = None + critic_ctx = None + critic_outputs: dict[str, Any] = {} + generator_metrics: dict[str, LogScalar] = {} + critic_metrics: dict[str, LogScalar] = {} + if update_student: + generator_pred_x0 = self.student.predict_x0( + noisy_latents, + forced_timestep, + training_batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="vsa", + ) + student_ctx = ( + training_batch.timesteps, + training_batch.attn_metadata_vsa, + ) + generator_loss, generator_metrics = self._dmd_loss(generator_pred_x0, training_batch) + training_batch.dmd_latent_vis_dict["generator_pred_video"] = generator_pred_x0.detach() + else: + with torch.no_grad(): + generator_pred_x0 = self.student.predict_x0( + noisy_latents, + forced_timestep, + training_batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="vsa", + ) + ( + fake_score_loss, + critic_ctx, + critic_outputs, + critic_metrics, + ) = self._critic_flow_matching_loss( + training_batch, + generator_pred_x0=generator_pred_x0, + ) + training_batch.dmd_latent_vis_dict["generator_timestep"] = forced_timestep.float().detach() + + total_loss = generator_loss + fake_score_loss + loss_map = { + "total_loss": total_loss, + "generator_loss": generator_loss, + "fake_score_loss": fake_score_loss, + } + outputs: dict[str, Any] = dict(critic_outputs) + outputs["_fv_backward"] = { + "update_student": update_student, + "student_ctx": student_ctx, + "critic_ctx": critic_ctx, + } + metrics: dict[str, LogScalar] = { + "update_student": float(update_student), + "data_forced": 1.0, + **generator_metrics, + **critic_metrics, + } + # Rank-local latent snapshots for LatentVisCallback. + self.latent_vis = { + **(training_batch.fake_score_latent_vis_dict or {}), + **(training_batch.dmd_latent_vis_dict or {}), + "_fv_latent_layout": + getattr(training_batch, "minimax_h3_dmd_layout", None), + } + return loss_map, outputs, metrics + + def _carry_snapshot_raw_batch( + self, + batch: dict[str, Any], + ) -> dict[str, Any]: + """Adopt the incoming loader batch as a trajectory's conditioning. + + A trajectory keeps the prompt it set out with for its whole walk, + so the raw dict is snapshotted (tensors detached and kept on the + student device) and mid-walk calls rebuild the training batch from + it via ``prepare_batch``; H3's ``prepare_batch`` reads the dict + without mutating it and rebuilds the packed layout and VSA + attention metadata deterministically on every call. + """ + device = self.student.device + snapshot: dict[str, Any] = {} + for key, value in batch.items(): + if isinstance(value, torch.Tensor): + snapshot[key] = value.detach().to(device) + else: + snapshot[key] = value + return snapshot + + def _staggered_start( + self, + state: torch.Tensor, + training_batch: Any, + step_list: torch.Tensor, + slot: int, + ) -> tuple[torch.Tensor, int]: + """Pre-walk a fresh trajectory and snapshot it at this stream's rung. + + First-ever fill of a slot only. Each (rank, slot) stream starts at + ``(stagger_rank * slots + slot) % len(grid)``, spreading the + streams evenly across the grid so student updates (which revisit + rungs modulo ``gcd(len(grid), generator_update_interval)``) see + every rung. Every rank walks the whole grid under ``no_grad`` + regardless of its offset so the FSDP forwards issue a uniform + collective count — a rank-dependent count would desynchronize the + all-gathers and hang; only the kept snapshot differs per rank. + """ + rank, _ = self._rollout_carry_rank_world() + grid_len = len(step_list) + # Exact-shape batches must remain shape-synchronous across every rank. + # Rank-staggered clears would let one rank adopt the next loader bucket + # while its peers were still carrying the previous geometry. Keep the + # slot staggering, but make it rank-independent for native-shape data. + data_config = getattr(self.training_config, "data", None) + stagger_rank = 0 if bool(getattr(data_config, "native_shape_bucketing", False)) else rank + offset = (stagger_rank * self._rollout_carry_slot_count + slot) % grid_len + device = state.device + snapshot = state + with torch.no_grad(): + for rung in range(grid_len - 1): + timestep = step_list[rung] * torch.ones( + 1, + device=device, + dtype=torch.long, + ) + pred_x0 = self.student.predict_x0( + state, + timestep, + training_batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="vsa", + ) + state = self._renoise( + state, + pred_x0.detach(), + timestep, + rung + 1, + step_list, + training_batch, + ) + if rung + 1 == offset: + snapshot = state + return snapshot, offset + + def _renoise( + self, + state: torch.Tensor, + pred_x0: torch.Tensor, + timestep: torch.Tensor, + next_rung: int, + step_list: torch.Tensor, + batch: Any | None = None, + ) -> torch.Tensor: + """Re-noise an x0 prediction made at ``timestep`` onto the next rung. + + ``sde`` draws fresh noise — the existing full-rollout hop. + ``ode`` reuses the noise the current state implies per modality + (``eps_m = (x_t - alpha_m(t) x0) / sigma_m(t)`` with each + modality's shifted sigma, via the adapter's ``extract_eps``), the + deterministic step the FastGen H3 recipe uses. + """ + device = state.device + next_timestep = step_list[next_rung] * torch.ones( + 1, + device=device, + dtype=torch.long, + ) + if self._rollout_sample_type == "ode": + eps = self._extract_eps_for_batch(self.student, state, pred_x0, timestep, batch) + else: + eps = torch.randn( + state.shape, + device=device, + dtype=pred_x0.dtype, + generator=self.cuda_generator, + ) + return self._add_noise_for_batch(self.student, pred_x0, eps, next_timestep, batch) + + def _advance_carry( + self, + slot: int, + state: torch.Tensor, + generator_pred_x0: torch.Tensor, + timestep: torch.Tensor, + rung: int, + step_list: torch.Tensor, + training_batch: Any, + raw_batch: dict[str, Any], + ) -> None: + """Hand the one paid-for step to the slot, or clear a finished walk. + + The advanced state is detached and produced under ``no_grad``: it + feeds a later call, not a gradient path. Walking past the last rung + ends the trajectory (the terminal clean sample is never trained + on), so the slot empties and the next call starts fresh at rung 0. + """ + if rung + 1 >= len(step_list): + self._carry_slots[slot] = None + return + with torch.no_grad(): + next_state = self._renoise( + state, + generator_pred_x0.detach(), + timestep, + rung + 1, + step_list, + training_batch, + ) + self._carry_slots[slot] = { + "state": next_state.detach(), + "rung": rung + 1, + "raw_batch": raw_batch, + } + def _critic_flow_matching_loss( self, batch: Any, - ) -> tuple[torch.Tensor, Any, dict[str, Any]]: - with torch.no_grad(): - generator_pred_x0 = self._student_rollout(batch, with_grad=False) + *, + generator_pred_x0: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, Any, dict[str, Any], dict[str, LogScalar]]: + if generator_pred_x0 is None: + with torch.no_grad(): + generator_pred_x0 = self._student_rollout(batch, with_grad=False) device = generator_pred_x0.device fake_score_timestep = self._sample_score_timestep(device) @@ -614,18 +1605,72 @@ def _critic_flow_matching_loss( dtype=generator_pred_x0.dtype, generator=self.cuda_generator, ) - noisy_x0 = self.student.add_noise(generator_pred_x0, noise, fake_score_timestep) - - pred_noise = self.critic.predict_noise( - noisy_x0, + noisy_x0 = self._add_noise_for_batch( + self.student, + generator_pred_x0, + noise, fake_score_timestep, batch, - conditional=True, - cfg_uncond=self._cfg_uncond, - attn_kind="dense", ) - target = noise - generator_pred_x0 - flow_matching_loss = torch.mean((pred_noise - target)**2) + + slices = self._modality_slices(batch) + emit_modality_metrics = slices is not None + if slices is None: + slices = (("packed", slice(None)), ) + all_x0 = all(self._fake_score_space_for(name) == "x0" for name, _ in slices) + + pred_x0: torch.Tensor | None = None + pred_noise: torch.Tensor | None = None + target: torch.Tensor | None = None + if all_x0: + # Match FastGen's fake_score_pred_type=x0 objective directly. + # Reweighting raw velocity MSE by sigma^2 is algebraically equal + # before rounding, but estimating sigma from BF16 x_t introduces a + # low-noise bias (especially for H3 audio). + pred_x0 = self.critic.predict_x0( + noisy_x0, + fake_score_timestep, + batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="dense", + ) + else: + # Retain the single-forward legacy path for mixed per-modality + # velocity/x0 configurations. Strict H3 parity uses global x0 and + # therefore always takes the direct branch above. + pred_noise = self.critic.predict_noise( + noisy_x0, + fake_score_timestep, + batch, + conditional=True, + cfg_uncond=self._cfg_uncond, + attn_kind="dense", + ) + target = noise - generator_pred_x0 + + flow_matching_loss = torch.zeros((), device=device, dtype=torch.float32) + metrics: dict[str, LogScalar] = {} + for name, modality in slices: + if all_x0: + assert pred_x0 is not None + loss_m = torch.mean((pred_x0[:, modality].float() - generator_pred_x0[:, modality].float())**2) + else: + assert pred_noise is not None and target is not None + loss_m = torch.mean((pred_noise[:, modality].float() - target[:, modality].float())**2) + if not all_x0 and self._fake_score_space_for(name) == "x0": + # For affine rectified flow, x0 MSE is sigma_m(t)^2 times + # velocity MSE. This compatibility path is retained only for + # legacy mixed-space recipes. + assert target is not None + with torch.no_grad(): + num = torch.mean((noisy_x0[:, modality].float() - generator_pred_x0[:, modality].float())**2) + den = torch.mean(target[:, modality].float()**2) + sigma_sq = num / den + loss_m = sigma_sq * loss_m + flow_matching_loss = flow_matching_loss + self._modality_weight(name) * loss_m + if emit_modality_metrics: + metrics[f"fake_score_loss_{name}"] = loss_m.detach() batch.fake_score_latent_vis_dict = { "generator_pred_video": generator_pred_x0, @@ -636,13 +1681,14 @@ def _critic_flow_matching_loss( flow_matching_loss, (batch.timesteps, batch.attn_metadata), outputs, + metrics, ) def _dmd_loss( self, generator_pred_x0: torch.Tensor, batch: Any, - ) -> torch.Tensor: + ) -> tuple[torch.Tensor, dict[str, LogScalar]]: guidance_scale = get_optional_float( self.method_config, "real_score_guidance_scale", @@ -661,7 +1707,13 @@ def _dmd_loss( dtype=generator_pred_x0.dtype, generator=self.cuda_generator, ) - noisy_latents = self.student.add_noise(generator_pred_x0, noise, timestep) + noisy_latents = self._add_noise_for_batch( + self.student, + generator_pred_x0, + noise, + timestep, + batch, + ) faker_x0 = self.critic.predict_x0( noisy_latents, @@ -679,22 +1731,46 @@ def _dmd_loss( cfg_uncond=self._cfg_uncond, attn_kind="dense", ) - real_uncond_x0 = self.teacher.predict_x0( - noisy_latents, - timestep, - batch, - conditional=False, - cfg_uncond=self._cfg_uncond, - attn_kind="dense", - ) - real_cfg_x0 = real_uncond_x0 + (real_cond_x0 - real_uncond_x0) * guidance_scale - - denom = torch.abs(generator_pred_x0 - real_cfg_x0).mean() - grad = (faker_x0 - real_cfg_x0) / denom - grad = torch.nan_to_num(grad) - - loss = 0.5 * F.mse_loss( - generator_pred_x0.float(), - (generator_pred_x0.float() - grad.float()).detach(), - ) - return loss + if float(guidance_scale) == 1.0: + # Scale 1 is the conditional prediction and needs no + # unconditional forward. + real_cfg_x0 = real_cond_x0 + else: + real_uncond_x0 = self.teacher.predict_x0( + noisy_latents, + timestep, + batch, + conditional=False, + cfg_uncond=self._cfg_uncond, + attn_kind="dense", + ) + real_cfg_x0 = real_uncond_x0 + (real_cond_x0 - real_uncond_x0) * guidance_scale + + # LatentVisCallback decodes these estimates on rank 0. + batch.dmd_latent_vis_dict.update({ + "real_score_pred_video": real_cfg_x0.detach(), + "faker_score_pred_video": faker_x0.detach(), + "dmd_timestep": timestep.detach(), + }) + + slices = self._modality_slices(batch) + emit_modality_metrics = slices is not None + if slices is None: + slices = (("packed", slice(None)), ) + loss = torch.zeros((), device=device, dtype=torch.float32) + metrics: dict[str, LogScalar] = {} + for name, modality in slices: + gen_m = generator_pred_x0[:, modality].float() + with torch.no_grad(): + real_m = real_cfg_x0[:, modality].float() + denom = (gen_m - real_m).abs().mean() + 1e-6 + grad = (faker_x0[:, modality].float() - real_m) / denom + if not bool(torch.isfinite(grad).all()): + raise RuntimeError( + f"Nonfinite DMD2 distribution direction for {name}" + ) + loss_m = 0.5 * F.mse_loss(gen_m, (gen_m - grad).detach()) + loss = loss + self._modality_weight(name) * loss_m + if emit_modality_metrics: + metrics[f"generator_loss_{name}"] = loss_m.detach() + return loss, metrics diff --git a/fastvideo/train/models/base.py b/fastvideo/train/models/base.py index 4ef95a3dcd..1a6357fd7d 100644 --- a/fastvideo/train/models/base.py +++ b/fastvideo/train/models/base.py @@ -153,6 +153,22 @@ def add_noise( ) -> torch.Tensor: """Apply forward-process noise at *timestep*.""" + def add_noise_for_batch( + self, + clean_latents: torch.Tensor, + noise: torch.Tensor, + timestep: torch.Tensor, + batch: TrainingBatch, + ) -> torch.Tensor: + """Apply noise with optional immutable geometry carried by ``batch``. + + Video-only and fixed-shape models keep their existing behavior. Joint + packed models override this hook when splitting the tensor requires + batch-local shape context. + """ + del batch + return self.add_noise(clean_latents, noise, timestep) + @abstractmethod def predict_noise( self, diff --git a/fastvideo/train/models/minimax_h3/__init__.py b/fastvideo/train/models/minimax_h3/__init__.py index acbdc51ae2..006658f47c 100644 --- a/fastvideo/train/models/minimax_h3/__init__.py +++ b/fastvideo/train/models/minimax_h3/__init__.py @@ -3,3 +3,7 @@ from fastvideo.train.models.minimax_h3.minimax_h3 import ( MiniMaxH3Model as MiniMaxH3Model, ) +from fastvideo.train.models.minimax_h3.minimax_h3_dmd import ( + MiniMaxH3DMDLatentLayout as MiniMaxH3DMDLatentLayout, + MiniMaxH3DMDModel as MiniMaxH3DMDModel, +) diff --git a/fastvideo/train/models/minimax_h3/minimax_h3.py b/fastvideo/train/models/minimax_h3/minimax_h3.py index 0788880182..adf30deb89 100644 --- a/fastvideo/train/models/minimax_h3/minimax_h3.py +++ b/fastvideo/train/models/minimax_h3/minimax_h3.py @@ -3,16 +3,22 @@ from __future__ import annotations +import math from typing import Any, Literal, TYPE_CHECKING import torch +import fastvideo.envs as envs +from fastvideo.attention.backends.video_sparse_attn_h3 import ( + MiniMaxH3VSAMetadataBuilder, ) +from fastvideo.dataset.shape_bucket import parse_video_shape_bucket_id from fastvideo.distributed import get_sp_group from fastvideo.forward_context import set_forward_context from fastvideo.models.schedulers.scheduling_minimax_h3 import MiniMaxH3Scheduler from fastvideo.pipelines import TrainingBatch from fastvideo.pipelines.basic.minimax_h3.packing import ( MINIMAX_H3_AUDIO_CHANNELS, + MINIMAX_H3_CANVAS_MULTIPLE, MINIMAX_H3_TEXT_TAG, MiniMaxH3PackedLayout, audio_latent_num_frames, @@ -21,7 +27,10 @@ patchify_video_latents, unpack_audio_tokens, unpatchify_video_tokens, + video_latent_num_frames, ) +from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_denoising import ( + _h3_vsa_prefix_segments, ) from fastvideo.platforms import AttentionBackendEnum from fastvideo.train.models.base import ModelBase, NoisePrediction @@ -38,13 +47,31 @@ _AUDIO_SCHEDULER_SHIFT = 3.0 _VIDEO_LATENT_CHANNELS = 24 _AUDIO_LATENT_CHANNELS = 32 +_AUDIO_SAMPLE_RATE = 32_000 + +# Dense TORCH_SDPA is the default; per-role overrides allow FLASH_ATTN +# (teacher/critic, FA4 via FASTVIDEO_FA4=1) and the packed-sequence VSA-H3 +# backend (distillation student). +_ALLOWED_ATTENTION_BACKENDS = ( + AttentionBackendEnum.TORCH_SDPA, + AttentionBackendEnum.FLASH_ATTN, + AttentionBackendEnum.VIDEO_SPARSE_ATTN_H3, +) -def shift_noise_amount(base_noise_amount: torch.Tensor, shift: float) -> torch.Tensor: - """Apply the MiniMax H3 rational shift to a unit noise amount.""" +def shift_noise_amount( + base_noise_amount: torch.Tensor, + shift: float, + *, + max_noise_amount: float = 1.0, +) -> torch.Tensor: + """Apply the MiniMax H3 rational shift on ``[0, max_noise_amount]``.""" if shift <= 0: raise ValueError(f"shift must be positive, got {shift}") - return shift * base_noise_amount / (1.0 + (shift - 1.0) * base_noise_amount) + if not 0.0 < max_noise_amount <= 1.0: + raise ValueError("max_noise_amount must satisfy 0 < max <= 1, " + f"got {max_noise_amount}") + return (shift * base_noise_amount * max_noise_amount / (base_noise_amount * (shift - 1.0) + max_noise_amount)) class MiniMaxH3Model(ModelBase): @@ -62,16 +89,19 @@ def __init__( enable_gradient_checkpointing_type: str | None = None, transformer_override_safetensor: str | None = None, attention_backend: AttentionBackendEnum | str | None = AttentionBackendEnum.TORCH_SDPA, + construction_precision: str | None = None, ) -> None: """Validate the single-document T2VA contract and load the transformer.""" super().__init__( trainable=trainable, attention_backend=attention_backend, ) - # PyTorch scaled dot product attention (SDPA) provides dense attention - # without adding another attention-kernel dependency to H3 training. - if self.attention_backend != AttentionBackendEnum.TORCH_SDPA: - raise ValueError("MiniMaxH3Model requires the TORCH_SDPA attention backend") + # Attention layers bind their backend during construction, so this is + # the per-role selection point (student/teacher/critic can differ). + if self.attention_backend not in _ALLOWED_ATTENTION_BACKENDS: + allowed = ", ".join(b.name for b in _ALLOWED_ATTENTION_BACKENDS) + raise ValueError("MiniMaxH3Model supports the attention backends " + f"{{{allowed}}}, got {self.attention_backend}") if training_config.pipeline_config is None: raise ValueError("MiniMaxH3Model requires a resolved MiniMax H3 pipeline config") # Packed row indices describe one text-video-audio document without a @@ -83,15 +113,19 @@ def __init__( if float(training_config.data.training_cfg_rate) != 0.0: raise ValueError("MiniMaxH3Model requires training.data.training_cfg_rate=0.0") # Joint supervision requires paired video and stereo-audio latents from - # every parquet row. - if str(training_config.data.preprocessed_data_type) != "t2va": - raise ValueError("MiniMaxH3Model requires training.data.preprocessed_data_type='t2va'") - - # FastVideo's Fully Sharded Data Parallel loading path requires one BF16 - # parameter dtype, including modules that H3 inference keeps in FP32. - training_config.pipeline_config.dit_config.uniform_parameter_dtype = True # type: ignore[attr-defined] + # every parquet row ('t2va'). 'text_only' rows carry prompt conditioning + # alone and are valid only for data-free methods that synthesize latent + # shapes from config (DMD2 enforces rollout_mode='simulate' for them). + if str(training_config.data.preprocessed_data_type) not in ("t2va", "text_only"): + raise ValueError("MiniMaxH3Model requires training.data.preprocessed_data_type " + "'t2va' or 'text_only'") + configured_precision = str(getattr(training_config, "dit_precision", "fp32")) + if trainable and construction_precision not in (None, configured_precision): + raise ValueError("A trainable MiniMaxH3 role cannot override construction_precision; " + "FP32 optimizer masters must follow training.dit_precision") self._init_from = str(init_from) + self._construction_precision = construction_precision self.training_config = training_config self.transformer = self._load_transformer( trainable=trainable, @@ -123,6 +157,7 @@ def _load_transformer( override_transformer_cls_name=self._transformer_cls_name, transformer_override_safetensor=transformer_override_safetensor, attention_backend=self.attention_backend, + construction_precision=self._construction_precision, ) checkpointing_type = (enable_gradient_checkpointing_type or self.training_config.model.enable_gradient_checkpointing_type) @@ -135,15 +170,17 @@ def _load_transformer( def init_preprocessors(self, training_config: TrainingConfig) -> None: """Load precomputed text embeddings and paired video-audio latents.""" - from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2va + from fastvideo.dataset.dataloader.schema import (pyarrow_schema_t2va, pyarrow_schema_text_only) from fastvideo.train.utils.dataloader import build_parquet_t2v_train_dataloader self.sp_group = get_sp_group() text_config = training_config.pipeline_config.text_encoder_configs[0] # type: ignore[union-attr] + parquet_schema = (pyarrow_schema_text_only + if str(training_config.data.preprocessed_data_type) == "text_only" else pyarrow_schema_t2va) self.dataloader = build_parquet_t2v_train_dataloader( training_config.data, text_len=int(text_config.arch_config.text_len), - parquet_schema=pyarrow_schema_t2va, + parquet_schema=parquet_schema, ) self.start_step = 0 @@ -154,30 +191,57 @@ def _resolve_clean_latents( dtype: torch.dtype, device: torch.device, ) -> tuple[torch.Tensor, torch.Tensor]: - """Resolve fixed visual and stereo-audio latent tensors for one sample.""" + """Resolve fixed or native visual and stereo-audio latents.""" data_config = self.training_config.data + native_shapes = bool(getattr(data_config, "native_shape_bucketing", False)) if latents_source == "data": if "vae_latent" not in raw_batch or "audio_latent" not in raw_batch: raise ValueError("A T2VA batch requires vae_latent and audio_latent tensors") video_latents = raw_batch["vae_latent"] audio_latents = raw_batch["audio_latent"] elif latents_source == "zeros": + latent_frames = int(data_config.num_latent_t) + height = int(data_config.num_height) + width = int(data_config.num_width) + num_frames = int(data_config.num_frames) + if native_shapes: + bucket_id = raw_batch.get("_shape_bucket_id") + if not isinstance(bucket_id, str): + raise ValueError( + "Native-shape data-free batches require the exact-shape sampler to set _shape_bucket_id") + bucket = parse_video_shape_bucket_id(bucket_id) + width = bucket.width + height = bucket.height + num_frames = bucket.num_frames + latent_frames = video_latent_num_frames(num_frames) + if width % MINIMAX_H3_CANVAS_MULTIPLE or height % MINIMAX_H3_CANVAS_MULTIPLE: + raise ValueError(f"Native pixel geometry {width}x{height} must use the H3 canvas multiple " + f"{MINIMAX_H3_CANVAS_MULTIPLE}") + latent_geometry = (latent_frames, height // 16, width // 16) + patch_size = tuple(int(value) for value in self.transformer.patch_size) + if any(value % patch for value, patch in zip(latent_geometry, patch_size, strict=True)): + raise ValueError( + f"Native latent geometry {latent_geometry} is not divisible by transformer patch {patch_size}") video_latents = torch.zeros( 1, _VIDEO_LATENT_CHANNELS, - data_config.num_latent_t, - data_config.num_height // 16, - data_config.num_width // 16, + latent_frames, + height // 16, + width // 16, ) audio_latents = torch.zeros( 1, MINIMAX_H3_AUDIO_CHANNELS, _AUDIO_LATENT_CHANNELS, - audio_latent_num_frames(data_config.num_frames), + audio_latent_num_frames(num_frames), ) else: raise ValueError(f"Unknown latents_source: {latents_source!r}") + if not isinstance(video_latents, torch.Tensor): + raise ValueError(f"vae_latent must be a tensor, got {type(video_latents).__name__}") + if not isinstance(audio_latents, torch.Tensor): + raise ValueError(f"audio_latent must be a tensor, got {type(audio_latents).__name__}") if video_latents.ndim != 5 or tuple(video_latents.shape[:2]) != (1, _VIDEO_LATENT_CHANNELS): raise ValueError("vae_latent must have shape [1, 24, latent_frames, latent_height, latent_width], " f"got {tuple(video_latents.shape)}") @@ -188,19 +252,102 @@ def _resolve_clean_latents( ): raise ValueError("audio_latent must have shape [1, 2, 32, audio_frames], " f"got {tuple(audio_latents.shape)}") - if data_config.num_latent_t > 0: - video_latents = video_latents[:, :, :data_config.num_latent_t] - expected_audio_frames = audio_latent_num_frames(data_config.num_frames) - audio_latents = audio_latents[:, :, :, :expected_audio_frames] - if video_latents.shape[2] != data_config.num_latent_t: - raise ValueError("vae_latent contains fewer frames than training.data.num_latent_t") - if audio_latents.shape[-1] != expected_audio_frames: - raise ValueError("audio_latent length does not match training.data.num_frames") + + if latents_source == "data" and native_shapes: + self._validate_native_latents(raw_batch, video_latents, audio_latents) + elif not native_shapes: + # Preserve the legacy fixed-shape contract for configs that have + # not opted into exact-shape bucketing. + if data_config.num_latent_t > 0: + video_latents = video_latents[:, :, :data_config.num_latent_t] + expected_audio_frames = audio_latent_num_frames(data_config.num_frames) + audio_latents = audio_latents[:, :, :, :expected_audio_frames] + if video_latents.shape[2] != data_config.num_latent_t: + raise ValueError("vae_latent contains fewer frames than training.data.num_latent_t") + if audio_latents.shape[-1] != expected_audio_frames: + raise ValueError("audio_latent length does not match training.data.num_frames") return ( video_latents.to(device=device, dtype=dtype), audio_latents.to(device=device, dtype=dtype), ) + def _validate_native_latents( + self, + raw_batch: dict[str, Any], + video_latents: torch.Tensor, + audio_latents: torch.Tensor, + ) -> None: + """Cross-check the canonical bucket, row metadata, and latent clocks.""" + bucket_id = raw_batch.get("_shape_bucket_id") + if not isinstance(bucket_id, str): + raise ValueError("Native-shape T2VA batches require the exact-shape sampler to set _shape_bucket_id") + bucket = parse_video_shape_bucket_id(bucket_id) + + infos = raw_batch.get("info_list") + if not isinstance(infos, list) or len(infos) != 1 or not isinstance(infos[0], dict): + raise ValueError("Native-shape T2VA batches require exactly one info_list metadata record") + info = infos[0] + + def _metadata_int(name: str) -> int: + value = info.get(name) + if value is None or isinstance(value, bool): + raise ValueError(f"T2VA metadata {name!r} must be a positive integer, got {value!r}") + try: + result = int(value) + except (TypeError, ValueError) as exc: + raise ValueError(f"T2VA metadata {name!r} must be a positive integer, got {value!r}") from exc + if result <= 0 or result != value: + raise ValueError(f"T2VA metadata {name!r} must be a positive integer, got {value!r}") + return result + + width = _metadata_int("width") + height = _metadata_int("height") + num_frames = _metadata_int("num_frames") + audio_sample_rate = _metadata_int("audio_sample_rate") + if (width, height, num_frames) != (bucket.width, bucket.height, bucket.num_frames): + raise ValueError(f"Shape bucket {bucket_id!r} disagrees with row metadata " + f"width={width}, height={height}, num_frames={num_frames}") + if audio_sample_rate != _AUDIO_SAMPLE_RATE: + raise ValueError( + f"MiniMax H3 T2VA audio must be encoded at {_AUDIO_SAMPLE_RATE} Hz, got {audio_sample_rate}") + fps_value = info.get("fps") + if fps_value is None: + raise ValueError("T2VA metadata 'fps' must be numeric, got None") + try: + fps = float(fps_value) + except (TypeError, ValueError) as exc: + raise ValueError(f"T2VA metadata 'fps' must be numeric, got {info.get('fps')!r}") from exc + if not math.isfinite(fps) or not math.isclose(fps, 24.0, rel_tol=0.0, abs_tol=0.05): + raise ValueError(f"MiniMax H3 T2VA video must use the 24 fps clock, got {fps}") + + if width % MINIMAX_H3_CANVAS_MULTIPLE or height % MINIMAX_H3_CANVAS_MULTIPLE: + raise ValueError(f"Native pixel geometry {width}x{height} must use the H3 canvas multiple " + f"{MINIMAX_H3_CANVAS_MULTIPLE}") + expected_video_shape = ( + 1, + _VIDEO_LATENT_CHANNELS, + video_latent_num_frames(num_frames), + height // 16, + width // 16, + ) + expected_audio_shape = ( + 1, + MINIMAX_H3_AUDIO_CHANNELS, + _AUDIO_LATENT_CHANNELS, + audio_latent_num_frames(num_frames), + ) + if tuple(video_latents.shape) != expected_video_shape: + raise ValueError(f"vae_latent shape does not match {bucket_id!r}: expected {expected_video_shape}, " + f"got {tuple(video_latents.shape)}") + if tuple(audio_latents.shape) != expected_audio_shape: + raise ValueError(f"audio_latent shape does not match the {num_frames}-frame H3 audio clock: " + f"expected {expected_audio_shape}, got {tuple(audio_latents.shape)}") + latent_geometry = (expected_video_shape[2], expected_video_shape[3], expected_video_shape[4]) + patch_size = tuple(int(value) for value in self.transformer.patch_size) + if any(value % patch for value, patch in zip(latent_geometry, patch_size, strict=True)): + raise ValueError( + f"Native latent geometry {latent_geometry} is not divisible by transformer patch {patch_size}") + def _sample_noise_amounts( self, generator: torch.Generator, @@ -294,8 +441,41 @@ def prepare_batch( training_batch.minimax_h3_layout = layout training_batch.attn_metadata = None training_batch.attn_metadata_vsa = None + self._maybe_build_vsa_metadata(training_batch) return training_batch + def _maybe_build_vsa_metadata(self, batch: TrainingBatch) -> None: + """Populate the VSA view when this model runs the VSA-H3 backend. + + Methods route it through ``predict_noise(attn_kind="vsa")``; dense + models keep ``attn_metadata_vsa`` as ``None`` so their forwards stay + dense. One builder per model keeps the padded tile buffer reused + across steps (see ``MiniMaxH3VSAMetadata.tile_buf_holder``). + """ + backend = (getattr(self, "attention_backend_name", None) or envs.FASTVIDEO_ATTENTION_BACKEND) + if backend != "VIDEO_SPARSE_ATTN_H3": + return + layout = batch.minimax_h3_layout + patch_size = tuple(self.transformer.patch_size) + builder = getattr(self, "_vsa_metadata_builder", None) + if builder is None: + builder = self._vsa_metadata_builder = MiniMaxH3VSAMetadataBuilder() + batch.attn_metadata_vsa = builder.build( + # Training builds one metadata per batch and reuses it across the + # step's forwards; the step index only feeds probe bookkeeping. + current_timestep=0, + raw_latent_shape=( + layout.num_video_latent_frames, + layout.latent_height, + layout.latent_width, + ), + patch_size=patch_size, + VSA_sparsity=float(self.training_config.vsa_sparsity), + prefix_segments=_h3_vsa_prefix_segments(layout, patch_size), + device=self.device, + tile_size=int(self.training_config.vsa_tile_size), + ) + def add_noise( self, clean_latents: torch.Tensor, @@ -320,10 +500,12 @@ def predict_noise( ) -> NoisePrediction: """Pack modality timesteps and convert H3 outputs to noise-minus-clean.""" del timestep - if not conditional or cfg_uncond is not None: - raise ValueError("MiniMaxH3Model predicts one conditional T2VA sample") - if attn_kind != "dense": - raise ValueError("MiniMaxH3Model supports dense attention for training") + # Under dense backends both metadata views are None, so "vsa" silently + # means dense (mirrors WanModel). MiniMaxH3DMDModel.prepare_batch + # populates attn_metadata_vsa when the role runs VIDEO_SPARSE_ATTN_H3. + if attn_kind not in ("dense", "vsa"): + raise ValueError(f"Unknown attn_kind: {attn_kind!r}") + attn_metadata = (batch.attn_metadata_vsa if attn_kind == "vsa" else batch.attn_metadata) layout = batch.minimax_h3_layout if not isinstance(layout, MiniMaxH3PackedLayout): raise RuntimeError("prepare_batch() must set TrainingBatch.minimax_h3_layout") @@ -332,8 +514,20 @@ def predict_noise( if batch.timesteps is None or batch.audio_timesteps is None: raise RuntimeError("prepare_batch() must set video and audio timesteps") + encoder_hidden_states = batch.encoder_hidden_states + if not conditional: + # H3 has no negative-prompt encoder at training time, so the only + # supported unconditional branch (teacher CFG in distillation) + # zeroes the text embeddings. + if (cfg_uncond or {}).get("text") != "zero": + raise ValueError("MiniMaxH3Model unconditional forwards require " + "method.cfg_uncond={'text': 'zero'}") + encoder_hidden_states = torch.zeros_like(encoder_hidden_states) + dtype = torch.bfloat16 device = self.device + video_input_dtype = noisy_latents.dtype + audio_input_dtype = batch.audio_noisy_model_input.dtype video_bcthw = noisy_latents.permute(0, 2, 1, 3, 4).to(dtype) # Match H3 checkpoint token order: video rows flatten # (C, patch_t, patch_h, patch_w), while audio rows flatten stereo @@ -352,14 +546,14 @@ def predict_noise( unique_timesteps = unique_timesteps.to(device) timestep_indices = timestep_indices.to(device) - with torch.autocast(device.type, dtype=dtype), set_forward_context( + with set_forward_context( current_timestep=unique_timesteps, - attn_metadata=None, + attn_metadata=attn_metadata, ): video_velocity, audio_velocity = self.transformer( hidden_states=video_rows[None], audio_hidden_states=audio_rows[None], - encoder_hidden_states=batch.encoder_hidden_states, + encoder_hidden_states=encoder_hidden_states, timestep=unique_timesteps, timestep_indices=timestep_indices, token_tags=layout.token_tags.to(device), @@ -379,7 +573,10 @@ def predict_noise( self.transformer.patch_size, ).permute(0, 2, 1, 3, 4) audio_prediction = unpack_audio_tokens(audio_velocity[0], num_audio_latents)[None] - return -video_prediction, -audio_prediction + return ( + (-video_prediction).to(video_input_dtype), + (-audio_prediction).to(audio_input_dtype), + ) def backward( self, @@ -390,6 +587,19 @@ def backward( ) -> None: """Restore the forward context and average accumulated microbatch gradients.""" timesteps, attn_metadata = ctx + expected_device_type = self.device.type + offenders = [] + for name, param in self.transformer.named_parameters(): + local = getattr(param, "_local_tensor", param) + if local.device.type != expected_device_type: + offenders.append(f"param {name} on {local.device}") + if param.grad is not None: + grad_local = getattr(param.grad, "_local_tensor", param.grad) + if grad_local.device.type != expected_device_type: + offenders.append(f"grad {name} on {grad_local.device}") + if offenders: + raise RuntimeError(f"{len(offenders)} training tensors off-CUDA before backward; " + f"first offenders: {offenders[:8]}") with set_forward_context( current_timestep=timesteps, attn_metadata=attn_metadata, diff --git a/fastvideo/train/models/minimax_h3/minimax_h3_dmd.py b/fastvideo/train/models/minimax_h3/minimax_h3_dmd.py new file mode 100644 index 0000000000..480b2b8ce6 --- /dev/null +++ b/fastvideo/train/models/minimax_h3/minimax_h3_dmd.py @@ -0,0 +1,522 @@ +# SPDX-License-Identifier: Apache-2.0 +"""MiniMax H3 distribution-matching adapter (packed dual-modality latents).""" + +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import Any, cast, Literal + +import torch + +from fastvideo.pipelines import TrainingBatch +from fastvideo.pipelines.basic.minimax_h3.packing import ( + MINIMAX_H3_AUDIO_CHANNELS, + audio_latent_num_frames, +) +from fastvideo.train.models.minimax_h3.minimax_h3 import ( + _AUDIO_LATENT_CHANNELS, + _AUDIO_SCHEDULER_SHIFT, + _VIDEO_LATENT_CHANNELS, + _VIDEO_SCHEDULER_SHIFT, + MiniMaxH3Model, + shift_noise_amount, +) + +# DMD2 expresses score time in timestep units on [0, 1000]. Strict FastGen +# parity keeps that coordinate continuous and applies shifts on max_t=0.999. +_DMD_TIMESTEP_SCALE = 1000 +_FASTGEN_MAX_T = 0.999 + + +@dataclass(frozen=True, slots=True) +class MiniMaxH3DMDLatentLayout: + """Exact batch-local shapes behind one packed H3 DMD latent tensor.""" + + video_shape: tuple[int, int, int, int, int] + audio_shape: tuple[int, int, int, int] + + def __post_init__(self) -> None: + if self.video_shape[0] != 1 or self.video_shape[2] != _VIDEO_LATENT_CHANNELS: + raise ValueError("DMD video latents must have shape [1, T, 24, H, W], got " + f"{self.video_shape}") + if self.audio_shape[:3] != (1, MINIMAX_H3_AUDIO_CHANNELS, _AUDIO_LATENT_CHANNELS): + raise ValueError("DMD audio latents must have shape [1, 2, 32, Ta], got " + f"{self.audio_shape}") + if any(value <= 0 for value in self.video_shape + self.audio_shape): + raise ValueError("DMD latent layout dimensions must all be positive") + + @classmethod + def from_latents( + cls, + video_latents: torch.Tensor, + audio_latents: torch.Tensor, + ) -> MiniMaxH3DMDLatentLayout: + if video_latents.ndim != 5 or audio_latents.ndim != 4: + raise ValueError("H3 DMD pack expects video [1,T,24,H,W] and audio [1,2,32,Ta], " + f"got {tuple(video_latents.shape)} and {tuple(audio_latents.shape)}") + return cls( + video_shape=cast(tuple[int, int, int, int, int], tuple(int(value) for value in video_latents.shape)), + audio_shape=cast(tuple[int, int, int, int], tuple(int(value) for value in audio_latents.shape)), + ) + + @property + def video_numel(self) -> int: + return math.prod(self.video_shape) + + @property + def audio_numel(self) -> int: + return math.prod(self.audio_shape) + + @property + def packed_numel(self) -> int: + return self.video_numel + self.audio_numel + + def modality_slices(self) -> tuple[tuple[str, slice], ...]: + return ( + ("video", slice(0, self.video_numel)), + ("audio", slice(self.video_numel, self.packed_numel)), + ) + + +class MiniMaxH3DMDModel(MiniMaxH3Model): + """Present H3's dual (video, audio) streams to DMD2 as one packed tensor. + + ``DMD2Method``'s rollout and loss math assume one latent tensor per + sample. This adapter flattens both modality latents into one ``[1, N]`` + tensor (video's ``[1, T, 24, H, W]`` elements first, stereo audio's + ``[1, 2, 32, Ta]`` elements after) so ``dmd2.py`` stays model-agnostic. + Method timesteps become one shared base noise amount that is shifted per + modality (video 12.0, audio 3.0). Continuous score times use FastGen's + ``max_t=0.999`` domain; integer rollout rungs preserve the release + pipeline's unit-domain grid. + + ``modality_slices()`` exposes the packed video/audio column ranges so + DMD2 computes losses and normalizers per modality instead of one packed + mean that video's element count would dominate. + """ + + @property + def num_train_timesteps(self) -> int: + return _DMD_TIMESTEP_SCALE + + # ------------------------------------------------------------------ + # Packed dual-modality helpers + # ------------------------------------------------------------------ + + def _modality_shapes(self) -> tuple[tuple[int, int, int, int, int], tuple[int, int, int, int]]: + """Return the ``[1, T, C, H, W]`` video and ``[1, 2, 32, Ta]`` audio shapes.""" + data = self.training_config.data + video_shape = ( + 1, + int(data.num_latent_t), + _VIDEO_LATENT_CHANNELS, + int(data.num_height) // 16, + int(data.num_width) // 16, + ) + audio_shape = ( + 1, + MINIMAX_H3_AUDIO_CHANNELS, + _AUDIO_LATENT_CHANNELS, + audio_latent_num_frames(int(data.num_frames)), + ) + return video_shape, audio_shape + + def _fixed_latent_layout(self) -> MiniMaxH3DMDLatentLayout: + video_shape, audio_shape = self._modality_shapes() + return MiniMaxH3DMDLatentLayout( + video_shape=video_shape, + audio_shape=audio_shape, + ) + + def _batch_latent_layout(self, batch: TrainingBatch) -> MiniMaxH3DMDLatentLayout: + layout = batch.minimax_h3_dmd_layout + if not isinstance(layout, MiniMaxH3DMDLatentLayout): + raise RuntimeError("prepare_batch() must set TrainingBatch.minimax_h3_dmd_layout") + return layout + + def pack_latents( + self, + video_latents: torch.Tensor, + audio_latents: torch.Tensor, + *, + layout: MiniMaxH3DMDLatentLayout | None = None, + ) -> torch.Tensor: + """Flatten both modality latents into one ``[1, N]`` tensor.""" + actual_layout = MiniMaxH3DMDLatentLayout.from_latents(video_latents, audio_latents) + if layout is not None and actual_layout != layout: + raise ValueError(f"Latent tensors do not match their batch layout: {actual_layout} != {layout}") + return torch.cat( + (video_latents.reshape(1, -1), audio_latents.reshape(1, -1)), + dim=1, + ) + + def modality_slices( + self, + *, + layout: MiniMaxH3DMDLatentLayout | None = None, + ) -> tuple[tuple[str, slice], ...]: + """Named packed-latent column slices for per-modality DMD2 losses. + + Video is ~3.6M packed elements against audio's ~15-30k, so a single + global mean would give audio <1% of the distillation signal; DMD2 + consumes these slices to normalize and weight each stream separately. + """ + return (layout or self._fixed_latent_layout()).modality_slices() + + def modality_slices_for_batch(self, batch: TrainingBatch) -> tuple[tuple[str, slice], ...]: + """Return per-modality slices for this batch's native shape.""" + return self.modality_slices(layout=self._batch_latent_layout(batch)) + + def unpack_latents( + self, + packed: torch.Tensor, + *, + layout: MiniMaxH3DMDLatentLayout | None = None, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Split one packed ``[1, N]`` tensor back into (video, audio) latents.""" + resolved = layout or self._fixed_latent_layout() + if packed.shape != (1, resolved.packed_numel): + raise ValueError(f"Packed latents must have shape [1, {resolved.packed_numel}], got {tuple(packed.shape)}") + return ( + packed[:, :resolved.video_numel].reshape(resolved.video_shape), + packed[:, resolved.video_numel:].reshape(resolved.audio_shape), + ) + + def _noise_amounts(self, timestep: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Map one method timestep to both modality noise amounts in FP64.""" + # Strict FastGen score times are continuous and live on max_t=0.999. + # Integer rollout rungs retain the release pipeline's unit-domain + # schedule so training and offline validation walk identical states. + warp_max = (_FASTGEN_MAX_T if timestep.is_floating_point() else 1.0) + base = (timestep.reshape(-1)[:1].to(torch.float64) / _DMD_TIMESTEP_SCALE) + base = base.clamp(0.0, warp_max) + return ( + shift_noise_amount( + base, + _VIDEO_SCHEDULER_SHIFT, + max_noise_amount=warp_max, + ), + shift_noise_amount( + base, + _AUDIO_SCHEDULER_SHIFT, + max_noise_amount=warp_max, + ), + ) + + # ------------------------------------------------------------------ + # ModelBase overrides (packed convention) + # ------------------------------------------------------------------ + + def set_requires_negative_conditioning(self, requires: bool) -> None: + """Fail fast: H3 cannot encode negative prompts at training time.""" + if requires: + raise ValueError("MiniMaxH3DMDModel has no negative-prompt encoder; set " + "method.cfg_uncond={'text': 'zero'} for unconditional forwards") + + def prepare_batch( + self, + raw_batch: dict[str, Any], + *, + generator: torch.Generator, + latents_source: Literal["data", "zeros"] = "data", + ) -> TrainingBatch: + """Prepare the T2VA batch, then expose clean latents in packed form.""" + batch = super().prepare_batch( + raw_batch, + generator=generator, + latents_source=latents_source, + ) + # DMD2 draws its own noise and timesteps per forward; only the packed + # clean latents matter here (the base prepare_batch already built the + # VSA metadata view for VSA-H3 roles). The fine-tuning noisy fields + # are refreshed by predict_noise on every call. + if batch.latents is None or batch.audio_latents is None: + raise RuntimeError("MiniMax H3 batch preparation did not produce paired latents") + layout = MiniMaxH3DMDLatentLayout.from_latents(batch.latents, batch.audio_latents) + batch.minimax_h3_dmd_layout = layout + batch.latents = self.pack_latents(batch.latents, batch.audio_latents, layout=layout) + return batch + + def add_noise( + self, + clean_latents: torch.Tensor, + noise: torch.Tensor, + timestep: torch.Tensor, + ) -> torch.Tensor: + """Noise legacy fixed-shape packed latents at one shared timestep.""" + return self._add_noise_with_layout( + clean_latents, + noise, + timestep, + self._fixed_latent_layout(), + ) + + def add_noise_for_batch( + self, + clean_latents: torch.Tensor, + noise: torch.Tensor, + timestep: torch.Tensor, + batch: TrainingBatch, + ) -> torch.Tensor: + """Noise packed latents using immutable geometry from ``batch``.""" + layout = self._batch_latent_layout(batch) + if not bool(getattr(self.training_config.data, "native_shape_bucketing", False)): + # Keep the established fixed-shape call path observable and + # byte-identical for existing data-free/data-forcing recipes. + return self.add_noise(clean_latents, noise, timestep) + return self._add_noise_with_layout( + clean_latents, + noise, + timestep, + layout, + ) + + def _add_noise_with_layout( + self, + clean_latents: torch.Tensor, + noise: torch.Tensor, + timestep: torch.Tensor, + layout: MiniMaxH3DMDLatentLayout, + ) -> torch.Tensor: + sigma_video, sigma_audio = self._noise_amounts(timestep) + clean_video, clean_audio = self.unpack_latents(clean_latents, layout=layout) + noise_video, noise_audio = self.unpack_latents(noise, layout=layout) + return self.pack_latents( + self._mix(clean_video, noise_video, sigma_video), + self._mix(clean_audio, noise_audio, sigma_audio), + layout=layout, + ) + + @staticmethod + def _mix( + clean: torch.Tensor, + noise: torch.Tensor, + sigma: torch.Tensor, + ) -> torch.Tensor: + original_dtype = clean.dtype + clean_fp64 = clean.to(torch.float64) + noise_fp64 = noise.to(torch.float64) + sigma_fp64 = sigma.to(device=clean.device, dtype=torch.float64) + return ((1.0 - sigma_fp64) * clean_fp64 + sigma_fp64 * noise_fp64).to(original_dtype) + + def extract_eps( + self, + noisy_latents: torch.Tensor, + clean_latents: torch.Tensor, + timestep: torch.Tensor, + ) -> torch.Tensor: + """Recover noise for legacy fixed-shape packed latents. + + The inverse of :meth:`add_noise` at the same shared timestep: with + ``x_t = (1 - sigma_m) x0 + sigma_m eps`` under each modality's + shifted sigma (video 12.0, audio 3.0), the implied noise is + ``eps_m = (x_t - (1 - sigma_m) x0) / sigma_m``. DMD2's ODE renoise + uses this to step a carried trajectory deterministically between + grid rungs. + """ + return self._extract_eps_with_layout( + noisy_latents, + clean_latents, + timestep, + self._fixed_latent_layout(), + ) + + def extract_eps_for_batch( + self, + noisy_latents: torch.Tensor, + clean_latents: torch.Tensor, + timestep: torch.Tensor, + batch: TrainingBatch, + ) -> torch.Tensor: + """Recover packed noise using immutable geometry from ``batch``.""" + layout = self._batch_latent_layout(batch) + if not bool(getattr(self.training_config.data, "native_shape_bucketing", False)): + return self.extract_eps(noisy_latents, clean_latents, timestep) + return self._extract_eps_with_layout( + noisy_latents, + clean_latents, + timestep, + layout, + ) + + def _extract_eps_with_layout( + self, + noisy_latents: torch.Tensor, + clean_latents: torch.Tensor, + timestep: torch.Tensor, + layout: MiniMaxH3DMDLatentLayout, + ) -> torch.Tensor: + sigma_video, sigma_audio = self._noise_amounts(timestep) + noisy_video, noisy_audio = self.unpack_latents(noisy_latents, layout=layout) + clean_video, clean_audio = self.unpack_latents(clean_latents, layout=layout) + return self.pack_latents( + self._unmix(noisy_video, clean_video, sigma_video), + self._unmix(noisy_audio, clean_audio, sigma_audio), + layout=layout, + ) + + @staticmethod + def _unmix( + noisy: torch.Tensor, + clean: torch.Tensor, + sigma: torch.Tensor, + ) -> torch.Tensor: + original_dtype = noisy.dtype + noisy_fp64 = noisy.to(torch.float64) + clean_fp64 = clean.to(torch.float64) + sigma_fp64 = sigma.to(device=noisy.device, dtype=torch.float64) + # The DMD grid never renoises from t=0, but clamp so a degenerate + # call cannot divide by zero. + eps = ((noisy_fp64 - (1.0 - sigma_fp64) * clean_fp64) / sigma_fp64.clamp_min(1e-6)) + return eps.to(original_dtype) + + def predict_noise( + self, + noisy_latents: torch.Tensor, + timestep: torch.Tensor, + batch: TrainingBatch, + *, + conditional: bool, + cfg_uncond: dict[str, Any] | None = None, + attn_kind: Literal["dense", "vsa"] = "dense", + ) -> torch.Tensor: + """Run one packed joint forward at an explicit method timestep. + + Both modality clean-time fields on ``batch`` are rewritten from + ``timestep`` so the packed-row timestep plan and the backward + forward-context stay coherent with this call. + """ + sigma_video, sigma_audio = self._noise_amounts(timestep) + layout = self._batch_latent_layout(batch) + noisy_video, noisy_audio = self.unpack_latents(noisy_latents, layout=layout) + batch.timesteps = (1.0 - sigma_video).to(noisy_latents.device) + batch.audio_timesteps = (1.0 - sigma_audio).to(noisy_latents.device) + batch.audio_noisy_model_input = noisy_audio + video_pred, audio_pred = super().predict_noise( + noisy_video, + timestep, + batch, + conditional=conditional, + cfg_uncond=cfg_uncond, + attn_kind=attn_kind, + ) + return self.pack_latents(video_pred, audio_pred, layout=layout) + + def predict_x0( + self, + noisy_latents: torch.Tensor, + timestep: torch.Tensor, + batch: TrainingBatch, + *, + conditional: bool, + cfg_uncond: dict[str, Any] | None = None, + attn_kind: Literal["dense", "vsa"] = "dense", + ) -> torch.Tensor: + """Convert packed noise-minus-clean predictions to packed clean latents.""" + pred_noise = self.predict_noise( + noisy_latents, + timestep, + batch, + conditional=conditional, + cfg_uncond=cfg_uncond, + attn_kind=attn_kind, + ) + sigma_video, sigma_audio = self._noise_amounts(timestep) + layout = self._batch_latent_layout(batch) + noisy_video, noisy_audio = self.unpack_latents(noisy_latents, layout=layout) + pred_video, pred_audio = self.unpack_latents(pred_noise, layout=layout) + return self.pack_latents( + self._to_x0(noisy_video, pred_video, sigma_video), + self._to_x0(noisy_audio, pred_audio, sigma_audio), + layout=layout, + ) + + @staticmethod + def _to_x0( + noisy: torch.Tensor, + pred_noise: torch.Tensor, + sigma: torch.Tensor, + ) -> torch.Tensor: + # noisy = (1 - sigma) * clean + sigma * noise and pred approximates + # noise - clean, so clean = noisy - sigma * pred. + original_dtype = noisy.dtype + noisy_fp64 = noisy.to(torch.float64) + pred_noise_fp64 = pred_noise.to(torch.float64) + sigma_fp64 = sigma.to(device=noisy.device, dtype=torch.float64) + return (noisy_fp64 - sigma_fp64 * pred_noise_fp64).to(original_dtype) + + # ------------------------------------------------------------------ + # Intermediate-latent visualization (LatentVisCallback) + # ------------------------------------------------------------------ + + def _load_vis_vae(self) -> Any: + """Lazily load the H3 video VAE for visualization decodes. + + The module stays CPU-resident between decodes; ``decode_vis_latents`` + moves it to the GPU per call. Loading mirrors the H3 preprocess + scripts: the inference component registry keeps precision policy and + normalization identical to the published decode recipe. + """ + vae = getattr(self, "_vis_vae_module", None) + if vae is not None: + return vae + import os + + from fastvideo.fastvideo_args import FastVideoArgs + from fastvideo.models.loader.component_loader import PipelineComponentLoader + from fastvideo.utils import verify_model_config_and_directory + + model_index = verify_model_config_and_directory(self._init_from) + transformers_or_diffusers, _ = model_index["vae"][:2] + args = FastVideoArgs( + model_path=self._init_from, + pipeline_config=self.training_config.pipeline_config, + num_gpus=1, + tp_size=1, + sp_size=1, + hsdp_shard_dim=1, + use_fsdp_inference=False, + vae_cpu_offload=True, + text_encoder_cpu_offload=True, + ) + vae = PipelineComponentLoader.load_module( + module_name="vae", + component_model_path=os.path.join(self._init_from, "vae"), + transformers_or_diffusers=transformers_or_diffusers, + fastvideo_args=args, + ) + vae.to("cpu") + self._vis_vae_module = vae + return vae + + @torch.no_grad() + def decode_vis_latents( + self, + packed: torch.Tensor, + *, + layout: MiniMaxH3DMDLatentLayout | None = None, + ) -> Any: + """Decode the packed video stream into a uint8 ``[B, T, C, H, W]`` clip. + + Follows ``MiniMaxH3VideoDecodingStage``: denormalize latents, decode + under the published FP16-autocast-over-FP32 recipe, denormalize + pixels. The audio stream is dropped — the tracker artifact is a + silent video. + """ + video_latents, _ = self.unpack_latents(packed.detach(), layout=layout) + latents = video_latents.permute(0, 2, 1, 3, 4).to(device=self.device, dtype=torch.float32) + vae = self._load_vis_vae() + vae.to(self.device) + try: + latents = vae.denormalize_latents(latents) + with torch.autocast(self.device.type, dtype=torch.float16, enabled=self.device.type == "cuda"): + video = vae.decode(latents).sample + video = vae.denormalize_pixels(video.float()).clamp_(0.0, 1.0).cpu() + finally: + vae.to("cpu") + video = video.permute(0, 2, 1, 3, 4) + return (video * 255.0).to(torch.uint8).numpy() + + +__all__ = ["MiniMaxH3DMDLatentLayout", "MiniMaxH3DMDModel"] diff --git a/fastvideo/train/trainer.py b/fastvideo/train/trainer.py index be6054f772..e09887fe25 100644 --- a/fastvideo/train/trainer.py +++ b/fastvideo/train/trainer.py @@ -11,6 +11,7 @@ from tqdm.auto import tqdm from fastvideo.distributed import get_sp_group, get_world_group +from fastvideo.logger import init_logger from fastvideo.train.callbacks.callback import CallbackDict from fastvideo.train.methods.base import LogScalar, TrainingMethod from fastvideo.train.utils.tracking import build_tracker @@ -19,6 +20,42 @@ from fastvideo.train.utils.training_config import ( TrainingConfig, ) +logger = init_logger(__name__) + + +def _verify_master_weight_precision(method: TrainingMethod, tc: TrainingConfig) -> None: + """Refuse to train on low-precision master weights unless opted in. + + Optimizer steps applied in-place to bf16/fp16 parameters round away + updates below ~half an ulp of each weight's magnitude; O(1)-magnitude + parameters (norm gains) freeze entirely at typical distillation learning + rates, and ``zeros_like``-allocated optimizer state inherits the same + starved dtype. FP32 sharded masters (``training.dit_precision: fp32``) + fix both; ordinary FSDP groups still compute in BF16, while models may + declare narrower FP32 compute boundaries. + """ + if bool(getattr(tc.model, "allow_low_precision_master_weights", False)): + return + offenders: list[str] = [] + for role, model in getattr(method, "_role_models", {}).items(): + if not getattr(model, "_trainable", False): + continue + transformer = getattr(model, "transformer", None) + if transformer is None: + continue + for name, param in transformer.named_parameters(): + if param.requires_grad and param.dtype != torch.float32: + offenders.append(f"{role}:{name} ({param.dtype})") + break + if offenders: + raise RuntimeError("Trainable master weights are not fp32: " + f"{offenders}. bf16/fp16 parameter storage silently rounds away " + "optimizer updates below ~half an ulp per weight (norm-scale " + "parameters freeze completely). Set training.dit_precision: fp32 " + "(FP32 sharded masters; ordinary groups compute in BF16), " + "or acknowledge the effect explicitly with " + "training.model.allow_low_precision_master_weights: true.") + def _coerce_log_scalar( value: Any, @@ -114,6 +151,7 @@ def run( ) method.set_tracker(self.tracker) + _verify_master_weight_precision(method, tc) method.on_train_start() self.callbacks.on_train_start( method, @@ -127,6 +165,17 @@ def run( resumed_step = (checkpoint_manager.maybe_resume(resume_from_checkpoint=(resume_from_checkpoint))) if resumed_step is not None: start_step = int(resumed_step) + if bool(getattr(tc.checkpoint, "reset_lr_on_resume", False)): + # The DCP load above restored the checkpoint's optimizer + # LRs and scheduler base_lrs; re-apply the YAML's values. + method.apply_configured_lrs() + logger.info("reset_lr_on_resume: re-applied configured learning rates at step %s", start_step) + initial_validation_scheduled = self.callbacks.will_run_validation(iteration=start_step) + if checkpoint_manager is not None: + checkpoint_manager.maybe_save_inference( + start_step, + validation_scheduled=initial_validation_scheduled, + ) self.callbacks.on_validation_begin( method, iteration=start_step, @@ -230,7 +279,14 @@ def run( iteration=step, ) + validation_scheduled = self.callbacks.will_run_validation(iteration=step) if checkpoint_manager is not None: + # The deployable checkpoint is preserved first and corresponds + # exactly to the model this validation event will evaluate. + checkpoint_manager.maybe_save_inference( + step, + validation_scheduled=validation_scheduled, + ) checkpoint_manager.maybe_save(step) self.callbacks.on_validation_begin( diff --git a/fastvideo/train/utils/checkpoint.py b/fastvideo/train/utils/checkpoint.py index 7a3d6c1055..efc4d9747e 100644 --- a/fastvideo/train/utils/checkpoint.py +++ b/fastvideo/train/utils/checkpoint.py @@ -2,11 +2,13 @@ from __future__ import annotations +import contextlib import json import os import random import re import shutil +import time from dataclasses import dataclass from pathlib import Path from typing import Any @@ -27,6 +29,8 @@ logger = init_logger(__name__) _CHECKPOINT_DIR_RE = re.compile(r"^checkpoint-(\d+)$") +_TRAINING_CHECKPOINT_COMPLETE_MARKER = ".complete" +_RANK_RNG_STATE_RE = re.compile(r"^rng_state_rank(\d+)\.pt$") def _is_stateful(obj: Any) -> bool: @@ -52,7 +56,62 @@ def _parse_step_from_dir(checkpoint_dir: Path) -> int: return int(match.group(1)) -def _find_latest_checkpoint(output_dir: Path) -> Path | None: +def _saved_checkpoint_world_size(metadata: dict[str, Any]) -> int | None: + try: + world_size = metadata["config"]["training"]["distributed"]["num_gpus"] + except (KeyError, TypeError): + return None + if isinstance(world_size, bool) or not isinstance(world_size, int) or world_size <= 0: + return None + return world_size + + +def _is_complete_training_checkpoint( + checkpoint_dir: Path, + *, + require_complete_marker: bool, +) -> bool: + """Return whether ``checkpoint_dir`` is safe to select for resume. + + ``dcp/.metadata`` is the historical completion contract. Strict callers + additionally require the marker published after every rank has written its + RNG snapshot. Keeping strictness opt-in preserves compatibility with + checkpoints created before the stronger marker existed. + """ + dcp_metadata = checkpoint_dir / "dcp" / ".metadata" + if not dcp_metadata.is_file(): + return False + if not require_complete_marker: + return True + + try: + step = _parse_step_from_dir(checkpoint_dir) + marker = (checkpoint_dir / _TRAINING_CHECKPOINT_COMPLETE_MARKER).read_text(encoding="utf-8") + metadata = json.loads((checkpoint_dir / "metadata.json").read_text(encoding="utf-8")) + except (OSError, TypeError, ValueError): + return False + if not isinstance(metadata, dict) or marker != "complete\n" or metadata.get("step") != step: + return False + + world_size = _saved_checkpoint_world_size(metadata) + if world_size is None: + return False + expected_rng_names = {f"rng_state_rank{rank}.pt" for rank in range(world_size)} + try: + actual_rng_paths = list(checkpoint_dir.glob("rng_state_rank*.pt")) + actual_rng_names = {path.name for path in actual_rng_paths if _RANK_RNG_STATE_RE.match(path.name)} + rng_files_complete = all(path.is_file() and path.stat().st_size > 0 for path in actual_rng_paths) + except OSError: + return False + return (actual_rng_names == expected_rng_names and len(actual_rng_paths) == len(expected_rng_names) + and rng_files_complete) + + +def _find_latest_checkpoint( + output_dir: Path, + *, + require_complete_marker: bool = False, +) -> Path | None: if not output_dir.exists(): return None @@ -62,7 +121,10 @@ def _find_latest_checkpoint(output_dir: Path) -> Path | None: continue if not _CHECKPOINT_DIR_RE.match(child.name): continue - if not (child / "dcp").is_dir(): + if not _is_complete_training_checkpoint( + child, + require_complete_marker=require_complete_marker, + ): continue try: step = _parse_step_from_dir(child) @@ -76,7 +138,27 @@ def _find_latest_checkpoint(output_dir: Path) -> Path | None: return candidates[-1][1] -def _resolve_resume_checkpoint(resume_from_checkpoint: str, *, output_dir: str) -> Path | None: +def _publish_training_checkpoint_complete(checkpoint_dir: Path) -> None: + """Atomically publish the marker that makes a training checkpoint visible.""" + marker = checkpoint_dir / _TRAINING_CHECKPOINT_COMPLETE_MARKER + temporary = checkpoint_dir / f"{_TRAINING_CHECKPOINT_COMPLETE_MARKER}.tmp-{os.getpid()}-{time.time_ns()}" + try: + with temporary.open("w", encoding="utf-8") as handle: + handle.write("complete\n") + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary, marker) + finally: + with contextlib.suppress(FileNotFoundError): + temporary.unlink() + + +def _resolve_resume_checkpoint( + resume_from_checkpoint: str, + *, + output_dir: str, + require_complete_marker: bool = False, +) -> Path | None: """Resolve a user-provided resume path to a concrete checkpoint dir. Accepted values: @@ -89,8 +171,16 @@ def _resolve_resume_checkpoint(resume_from_checkpoint: str, *, output_dir: str) if str(resume_from_checkpoint).strip().lower() == "latest": out = Path(os.path.expanduser(str(output_dir))).resolve() - latest = _find_latest_checkpoint(out) + latest = _find_latest_checkpoint( + out, + require_complete_marker=require_complete_marker, + ) if latest is None: + has_checkpoint_dirs = out.is_dir() and any(child.is_dir() and _CHECKPOINT_DIR_RE.match(child.name) + for child in out.iterdir()) + if require_complete_marker and has_checkpoint_dirs: + raise ValueError(f"No complete resumable checkpoint found under {out}; " + "refusing to start from scratch in a non-empty training namespace") logger.info( "resume_from_checkpoint='latest' but no " "checkpoints found under %s; starting from " @@ -110,10 +200,18 @@ def _resolve_resume_checkpoint(resume_from_checkpoint: str, *, output_dir: str) if path.is_dir() and _CHECKPOINT_DIR_RE.match(path.name): if not (path / "dcp").is_dir(): raise FileNotFoundError(f"Missing dcp dir under checkpoint: {path / 'dcp'}") + if not _is_complete_training_checkpoint( + path, + require_complete_marker=require_complete_marker, + ): + raise ValueError(f"Checkpoint is incomplete under the configured resume policy: {path}") return path # Treat as output_dir -> pick latest. - latest = _find_latest_checkpoint(path) + latest = _find_latest_checkpoint( + path, + require_complete_marker=require_complete_marker, + ) if latest is not None: return latest @@ -182,8 +280,20 @@ def load_state_dict( @dataclass(slots=True) class CheckpointConfig: + # Full distributed state used to resume training. save_steps: int keep_last: int + # Suppress periodic saves before this step (early checkpoints of a long + # run are rarely useful and cost ~100 GiB each). 0 disables the gate. + start_step: int = 0 + # Deployable model-only checkpoints are written for every validation event + # and are never removed by ``keep_last``. + save_inference_on_validation: bool = False + inference_role: str = "student" + inference_dtype: str = "bfloat16" + # Require the post-RNG completion marker when resolving resumable state. + # False preserves checkpoints written before that marker was introduced. + require_complete_training_checkpoint: bool = False class CheckpointManager: @@ -207,9 +317,23 @@ def __init__( self.dataloader = dataloader self.output_dir = str(output_dir) self.config = config + save_inference = bool(config.save_inference_on_validation) + inference_role = str(config.inference_role or "") + if save_inference and (not inference_role or "." in inference_role): + raise ValueError("inference_role must be a non-empty DCP key segment when inference saving is enabled") + if save_inference and str(config.inference_dtype) not in {"bfloat16", "float16", "float32"}: + raise ValueError("inference_dtype must be bfloat16, float16, or float32") + if config.require_complete_training_checkpoint: + metadata = {"config": raw_config} + if _saved_checkpoint_world_size(metadata) is None: + raise ValueError("require_complete_training_checkpoint needs a positive " + "training.distributed.num_gpus value in the saved raw config") self._callbacks = callbacks self._raw_config = raw_config + # Training-state and inference checkpoints have independent policies + # and deduplication. self._last_saved_step: int | None = None + self._last_inference_saved_step: int | None = None def _build_states(self) -> dict[str, Any]: states: dict[str, Any] = self.method.checkpoint_state() @@ -230,31 +354,53 @@ def _checkpoint_dir(self, step: int) -> Path: def _dcp_dir(self, step: int) -> Path: return self._checkpoint_dir(step) / "dcp" + def _inference_checkpoint_dir(self, step: int) -> Path: + return Path(self.output_dir) / "inference" / f"checkpoint-{step}" + + def _inference_staging_dir(self, step: int) -> Path: + return Path(self.output_dir) / ".inference-staging" / f"checkpoint-{step}" + def maybe_save(self, step: int) -> None: - save_steps = int(self.config.save_steps or 0) - if save_steps <= 0: + if step < int(self.config.start_step or 0): return - if step % save_steps != 0: + + save_steps = int(self.config.save_steps or 0) + if save_steps > 0 and step % save_steps == 0 and self._last_saved_step != step: + self.save(step) + + def maybe_save_inference(self, step: int, *, validation_scheduled: bool) -> None: + """Save the model evaluated by one scheduled validation event. + + This event-driven policy includes step-zero validation and deliberately + ignores the start gate and cadence used for resumable training state. + """ + if not validation_scheduled or not bool(self.config.save_inference_on_validation): return - if self._last_saved_step == step: + if self._last_inference_saved_step == step: return - self.save(step) + self.save_inference(step) def save_final(self, step: int) -> None: - save_steps = int(self.config.save_steps or 0) - if save_steps <= 0: - return - self.save(step) + if int(self.config.save_steps or 0) > 0 and self._last_saved_step != step: + self.save(step) def save(self, step: int) -> None: checkpoint_dir = self._checkpoint_dir(step) dcp_dir = self._dcp_dir(step) os.makedirs(dcp_dir, exist_ok=True) + # A retry may target a directory whose previous DCP save completed but + # whose RNG snapshots did not. Remove the publication marker before + # overwriting any state so strict readers can never select stale data. + if _rank() == 0: + with contextlib.suppress(FileNotFoundError): + (checkpoint_dir / _TRAINING_CHECKPOINT_COMPLETE_MARKER).unlink() + _barrier() + states = self._build_states() if _rank() == 0: logger.info( - "Saving checkpoint to %s", + "Saving resumable training checkpoint to %s", checkpoint_dir, ) self._write_metadata(checkpoint_dir, step) @@ -269,10 +415,145 @@ def save(self, step: int) -> None: self._save_rng_snapshot(checkpoint_dir) _barrier() + if _rank() == 0: + _publish_training_checkpoint_complete(checkpoint_dir) + _barrier() + self._last_saved_step = step self._cleanup_old_checkpoints() + def save_inference(self, step: int) -> None: + """Save one deployable inference checkpoint for the configured role. + + DCP is used only as a temporary, distributed staging format so FSDP2 + ranks never gather the full fp32 model into one process. Rank zero then + streams bounded tensor groups into a bf16/fp16/fp32 modular model + directory and publishes it atomically. + """ + role = str(self.config.inference_role or "student") + modules = self.method.inference_checkpoint_modules(role) + base_model_path = self.method.inference_checkpoint_base_model_path(role) + checkpoint_dir = self._inference_checkpoint_dir(step) + + already_complete: bool | None = None + existing_error: str | None = None + if _rank() == 0: + try: + from fastvideo.train.utils.inference_checkpoint import ( + validate_complete_inference_checkpoint, ) + + already_complete = (validate_complete_inference_checkpoint(checkpoint_dir, step=step) is not None) + except Exception as error: + existing_error = f"{type(error).__name__}: {error}" + if dist.is_available() and dist.is_initialized(): + complete_payload: list[Any] = [already_complete, existing_error] + dist.broadcast_object_list(complete_payload, src=0) + already_complete = bool(complete_payload[0]) + existing_error = complete_payload[1] + if existing_error is not None: + raise RuntimeError(f"Existing inference checkpoint failed validation at step {step}: {existing_error}") + if already_complete: + if _rank() == 0: + logger.info("Inference checkpoint already complete at %s; skipping", checkpoint_dir) + self._last_inference_saved_step = step + return + + staging_dir = self._inference_staging_dir(step) + dcp_dir = staging_dir / "dcp" + export_status_path = staging_dir / "export-status.json" + if _rank() == 0: + # A prior failed save is never a valid source: DCP writes + # ``.metadata`` last, and the exporter publishes independently. + shutil.rmtree(staging_dir, ignore_errors=True) + os.makedirs(dcp_dir, exist_ok=True) + _barrier() + + states = {f"roles.{role}.{module_name}": _FullModelState(module) for module_name, module in modules.items()} + if not states: + raise ValueError(f"Inference checkpoint role {role!r} exposes no modules") + + # Saving weights must not perturb the training trajectory. The regular + # resumable checkpoint intentionally snapshots its post-save RNG state; + # this model-only staging save instead restores the pre-save state. + torch_rng = torch.get_rng_state() + python_rng = random.getstate() + numpy_rng = np.random.get_state() + cuda_rng = torch.cuda.get_rng_state() if torch.cuda.is_available() else None + generator = getattr(self.method, "cuda_generator", None) + generator_rng = generator.get_state() if generator is not None else None + try: + if _rank() == 0: + logger.info("Staging inference role %s with DCP at %s", role, dcp_dir) + dcp.save(states, checkpoint_id=str(dcp_dir)) + _barrier() + finally: + torch.set_rng_state(torch_rng) + random.setstate(python_rng) + np.random.set_state(numpy_rng) + if cuda_rng is not None: + torch.cuda.set_rng_state(cuda_rng) + if generator is not None and generator_rng is not None: + generator.set_state(generator_rng) + + export_error: str | None = None + if _rank() == 0: + try: + from fastvideo.train.utils.inference_checkpoint import ( + export_inference_checkpoint, ) + + export_inference_checkpoint( + dcp_dir=dcp_dir, + output_dir=checkpoint_dir, + base_model_path=base_model_path, + role=role, + modules=modules, + dtype=str(self.config.inference_dtype), + step=step, + raw_config=self._raw_config, + ) + except Exception as error: # propagate the rank-zero failure collectively + logger.exception("Inference checkpoint export failed at step %s", step) + export_error = f"{type(error).__name__}: {error}" + status_tmp = export_status_path.with_suffix(".tmp") + status_tmp.write_text( + json.dumps({ + "complete": export_error is None, + "error": export_error + }) + "\n", + encoding="utf-8", + ) + os.replace(status_tmp, export_status_path) + else: + # Do not enter a collective while rank zero performs a multi-minute + # CPU/Lustre export: an outstanding NCCL operation can trip the + # process-group watchdog. The atomically published shared-FS result + # gives every rank the same terminal outcome before any barrier. + last_log = time.monotonic() + while not export_status_path.is_file(): + time.sleep(2.0) + now = time.monotonic() + if now - last_log >= 60.0: + logger.info("Waiting for rank-zero inference export at step %s", step) + last_log = now + try: + status = json.loads(export_status_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as error: + raise RuntimeError(f"Invalid inference export status at step {step}: {export_status_path}") from error + if status.get("complete") is not True: + export_error = str(status.get("error") or "rank-zero export failed without an error message") + if export_error is not None: + raise RuntimeError(f"Inference checkpoint export failed at step {step}: {export_error}") + + _barrier() + if _rank() == 0: + shutil.rmtree(staging_dir, ignore_errors=True) + staging_root = staging_dir.parent + with contextlib.suppress(OSError): + staging_root.rmdir() + _barrier() + self._last_inference_saved_step = step + def _write_metadata( self, checkpoint_dir: Path, @@ -328,6 +609,7 @@ def load_rng_snapshot( resolved = _resolve_resume_checkpoint( checkpoint_path, output_dir=self.output_dir, + require_complete_marker=self.config.require_complete_training_checkpoint, ) if resolved is None: return @@ -370,6 +652,7 @@ def maybe_resume(self, *, resume_from_checkpoint: str | None) -> int | None: resolved = _resolve_resume_checkpoint( resume_from_checkpoint, output_dir=self.output_dir, + require_complete_marker=self.config.require_complete_training_checkpoint, ) if resolved is None: return None @@ -398,6 +681,13 @@ def _cleanup_old_checkpoints(self) -> None: continue if not _CHECKPOINT_DIR_RE.match(child.name): continue + # In strict mode, a directory without the post-RNG publication + # marker is diagnostic debris, not one of the rolling resumable + # checkpoints. It must not consume ``keep_last`` and displace an + # older checkpoint that can actually be resumed. + if (self.config.require_complete_training_checkpoint + and not _is_complete_training_checkpoint(child, require_complete_marker=True)): + continue try: step = _parse_step_from_dir(child) except Exception: diff --git a/fastvideo/train/utils/config.py b/fastvideo/train/utils/config.py index d396b02585..5ffa4d7d51 100644 --- a/fastvideo/train/utils/config.py +++ b/fastvideo/train/utils/config.py @@ -382,6 +382,32 @@ def _build_training_config( "{'t2v', 't2va', 'text_only'}, got " f"{preprocessed_data_type!r}") + vsa_tile_size = int(vs.get("tile_size", 256) or 256) + if vsa_tile_size not in (64, 256): + raise ValueError(f"training.vsa.tile_size must be 64 or 256, got {vsa_tile_size!r}") + + save_inference_checkpoint_on_validation = require_bool( + ck, + "save_inference_checkpoint_on_validation", + default=False, + where="training.checkpoint.save_inference_checkpoint_on_validation", + ) + inference_checkpoint_role = str(ck.get("inference_checkpoint_role", "student") or "").strip() + if save_inference_checkpoint_on_validation and not inference_checkpoint_role: + raise ValueError("training.checkpoint.inference_checkpoint_role must be non-empty " + "when validation inference checkpointing is enabled") + inference_checkpoint_dtype = str(ck.get("inference_checkpoint_dtype", "bfloat16") or "bfloat16").strip().lower() + if inference_checkpoint_dtype not in {"bfloat16", "float16", "float32"}: + raise ValueError("training.checkpoint.inference_checkpoint_dtype must be one of " + "['bfloat16', 'float16', 'float32'], got " + f"{inference_checkpoint_dtype!r}") + require_complete_training_checkpoint = require_bool( + ck, + "require_complete_training_checkpoint", + default=False, + where="training.checkpoint.require_complete_training_checkpoint", + ) + return TrainingConfig( distributed=DistributedConfig( num_gpus=num_gpus, @@ -402,6 +428,7 @@ def _build_training_config( num_width=int(da.get("num_width", 0) or 0), num_latent_t=int(da.get("num_latent_t", 0) or 0), num_frames=int(da.get("num_frames", 0) or 0), + native_shape_bucketing=bool(da.get("native_shape_bucketing", False)), ), optimizer=OptimizerConfig( learning_rate=float(o.get("learning_rate", 0.0) or 0.0), @@ -420,8 +447,14 @@ def _build_training_config( checkpoint=CheckpointConfig( output_dir=str(ck.get("output_dir", "") or ""), resume_from_checkpoint=str(ck.get("resume_from_checkpoint", "") or ""), + save_inference_checkpoint_on_validation=save_inference_checkpoint_on_validation, + inference_checkpoint_role=inference_checkpoint_role or "student", + inference_checkpoint_dtype=inference_checkpoint_dtype, training_state_checkpointing_steps=int(ck.get("training_state_checkpointing_steps", 0) or 0), + require_complete_training_checkpoint=require_complete_training_checkpoint, checkpoints_total_limit=int(ck.get("checkpoints_total_limit", 0) or 0), + checkpointing_start_step=int(ck.get("checkpointing_start_step", 0) or 0), + reset_lr_on_resume=bool(ck.get("reset_lr_on_resume", False)), ), tracker=TrackerConfig( trackers=list(tr.get("trackers", []) or []), @@ -430,6 +463,7 @@ def _build_training_config( run_name=str(tr.get("run_name", "") or ""), ), vsa_sparsity=float(vs.get("sparsity", 0.0) or 0.0), + vsa_tile_size=vsa_tile_size, vsa_cache_tile_buf=bool(vs.get("cache_tile_buf", False) or False), model=ModelTrainingConfig( weighting_scheme=str(m.get("weighting_scheme", "uniform") or "uniform"), @@ -439,6 +473,9 @@ def _build_training_config( precondition_outputs=bool(m.get("precondition_outputs", False)), moba_config=dict(m.get("moba_config", {}) or {}), enable_gradient_checkpointing_type=(m.get("enable_gradient_checkpointing_type")), + allow_low_precision_master_weights=bool(m.get("allow_low_precision_master_weights", False)), + enable_torch_compile=bool(m.get("enable_torch_compile", False)), + torch_compile_kwargs=dict(m.get("torch_compile_kwargs", {}) or {}), ), pipeline_config=pipeline_config, model_path=model_path, diff --git a/fastvideo/train/utils/dataloader.py b/fastvideo/train/utils/dataloader.py index 44292fafd8..767f112f7d 100644 --- a/fastvideo/train/utils/dataloader.py +++ b/fastvideo/train/utils/dataloader.py @@ -29,6 +29,7 @@ def build_parquet_t2v_train_dataloader( drop_last=True, text_padding_length=int(text_len), seed=int(data_config.seed or 0), + native_shape_bucketing=bool(data_config.native_shape_bucketing), )) return dataloader @@ -53,4 +54,4 @@ def build_parquet_matrixgame2_train_dataloader( text_padding_length=512, seed=int(data_config.seed or 0), )) - return dataloader \ No newline at end of file + return dataloader diff --git a/fastvideo/train/utils/inference_checkpoint.py b/fastvideo/train/utils/inference_checkpoint.py new file mode 100644 index 0000000000..fc4c1f1925 --- /dev/null +++ b/fastvideo/train/utils/inference_checkpoint.py @@ -0,0 +1,627 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Bounded-memory export of model-only DCP state for inference. + +The modular trainer stores resumable state with PyTorch Distributed +Checkpoint (DCP), while FastVideo inference consumes a Diffusers-style model +directory. This module bridges those formats without gathering a full model +in memory: rank 0 reads one bounded shard at a time from an already-complete +role/module-only DCP checkpoint and publishes an immutable inference directory +with an atomic rename. + +The manager-facing entry point is :func:`export_inference_checkpoint`; the +lower-level :func:`export_inference_checkpoint_from_dcp` exposes the same +rank-0-only conversion with a run-root destination. The checkpoint manager +owns distributed coordination around the temporary DCP save and local export. +""" + +from __future__ import annotations + +import json +import os +import re +import shutil +import tempfile +from collections.abc import Mapping +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import torch +import torch.distributed as dist +import torch.distributed.checkpoint as dcp +from safetensors import safe_open +from safetensors.torch import save_file +from torch.distributed.checkpoint import FileSystemReader + +DEFAULT_MAX_SHARD_SIZE_BYTES = 5 * 1024**3 +_FORMAT_VERSION = 1 +_ALLOWED_NATIVE_EXTRA_KEYS = (re.compile(r"(?:^|\.)attn\.to_gate_compress\.weight$"), ) + + +class InferenceCheckpointExportError(RuntimeError): + """The model-only checkpoint could not be exported safely.""" + + +class UnsupportedMergedReverseMappingError(InferenceCheckpointExportError): + """A fused training parameter would need to be split for inference.""" + + +def _is_allowed_native_extra(key: str) -> bool: + return any(pattern.search(key) for pattern in _ALLOWED_NATIVE_EXTRA_KEYS) + + +@dataclass(frozen=True, slots=True) +class _TensorPlan: + checkpoint_key: str + internal_key: str + output_key: str + shape: tuple[int, ...] + source_dtype: torch.dtype + output_dtype: torch.dtype + output_nbytes: int + + +def _rank() -> int: + if dist.is_available() and dist.is_initialized(): + return int(dist.get_rank()) + return 0 + + +def _normalize_output_dtype(dtype: torch.dtype | str) -> torch.dtype: + if isinstance(dtype, str): + name = dtype.removeprefix("torch.") + resolved = getattr(torch, name, None) + if not isinstance(resolved, torch.dtype): + raise ValueError(f"Unsupported inference checkpoint dtype: {dtype!r}") + dtype = resolved + if not isinstance(dtype, torch.dtype): + raise TypeError("dtype must be a torch.dtype or torch dtype name") + if not torch.empty((), dtype=dtype).is_floating_point(): + raise ValueError(f"Inference checkpoint dtype must be floating point, got {dtype}") + return dtype + + +def _resolve_dcp_dir(checkpoint: str | os.PathLike[str]) -> Path: + path = Path(checkpoint).expanduser().resolve() + if path.name != "dcp" and (path / "dcp").is_dir(): + path = path / "dcp" + if not path.is_dir(): + raise FileNotFoundError(f"Inference checkpoint DCP directory not found: {path}") + if not (path / ".metadata").is_file(): + raise FileNotFoundError(f"Incomplete inference checkpoint DCP (missing .metadata): {path}") + return path + + +def validate_complete_inference_checkpoint(path: Path, *, step: int) -> Path | None: + """Validate a published inference checkpoint without loading its tensors.""" + if not (path.exists() or path.is_symlink()): + return None + complete_path = path / ".complete" + metadata_path = path / "metadata.json" + if not complete_path.is_file() or not metadata_path.is_file(): + raise InferenceCheckpointExportError(f"Refusing to overwrite incomplete inference checkpoint: {path}") + try: + complete_marker = complete_path.read_text(encoding="utf-8") + except OSError as exc: + raise InferenceCheckpointExportError(f"Cannot read inference completion marker: {complete_path}") from exc + if complete_marker != "complete\n": + raise InferenceCheckpointExportError(f"Invalid inference completion marker: {complete_path}") + try: + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise InferenceCheckpointExportError( + f"Invalid metadata for existing inference checkpoint: {metadata_path}") from exc + if (metadata.get("format_version") != _FORMAT_VERSION or metadata.get("kind") != "inference" + or metadata.get("step") != step): + raise InferenceCheckpointExportError( + f"Existing inference checkpoint metadata does not match step={step}: {metadata_path}") + role = metadata.get("role") + dtype = metadata.get("dtype") + if not isinstance(role, str) or not role or "." in role or dtype not in {"bfloat16", "float16", "float32"}: + raise InferenceCheckpointExportError(f"Inference checkpoint role/dtype metadata is invalid: {metadata_path}") + module_name = str(metadata.get("module") or "") + if not module_name or "/" in module_name or "\\" in module_name or module_name in {".", ".."}: + raise InferenceCheckpointExportError( + f"Inference checkpoint has an invalid module name {module_name!r}: {metadata_path}") + module_dir = path / module_name + index_path = module_dir / "diffusion_pytorch_model.safetensors.index.json" + if not index_path.is_file(): + raise InferenceCheckpointExportError(f"Inference checkpoint is missing its module index: {index_path}") + if not (module_dir / "config.json").is_file(): + raise InferenceCheckpointExportError(f"Inference checkpoint is missing its module config: {module_dir}") + if not any((path / name).is_file() for name in ("model_index.json", "modular_model_index.json")): + raise InferenceCheckpointExportError(f"Inference checkpoint is missing its model index: {path}") + try: + index = json.loads(index_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise InferenceCheckpointExportError(f"Invalid inference checkpoint index: {index_path}") from exc + weight_map = index.get("weight_map") + if not isinstance(weight_map, dict) or not weight_map: + raise InferenceCheckpointExportError(f"Inference checkpoint index has no weight map: {index_path}") + expected_by_shard: dict[str, set[str]] = {} + for key, filename in weight_map.items(): + if not isinstance(key, str) or not key or not isinstance(filename, str) or not filename: + raise InferenceCheckpointExportError(f"Invalid key/shard entry in {index_path}: {key!r} -> {filename!r}") + if Path(filename).name != filename: + raise InferenceCheckpointExportError( + f"Inference checkpoint index contains an invalid shard path: {filename!r}") + expected_by_shard.setdefault(filename, set()).add(key) + actual_shards = {shard.name for shard in module_dir.glob("*.safetensors") if shard.is_file()} + expected_shards = set(expected_by_shard) + if actual_shards != expected_shards: + raise InferenceCheckpointExportError( + f"Inference checkpoint shards differ from its index under {module_dir}: " + f"missing={sorted(expected_shards - actual_shards)} extra={sorted(actual_shards - expected_shards)}") + shard_sizes: list[int] = [] + output_shapes: dict[str, tuple[int, ...]] = {} + for filename in sorted(expected_by_shard): + expected_keys = expected_by_shard[filename] + shard_path = module_dir / filename + try: + with safe_open(str(shard_path), framework="pt", device="cpu") as handle: + actual_keys = set(handle.keys()) + for key in actual_keys: + output_shapes[key] = tuple(int(dim) for dim in handle.get_slice(key).get_shape()) + except Exception as exc: + raise InferenceCheckpointExportError(f"Cannot read inference checkpoint shard: {shard_path}") from exc + if actual_keys != expected_keys: + raise InferenceCheckpointExportError( + f"Inference checkpoint shard keys differ from its index for {shard_path}: " + f"missing={sorted(expected_keys - actual_keys)[:10]} extra={sorted(actual_keys - expected_keys)[:10]}") + shard_sizes.append(shard_path.stat().st_size) + if int(metadata.get("tensor_count", -1)) != len(weight_map): + raise InferenceCheckpointExportError( + f"Inference checkpoint tensor_count does not match its index: {metadata_path}") + if int(metadata.get("shard_count", -1)) != len(expected_shards): + raise InferenceCheckpointExportError( + f"Inference checkpoint shard_count does not match its index: {metadata_path}") + logical_shard_sizes = metadata.get("shard_sizes") + total_size = metadata.get("total_size") + max_shard_size = metadata.get("max_shard_size_bytes") + index_metadata = index.get("metadata") + index_total_size = index_metadata.get("total_size") if isinstance(index_metadata, dict) else None + if (not isinstance(logical_shard_sizes, list) or len(logical_shard_sizes) != len(expected_shards) + or any(not isinstance(size, int) or size < 0 + for size in logical_shard_sizes) or not isinstance(max_shard_size, int) or max_shard_size <= 0 + or any(size > max_shard_size for size in logical_shard_sizes) or not isinstance(total_size, int) + or total_size != sum(logical_shard_sizes) or index_total_size != total_size): + raise InferenceCheckpointExportError( + f"Inference checkpoint logical shard sizes are inconsistent: {metadata_path}") + recorded_shard_sizes = metadata.get("shard_file_sizes") + if not isinstance(recorded_shard_sizes, list) or recorded_shard_sizes != shard_sizes: + raise InferenceCheckpointExportError( + f"Inference checkpoint shard_file_sizes do not match files on disk: {metadata_path}") + + base_model_dir = metadata.get("base_model_dir") + if not isinstance(base_model_dir, str) or not base_model_dir: + raise InferenceCheckpointExportError(f"Inference checkpoint has no base_model_dir: {metadata_path}") + base_shapes = _component_tensor_shapes(Path(base_model_dir) / module_name) + for key, shape in output_shapes.items(): + expected_shape = base_shapes.get(key) + if expected_shape is None: + if not _is_allowed_native_extra(key): + raise InferenceCheckpointExportError( + f"Inference checkpoint contains unknown non-native tensor {key!r}: {path}") + elif shape != expected_shape: + raise InferenceCheckpointExportError( + f"Inference checkpoint tensor {key!r} shape {shape} != base transformer shape {expected_shape}") + missing_base_keys = set(base_shapes) - set(output_shapes) + if missing_base_keys: + raise InferenceCheckpointExportError( + f"Inference checkpoint is missing {len(missing_base_keys)} base transformer tensors; " + f"first={sorted(missing_base_keys)[:10]}") + return path + + +def _component_tensor_shapes(module_dir: Path) -> dict[str, tuple[int, ...]]: + """Read the base component's exact key/shape contract from safetensors headers.""" + index_candidates = ( + module_dir / "diffusion_pytorch_model.safetensors.index.json", + module_dir / "model.safetensors.index.json", + ) + index_path = next((path for path in index_candidates if path.is_file()), None) + expected_files: set[str] | None = None + indexed_keys: set[str] | None = None + if index_path is not None: + try: + index = json.loads(index_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise InferenceCheckpointExportError(f"Invalid base transformer index: {index_path}") from exc + weight_map = index.get("weight_map") + if not isinstance(weight_map, dict) or not weight_map: + raise InferenceCheckpointExportError(f"Base transformer index has no weight map: {index_path}") + indexed_keys = set(weight_map) + expected_files = {str(filename) for filename in weight_map.values()} + + files = sorted(module_dir.glob("*.safetensors")) + if expected_files is not None: + files = [module_dir / filename for filename in sorted(expected_files)] + if not files or any(not path.is_file() for path in files): + raise InferenceCheckpointExportError(f"Base transformer safetensors are incomplete under {module_dir}") + + shapes: dict[str, tuple[int, ...]] = {} + for path in files: + try: + with safe_open(str(path), framework="pt", device="cpu") as handle: + # ``safe_open`` exposes ``keys()`` but is not itself iterable. + for key in handle.keys(): # noqa: SIM118 + if key in shapes: + raise InferenceCheckpointExportError(f"Duplicate base transformer tensor key {key!r}") + shapes[key] = tuple(int(dim) for dim in handle.get_slice(key).get_shape()) + except InferenceCheckpointExportError: + raise + except Exception as exc: + raise InferenceCheckpointExportError(f"Cannot read base transformer safetensors: {path}") from exc + if indexed_keys is not None and set(shapes) != indexed_keys: + raise InferenceCheckpointExportError( + f"Base transformer index/header mismatch under {module_dir}: " + f"missing={sorted(indexed_keys - set(shapes))[:10]} extra={sorted(set(shapes) - indexed_keys)[:10]}") + return shapes + + +def _mapping_output_key( + internal_key: str, + reverse_mapping: Mapping[str, Any], +) -> str: + entry = reverse_mapping.get(internal_key) + if entry is None: + # FastVideo-native additions such as MiniMax-H3's learned + # ``attn.to_gate_compress`` VSA parameters intentionally have no key in + # the base checkpoint. The component loader accepts their native name. + return internal_key + if not isinstance(entry, tuple | list) or len(entry) != 3: + raise InferenceCheckpointExportError(f"Invalid reverse mapping for {internal_key!r}: expected " + "(output_key, merge_index, num_params_to_merge)") + output_key, merge_index, num_params_to_merge = entry + if merge_index is not None or num_params_to_merge not in (None, 1): + raise UnsupportedMergedReverseMappingError(f"Cannot stream merged reverse mapping for {internal_key!r}: " + f"output_key={output_key!r}, merge_index={merge_index!r}, " + f"num_params_to_merge={num_params_to_merge!r}. " + "A model-specific split exporter is required.") + if not isinstance(output_key, str) or not output_key: + raise InferenceCheckpointExportError( + f"Invalid output key in reverse mapping for {internal_key!r}: {output_key!r}") + return output_key + + +def _tensor_output_dtype(source_dtype: torch.dtype, configured_dtype: torch.dtype) -> torch.dtype: + if torch.empty((), dtype=source_dtype).is_floating_point(): + return configured_dtype + return source_dtype + + +def _build_tensor_plan( + *, + dcp_dir: Path, + state_prefix: str, + reverse_mapping: Mapping[str, Any], + base_shapes: Mapping[str, tuple[int, ...]], + output_dtype: torch.dtype, + max_shard_size_bytes: int, +) -> list[_TensorPlan]: + metadata = FileSystemReader(str(dcp_dir)).read_metadata() + plans: list[_TensorPlan] = [] + output_keys: set[str] = set() + + for checkpoint_key in sorted(metadata.state_dict_metadata): + if not checkpoint_key.startswith(state_prefix): + continue + tensor_metadata = metadata.state_dict_metadata[checkpoint_key] + properties = getattr(tensor_metadata, "properties", None) + shape = getattr(tensor_metadata, "size", None) + source_dtype = getattr(properties, "dtype", None) + if shape is None or not isinstance(source_dtype, torch.dtype): + raise InferenceCheckpointExportError(f"Inference state contains a non-tensor value at {checkpoint_key!r}; " + "safetensors exports support tensors only") + + internal_key = checkpoint_key[len(state_prefix):] + if not internal_key: + raise InferenceCheckpointExportError(f"Empty module key under DCP prefix {state_prefix!r}") + output_key = _mapping_output_key(internal_key, reverse_mapping) + tensor_shape = tuple(int(dim) for dim in shape) + expected_shape = base_shapes.get(output_key) + if expected_shape is None: + if internal_key in reverse_mapping or not _is_allowed_native_extra(output_key): + raise InferenceCheckpointExportError( + f"Inference tensor {internal_key!r} maps to unknown base key {output_key!r}; " + "only MiniMax-H3 attn.to_gate_compress.weight may be exported as a native extra") + elif tensor_shape != expected_shape: + raise InferenceCheckpointExportError( + f"Inference tensor {output_key!r} shape {tensor_shape} != base transformer shape {expected_shape}") + if output_key in output_keys: + raise InferenceCheckpointExportError(f"Reverse mapping produces duplicate inference key {output_key!r}") + output_keys.add(output_key) + + tensor_output_dtype = _tensor_output_dtype(source_dtype, output_dtype) + numel = 1 + for dim in tensor_shape: + numel *= dim + output_nbytes = numel * torch.empty((), dtype=tensor_output_dtype).element_size() + if output_nbytes > max_shard_size_bytes: + raise InferenceCheckpointExportError( + f"Tensor {checkpoint_key!r} requires {output_nbytes} bytes after casting, " + f"which exceeds max_shard_size_bytes={max_shard_size_bytes}") + plans.append( + _TensorPlan( + checkpoint_key=checkpoint_key, + internal_key=internal_key, + output_key=output_key, + shape=tensor_shape, + source_dtype=source_dtype, + output_dtype=tensor_output_dtype, + output_nbytes=output_nbytes, + )) + + if not plans: + raise InferenceCheckpointExportError(f"No tensor keys found under DCP prefix {state_prefix!r} in {dcp_dir}") + missing_base_keys = set(base_shapes) - output_keys + if missing_base_keys: + raise InferenceCheckpointExportError( + f"Inference checkpoint is missing {len(missing_base_keys)} base transformer tensors; " + f"first={sorted(missing_base_keys)[:10]}") + return plans + + +def _group_shards(plans: list[_TensorPlan], max_shard_size_bytes: int) -> list[list[_TensorPlan]]: + shards: list[list[_TensorPlan]] = [] + current: list[_TensorPlan] = [] + current_nbytes = 0 + for plan in plans: + if current and current_nbytes + plan.output_nbytes > max_shard_size_bytes: + shards.append(current) + current = [] + current_nbytes = 0 + current.append(plan) + current_nbytes += plan.output_nbytes + if current: + shards.append(current) + return shards + + +def _prepare_model_layout(temp_dir: Path, base_model_dir: Path, module_name: str) -> Path: + if not any((base_model_dir / name).is_file() for name in ("model_index.json", "modular_model_index.json")): + raise FileNotFoundError( + f"Base model directory has no model_index.json or modular_model_index.json: {base_model_dir}") + base_module_dir = base_model_dir / module_name + base_config = base_module_dir / "config.json" + if not base_config.is_file(): + raise FileNotFoundError(f"Base model component config not found: {base_config}") + + module_dir = temp_dir / module_name + module_dir.mkdir(parents=True) + shutil.copy2(base_config, module_dir / "config.json") + + reserved = {module_name, "metadata.json", ".complete"} + for entry in sorted(base_model_dir.iterdir(), key=lambda item: item.name): + if entry.name in reserved or entry.name == ".cache" or entry.name.startswith(".git"): + continue + target = temp_dir / entry.name + target.symlink_to(entry.resolve(), target_is_directory=entry.is_dir()) + return module_dir + + +def _write_tensor_shards( + *, + dcp_dir: Path, + module_dir: Path, + shards: list[list[_TensorPlan]], +) -> tuple[dict[str, str], int, list[int], list[int]]: + weight_map: dict[str, str] = {} + total_size = 0 + shard_sizes: list[int] = [] + shard_file_sizes: list[int] = [] + shard_count = len(shards) + + for shard_index, shard in enumerate(shards, start=1): + filename = (f"diffusion_pytorch_model-{shard_index:05d}-of-{shard_count:05d}.safetensors") + state = {plan.checkpoint_key: torch.empty(plan.shape, dtype=plan.source_dtype, device="cpu") for plan in shard} + # This export is deliberately rank-local. DCP assembles the full CPU + # tensors from its storage shards without using the live process group. + dcp.load(state, checkpoint_id=str(dcp_dir), no_dist=True) + + output_tensors: dict[str, torch.Tensor] = {} + shard_size = 0 + for plan in shard: + tensor = state[plan.checkpoint_key] + if tensor.is_floating_point(): + tensor = tensor.to(dtype=plan.output_dtype) + output_tensors[plan.output_key] = tensor.detach().cpu().contiguous() + weight_map[plan.output_key] = filename + shard_size += plan.output_nbytes + save_file(output_tensors, module_dir / filename) + total_size += shard_size + shard_sizes.append(shard_size) + shard_file_sizes.append((module_dir / filename).stat().st_size) + del output_tensors, state + + index = { + "metadata": { + "total_size": total_size + }, + "weight_map": weight_map, + } + index_path = module_dir / "diffusion_pytorch_model.safetensors.index.json" + index_path.write_text(json.dumps(index, indent=2, sort_keys=True) + "\n", encoding="utf-8") + return weight_map, total_size, shard_sizes, shard_file_sizes + + +def export_inference_checkpoint_from_dcp( + *, + dcp_checkpoint: str | os.PathLike[str], + output_dir: str | os.PathLike[str], + step: int, + module: torch.nn.Module, + base_model_dir: str | os.PathLike[str], + role: str = "student", + module_name: str = "transformer", + dtype: torch.dtype | str = torch.bfloat16, + max_shard_size_bytes: int = DEFAULT_MAX_SHARD_SIZE_BYTES, + raw_config: Mapping[str, Any] | None = None, +) -> Path: + """Export one role/module-only DCP as an immutable inference checkpoint. + + ``CheckpointManager`` is expected to call this function on rank 0 after a + collective model-only DCP save has completed. ``dcp_checkpoint`` may name + that DCP directory directly or its parent containing ``dcp/``. The source + is never modified or removed. + + Floating tensors are cast into ``dtype`` in independent CPU buffers; + integer and boolean tensors retain their source dtype. Parameter names are + converted with ``module.reverse_param_names_mapping``. Unmapped names are + retained for FastVideo-native inference parameters such as MiniMax-H3 VSA + gates. Merged mappings fail rather than silently writing incompatible + weights. + + The completed model is atomically renamed to + ``/inference/checkpoint-``. A valid existing completed + export is returned unchanged, making retries idempotent. + """ + + if _rank() != 0: + raise InferenceCheckpointExportError("export_inference_checkpoint_from_dcp is rank-0-only; " + "the checkpoint manager must coordinate other ranks") + if isinstance(step, bool) or not isinstance(step, int) or step < 0: + raise ValueError(f"step must be a non-negative integer, got {step!r}") + if not role or "." in role: + raise ValueError(f"role must be a non-empty DCP key segment, got {role!r}") + if not module_name or "." in module_name: + raise ValueError(f"module_name must be a non-empty DCP key segment, got {module_name!r}") + if not 0 < max_shard_size_bytes <= DEFAULT_MAX_SHARD_SIZE_BYTES: + raise ValueError("max_shard_size_bytes must be in " + f"[1, {DEFAULT_MAX_SHARD_SIZE_BYTES}], got {max_shard_size_bytes}") + + output_dtype = _normalize_output_dtype(dtype) + run_output_dir = Path(output_dir).expanduser().resolve() + inference_root = run_output_dir / "inference" + final_dir = inference_root / f"checkpoint-{step}" + existing = validate_complete_inference_checkpoint(final_dir, step=step) + if existing is not None: + return existing + + dcp_dir = _resolve_dcp_dir(dcp_checkpoint) + base_dir = Path(base_model_dir).expanduser().resolve() + base_shapes = _component_tensor_shapes(base_dir / module_name) + reverse_mapping = getattr(module, "reverse_param_names_mapping", {}) + if reverse_mapping is None: + reverse_mapping = {} + if not isinstance(reverse_mapping, Mapping): + raise InferenceCheckpointExportError("module.reverse_param_names_mapping must be a mapping") + + state_prefix = f"roles.{role}.{module_name}." + plans = _build_tensor_plan( + dcp_dir=dcp_dir, + state_prefix=state_prefix, + reverse_mapping=reverse_mapping, + base_shapes=base_shapes, + output_dtype=output_dtype, + max_shard_size_bytes=max_shard_size_bytes, + ) + shards = _group_shards(plans, max_shard_size_bytes) + + inference_root.mkdir(parents=True, exist_ok=True) + temp_dir = Path(tempfile.mkdtemp( + prefix=f".checkpoint-{step}.tmp-", + dir=str(inference_root), + )) + try: + module_dir = _prepare_model_layout(temp_dir, base_dir, module_name) + weight_map, total_size, shard_sizes, shard_file_sizes = _write_tensor_shards( + dcp_dir=dcp_dir, + module_dir=module_dir, + shards=shards, + ) + metadata = { + "format_version": _FORMAT_VERSION, + "kind": "inference", + "step": step, + "role": role, + "module": module_name, + "dtype": str(output_dtype).removeprefix("torch."), + "base_model_dir": str(base_dir), + "tensor_count": len(weight_map), + "total_size": total_size, + "shard_count": len(shards), + "shard_sizes": shard_sizes, + "shard_file_sizes": shard_file_sizes, + "max_shard_size_bytes": max_shard_size_bytes, + } + if raw_config is not None: + metadata["config"] = raw_config + (temp_dir / "metadata.json").write_text( + json.dumps(metadata, indent=2, sort_keys=True) + "\n", + encoding="utf-8", + ) + # Written last inside the private temp directory. The subsequent + # same-filesystem rename publishes the complete tree in one operation. + (temp_dir / ".complete").write_text("complete\n", encoding="utf-8") + try: + temp_dir.rename(final_dir) + except FileExistsError as exc: + raise InferenceCheckpointExportError(f"Inference checkpoint appeared concurrently: {final_dir}") from exc + except BaseException: + shutil.rmtree(temp_dir, ignore_errors=True) + raise + validated = validate_complete_inference_checkpoint(final_dir, step=step) + if validated is None: + raise InferenceCheckpointExportError(f"Published inference checkpoint disappeared: {final_dir}") + return validated + + +def export_inference_checkpoint( + *, + dcp_dir: str | os.PathLike[str], + output_dir: str | os.PathLike[str], + base_model_path: str | os.PathLike[str], + role: str, + modules: Mapping[str, torch.nn.Module], + dtype: torch.dtype | str, + step: int, + raw_config: Mapping[str, Any] | None = None, +) -> Path: + """CheckpointManager adapter for one deployable role/module checkpoint. + + ``output_dir`` is the manager's fully resolved target, + ``/inference/checkpoint-``. The temporary DCP may contain only + the role/module state selected by ``modules``. Multi-component deployment + is intentionally rejected until its component layout and atomicity + contract are defined. + + ``raw_config`` is persisted in ``metadata.json`` when supplied, matching + resumable checkpoint provenance. The ephemeral DCP staging path is not + persisted because the manager removes it after a successful export. + """ + + if len(modules) != 1: + raise InferenceCheckpointExportError("Inference checkpoint export currently supports exactly one module; " + f"got {sorted(modules)}") + module_name, module = next(iter(modules.items())) + if not isinstance(module, torch.nn.Module): + raise TypeError(f"Inference checkpoint module {module_name!r} must be a torch.nn.Module") + + final_dir = Path(output_dir).expanduser().resolve() + expected_name = f"checkpoint-{step}" + if final_dir.name != expected_name or final_dir.parent.name != "inference": + raise ValueError("CheckpointManager output_dir must be " + f"/inference/{expected_name}, got {final_dir}") + run_output_dir = final_dir.parent.parent + return export_inference_checkpoint_from_dcp( + dcp_checkpoint=dcp_dir, + output_dir=run_output_dir, + step=step, + module=module, + base_model_dir=base_model_path, + role=role, + module_name=module_name, + dtype=dtype, + raw_config=raw_config, + ) + + +__all__ = [ + "DEFAULT_MAX_SHARD_SIZE_BYTES", + "InferenceCheckpointExportError", + "UnsupportedMergedReverseMappingError", + "export_inference_checkpoint", + "export_inference_checkpoint_from_dcp", + "validate_complete_inference_checkpoint", +] diff --git a/fastvideo/train/utils/moduleloader.py b/fastvideo/train/utils/moduleloader.py index 822e41c17a..d2ecf00af6 100644 --- a/fastvideo/train/utils/moduleloader.py +++ b/fastvideo/train/utils/moduleloader.py @@ -5,12 +5,15 @@ import os from contextlib import nullcontext from typing import Any, TYPE_CHECKING +from collections.abc import Callable import torch from fastvideo.attention.selector import ( + NO_REQUEST, _component_attention_backend_scope, coerce_attn_backend, + component_attention_backend, ) from fastvideo.configs.pipelines.base import PipelineConfig from fastvideo.fastvideo_args import ExecutionMode, TrainingArgs @@ -60,7 +63,11 @@ def _make_training_args( text_encoder_cpu_offload=False, image_encoder_cpu_offload=False, use_fsdp_inference=False, - enable_torch_compile=False, + enable_torch_compile=tc.model.enable_torch_compile, + # Modular stack opts into regional fullgraph compile; the legacy + # stack keeps whole-model torch.compile semantics (default False). + regional_compile=True, + torch_compile_kwargs=tc.model.torch_compile_kwargs, ) @@ -73,8 +80,19 @@ def make_inference_args( args = _make_training_args(tc, model_path=model_path) args.inference_mode = True args.mode = ExecutionMode.INFERENCE - args.dit_cpu_offload = True + # Never CPU-offload the DiT here: validation samples with the LIVE + # training transformer, and the denoising stages' post-sampling + # ``transformer.to("cpu")`` strands the FSDP-sharded training params on + # CPU — the next training backward then dies assigning CUDA grads to + # CPU tensors (no-grad forwards keep working off gathered buffers, + # which is why it surfaces steps later). + args.dit_cpu_offload = False + # Validation must sample at the training attention contract: both the + # sparsity AND the tile geometry. Leaving the tile size at its + # FastVideoArgs default silently validates a tile-64-trained student at + # tile 256 (v8 shipped 2400 steps of validation that way). args.VSA_sparsity = tc.vsa_sparsity + args.VSA_tile_size = tc.vsa_tile_size return args @@ -92,6 +110,8 @@ def load_module_from_path( override_transformer_cls_name: str | None = None, transformer_override_safetensor: str | None = None, attention_backend: AttentionBackendEnum | str | None = None, + construction_precision: str | None = None, + pre_fsdp_transform: Callable[[torch.nn.Module], torch.nn.Module] | None = None, ) -> torch.nn.Module: """Load one pipeline component with its role-scoped attention policy. @@ -104,6 +124,12 @@ def load_module_from_path( scoped to this load call. """ fastvideo_args: Any = _make_training_args(training_config, model_path=model_path) + original_dit_precision = fastvideo_args.pipeline_config.dit_precision + if construction_precision is not None: + # A frozen role does not need FP32 optimizer masters. Its FSDP forward + # already casts parameters to BF16, so constructing/storing that role + # in BF16 removes memory with no change to the actual teacher compute. + fastvideo_args.pipeline_config.dit_precision = str(construction_precision) local_model_path = maybe_download_model(model_path) config = verify_model_config_and_directory(local_model_path) @@ -130,6 +156,12 @@ def load_module_from_path( if transformer_override_safetensor: fastvideo_args.init_weights_from_safetensors = str(transformer_override_safetensor) + if pre_fsdp_transform is not None: + if module_type != "transformer": + raise ValueError("pre_fsdp_transform can only be set when loading " + f"a transformer, got module_type={module_type!r}") + fastvideo_args._pre_fsdp_transform = pre_fsdp_transform + if attention_backend is not None and module_type != "transformer": raise ValueError("attention_backend can only be set when loading " f"a transformer, got module_type={module_type!r}") @@ -145,15 +177,31 @@ def load_module_from_path( # Attention implementations are bound while transformer layers are # constructed. Scope the override to this one role so student, # teacher, and critic can use independent backends in one process. - with attention_context: - module = PipelineComponentLoader.load_module( - module_name=module_type, - component_model_path=component_path, - transformers_or_diffusers=(transformers_or_diffusers), - fastvideo_args=fastvideo_args, - ) + try: + with attention_context: + module = PipelineComponentLoader.load_module( + module_name=module_type, + component_model_path=component_path, + transformers_or_diffusers=(transformers_or_diffusers), + fastvideo_args=fastvideo_args, + ) + finally: + # _make_training_args intentionally shares the resolved pipeline + # config. Do not leak a role-local construction choice to later roles. + fastvideo_args.pipeline_config.dit_precision = original_dit_precision if not isinstance(module, torch.nn.Module): raise TypeError(f"Loaded {module_type!r} is not a " f"torch.nn.Module: {type(module)}") + if resolved_attention_backend is not None: + receipt = component_attention_backend(module) + if receipt is NO_REQUEST: + raise RuntimeError(f"Loaded {module_type!r} from {model_path!r} did not record its " + f"requested attention backend {resolved_attention_backend.name}. " + "The component loader must stamp the construction decision on " + "module.config._resolved_attention_backend.") + if receipt is not resolved_attention_backend: + raise RuntimeError(f"Loaded {module_type!r} from {model_path!r} requested attention " + f"backend {resolved_attention_backend.name}, but recorded " + f"{receipt.name}.") return module diff --git a/fastvideo/train/utils/optimizer.py b/fastvideo/train/utils/optimizer.py index 43a79d98df..b036fd913d 100644 --- a/fastvideo/train/utils/optimizer.py +++ b/fastvideo/train/utils/optimizer.py @@ -18,6 +18,75 @@ ) +class AdamWBeta1Zero(torch.optim.Optimizer): + """AdamW specialized for ``beta1 == 0``: identical update, no ``exp_avg``. + + With ``beta1 = 0`` Adam's first moment reduces to the raw gradient + (``m_t = g_t``, bias correction 1), so the buffer only doubles optimizer + state for nothing — one full parameter-sized tensor per model. The op + sequence below mirrors ``torch.optim.AdamW``'s single-tensor path + exactly, so the parameter trajectory is bitwise-equivalent to + ``AdamW(betas=(0.0, beta2))``. + """ + + def __init__( + self, + params, + lr: float, + beta2: float, + eps: float = 1e-8, + weight_decay: float = 0.0, + ) -> None: + if not 0.0 <= beta2 < 1.0: + raise ValueError(f"Invalid beta2: {beta2}") + defaults = dict(lr=float(lr), beta2=float(beta2), eps=float(eps), weight_decay=float(weight_decay)) + super().__init__(params, defaults) + + @torch.no_grad() + def step(self, closure=None): + loss = None + if closure is not None: + with torch.enable_grad(): + loss = closure() + for group in self.param_groups: + lr = group["lr"] + beta2 = group["beta2"] + eps = group["eps"] + weight_decay = group["weight_decay"] + for p in group["params"]: + if p.grad is None: + continue + grad = p.grad + state = self.state[p] + if not state: + state["step"] = 0 + # Low-precision params must not starve the second moment: + # bf16 v loses g^2 increments below ~0.4% of its running + # magnitude. Keep v in fp32 whenever p is not fp32 (the + # fp32-master path is unchanged and stays bitwise-equal + # to torch.optim.AdamW). + if p.dtype == torch.float32: + state["exp_avg_sq"] = torch.zeros_like(p, memory_format=torch.preserve_format) + else: + state["exp_avg_sq"] = torch.zeros_like(p, dtype=torch.float32) + state["step"] += 1 + exp_avg_sq = state["exp_avg_sq"] + p.mul_(1 - lr * weight_decay) + if p.dtype == torch.float32: + exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2) + bias_correction2_sqrt = (1 - beta2**state["step"])**0.5 + denom = (exp_avg_sq.sqrt() / bias_correction2_sqrt).add_(eps) + p.addcdiv_(grad, denom, value=-lr) + else: + grad_f = grad.float() + exp_avg_sq.mul_(beta2).addcmul_(grad_f, grad_f, value=1 - beta2) + bias_correction2_sqrt = (1 - beta2**state["step"])**0.5 + denom = (exp_avg_sq.sqrt() / bias_correction2_sqrt).add_(eps) + update = grad_f.div_(denom).mul_(-lr) + p.add_(update.to(p.dtype)) + return loss + + def build_optimizer_and_scheduler( *, params: list[torch.nn.Parameter], @@ -36,13 +105,24 @@ def build_optimizer_and_scheduler( raise ValueError("No trainable parameters passed to " "build_optimizer_and_scheduler") - optimizer = torch.optim.AdamW( - params, - lr=float(learning_rate), - betas=betas, - weight_decay=float(optimizer_config.weight_decay), - eps=1e-8, - ) + if float(betas[0]) == 0.0: + # beta1 == 0 makes Adam's first-moment buffer redundant; skip it to + # save one parameter-sized state tensor per trainable model. + optimizer: torch.optim.Optimizer = AdamWBeta1Zero( + params, + lr=float(learning_rate), + beta2=float(betas[1]), + weight_decay=float(optimizer_config.weight_decay), + eps=1e-8, + ) + else: + optimizer = torch.optim.AdamW( + params, + lr=float(learning_rate), + betas=betas, + weight_decay=float(optimizer_config.weight_decay), + eps=1e-8, + ) scheduler = get_scheduler( str(scheduler_name), diff --git a/fastvideo/train/utils/tracking.py b/fastvideo/train/utils/tracking.py index 91e06889be..83dc2f6699 100644 --- a/fastvideo/train/utils/tracking.py +++ b/fastvideo/train/utils/tracking.py @@ -6,6 +6,7 @@ from typing import Any, TYPE_CHECKING from fastvideo.distributed import get_world_group +from fastvideo.logger import init_logger from fastvideo.training.trackers import ( initialize_trackers, Trackers, @@ -17,6 +18,25 @@ TrackerConfig, ) +logger = init_logger(__name__) + + +def _wandb_usable() -> bool: + """True when wandb can actually start a run (importable + credentials).""" + try: + import wandb + except Exception: + return False + if os.environ.get("WANDB_API_KEY"): + return True + if os.environ.get("WANDB_MODE") in ("offline", "disabled"): + return True + try: + # Covers ~/.netrc logins from a prior `wandb login`. + return wandb.api.api_key is not None + except Exception: + return False + def build_tracker( tracker_config: TrackerConfig, @@ -33,6 +53,11 @@ def build_tracker( trackers.append(Trackers.WANDB.value) if world_group.rank != 0: trackers = [] + if Trackers.WANDB.value in trackers and not _wandb_usable(): + logger.warning("wandb tracking requested but wandb is not importable " + "or no credentials are configured (WANDB_API_KEY); " + "continuing without the wandb tracker.") + trackers = [t for t in trackers if t != Trackers.WANDB.value] tracker_log_dir = (checkpoint_config.output_dir or os.getcwd()) if trackers: @@ -43,11 +68,40 @@ def build_tracker( tracker_run_name = tracker_config.run_name or None project = (tracker_config.project_name or "fastvideo") - return initialize_trackers( - trackers, - experiment_name=project, - config=tracker_config_dict, - log_dir=tracker_log_dir, - entity=tracker_entity, - run_name=tracker_run_name, - ) + try: + return initialize_trackers( + trackers, + experiment_name=project, + config=tracker_config_dict, + log_dir=tracker_log_dir, + entity=tracker_entity, + run_name=tracker_run_name, + ) + except Exception as exc: + if Trackers.WANDB.value not in trackers: + raise + # A revoked API key or unreachable api.wandb.ai passes _wandb_usable() + # (it only proves a key exists) and then throws inside wandb.init — + # which must not kill a multi-node run at boot. Offline init never + # contacts the API; the run stays syncable later via `wandb sync`. + logger.warning("Tracker init failed (%s); retrying wandb in offline mode.", exc) + os.environ["WANDB_MODE"] = "offline" + try: + return initialize_trackers( + trackers, + experiment_name=project, + config=tracker_config_dict, + log_dir=tracker_log_dir, + entity=tracker_entity, + run_name=tracker_run_name, + ) + except Exception as offline_exc: + logger.warning("Offline tracker init also failed (%s); continuing without trackers.", offline_exc) + return initialize_trackers( + [], + experiment_name=project, + config=None, + log_dir=tracker_log_dir, + entity=None, + run_name=None, + ) diff --git a/fastvideo/train/utils/training_config.py b/fastvideo/train/utils/training_config.py index 25a546cef0..f431a15e15 100644 --- a/fastvideo/train/utils/training_config.py +++ b/fastvideo/train/utils/training_config.py @@ -32,6 +32,10 @@ class DataConfig: num_width: int = 0 num_latent_t: int = 0 num_frames: int = 0 + # Preserve each T2VA row's native temporal/spatial latent geometry and + # schedule exact-shape global microbatches across data-parallel ranks. + # False retains the legacy fixed-shape/truncation contract. + native_shape_bucketing: bool = False @dataclass(slots=True) @@ -56,8 +60,23 @@ class TrainingLoopConfig: class CheckpointConfig: output_dir: str = "" resume_from_checkpoint: str = "" + # Deployable, model-only checkpoints are saved for every scheduled + # validation event and are independent from rolling resumable state. + save_inference_checkpoint_on_validation: bool = False + inference_checkpoint_role: str = "student" + inference_checkpoint_dtype: str = "bfloat16" training_state_checkpointing_steps: int = 0 + # Opt in to resumable checkpoints published only after DCP state and every + # rank's RNG snapshot are complete. False keeps legacy checkpoints usable. + require_complete_training_checkpoint: bool = False + # Applies only to resumable ``checkpoint-`` directories. Inference + # checkpoints are retained as the run's immutable model lineage. checkpoints_total_limit: int = 0 + checkpointing_start_step: int = 0 + # DCP checkpoints restore optimizer param-group LRs and scheduler + # base_lrs, silently overriding the YAML on resume. Set true to re-apply + # the configured learning rates after loading (LR-change experiments). + reset_lr_on_resume: bool = False @dataclass(slots=True) @@ -77,6 +96,15 @@ class ModelTrainingConfig: precondition_outputs: bool = False moba_config: dict = field(default_factory=dict) enable_gradient_checkpointing_type: str | None = None + # Optimizer steps applied to bf16/fp16 parameter storage round away + # updates smaller than ~half an ulp of each weight's magnitude — + # O(1)-magnitude parameters (norm gains) freeze entirely at typical + # distillation LRs. The trainer refuses to start unless master weights + # are fp32 (training.dit_precision: fp32) or this explicit opt-in + # acknowledges the effect (memory-constrained topologies). + allow_low_precision_master_weights: bool = False + enable_torch_compile: bool = False + torch_compile_kwargs: dict = field(default_factory=dict) @dataclass(slots=True) @@ -88,6 +116,12 @@ class TrainingConfig: checkpoint: CheckpointConfig = field(default_factory=CheckpointConfig) tracker: TrackerConfig = field(default_factory=TrackerConfig) vsa_sparsity: float = 0.0 + # Tokens per sparse-attention tile for the VSA student. 256 (default) + # keeps the (4,8,8) tiles and today's VSA-256 CuTe/Triton routing; 64 + # selects (4,4,4) tiles on the native 64-token Triton block-sparse + # kernels (forward and backward). Consumed by the VSA-H3 (MiniMax H3) + # backend; Wan's VSA path ignores it. + vsa_tile_size: int = 256 # Reuse the per-step padded VSA tile buffer across attention layers. # Defaults to False for training: under full activation checkpointing the # cached buffer survives into the backward recompute and inflates peak diff --git a/mkdocs.yml b/mkdocs.yml index 3a67cb01c4..0a6a63ce03 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -217,6 +217,11 @@ nav: - Video Sparse Attention: attention/vsa/index.md - Sliding Tile Attention (Archived): attention/sta/index.md - Backend Development: contributing/attention_backend.md + - Quantization: + - Quantized Checkpoint Loading: quantization/loader_quant_params.md + - Affine INT8 for MiniMax-H3: quantization/h3_int8_affine.md + - NVFP4 for MiniMax-H3: quantization/h3_nvfp4.md + - W4A16 (4-bit weight) for MiniMax-H3: quantization/h3_w4a16.md - Utilities: - LoRA: utilities/lora.md - Debugging: utilities/debugging.md diff --git a/scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py b/scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py index 4a5cdfb4cb..f3069f6780 100644 --- a/scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py +++ b/scripts/checkpoint_conversion/convert_minimax_h3_adaln_rank.py @@ -70,6 +70,7 @@ def parse_args() -> argparse.Namespace: p.add_argument("--grid", type=int, default=4096, help="Timestep samples used to fit the basis.") p.add_argument("--freq-dim", type=int, default=256) p.add_argument("--report-only", action="store_true", help="Fit and report error without writing.") + p.add_argument("--receipt", type=Path, help="Optional JSON receipt for automated parity gates.") return p.parse_args() @@ -94,7 +95,7 @@ def load_keys(src: Path, index: dict[str, str], keys: list[str]) -> dict[str, to def fit_basis(src: Path, index: dict[str, str], rank: int, grid: int, - freq_dim: int) -> tuple[torch.Tensor, torch.Tensor]: + freq_dim: int) -> tuple[torch.Tensor, torch.Tensor, float]: """Return (V [time_embed_dim, rank], U [grid, time_embed_dim]) in float64.""" embedder = load_keys(src, index, list(TIME_EMBEDDER_KEYS)) # t is the DiT's timestep input: scheduler.timesteps = 1 - sigmas, so t in [0, 1]. @@ -108,15 +109,26 @@ def fit_basis(src: Path, index: dict[str, str], rank: int, grid: int, _, s, vh = torch.linalg.svd(u, full_matrices=False) residual = ((s[rank:]**2).sum() / (s**2).sum()).sqrt() print(f"basis: U={tuple(u.shape)} rank={rank} relative residual ||U-U_r||/||U|| = {residual:.3e}") - return vh[:rank].T.contiguous(), u + return vh[:rank].T.contiguous(), u, float(residual) def main() -> None: args = parse_args() src, dst = Path(args.src), Path(args.dst) - index_map = json.loads((src / INDEX_NAME).read_text())["weight_map"] - - basis, u = fit_basis(src, index_map, args.rank, args.grid, args.freq_dim) + index_path = src / INDEX_NAME + if index_path.exists(): + index_map = json.loads(index_path.read_text())["weight_map"] + else: + # dcp_to_diffusers writes a single shard for compact students; support it. + single = src / "model.safetensors" + if not single.exists(): + raise SystemExit(f"{src} has neither {INDEX_NAME} nor model.safetensors") + with safe_open(str(single), framework="pt") as handle: + # `safe_open` is not itself iterable in the cluster's pinned + # safetensors build; `.keys()` works across both old and new APIs. + index_map = {key: "model.safetensors" for key in handle.keys()} + + basis, u, residual = fit_basis(src, index_map, args.rank, args.grid, args.freq_dim) # Worst-case induced error on the actual modulation outputs. worst = 0.0 @@ -129,7 +141,21 @@ def main() -> None: scale = max(scale, ref.abs().max().item()) print(f"modulation error over all projections: max|err|={worst:.3e} " f"(|Wu|max={scale:.3f}, relative={worst / scale:.3e})") + receipt = { + "schema_version": 1, + "source": str(src.resolve()), + "destination": None if args.report_only else str(dst.resolve()), + "rank": args.rank, + "grid_points": args.grid, + "basis_relative_residual": residual, + "modulation_max_abs_error": worst, + "modulation_reference_absmax": scale, + "modulation_relative_max_error": worst / scale, + } if args.report_only: + if args.receipt: + args.receipt.parent.mkdir(parents=True, exist_ok=True) + args.receipt.write_text(json.dumps(receipt, indent=2) + "\n") return dst.mkdir(parents=True, exist_ok=True) @@ -171,6 +197,14 @@ def main() -> None: print(f"\nparameters: {total_before / 1e9:.3f}B -> {total_after / 1e9:.3f}B " f"({100 * (1 - total_after / total_before):.1f}% removed)") print(f"bf16 footprint: {total_before * 2 / 1e9:.1f} GB -> ~{total_after * 2 / 1e9:.1f} GB") + receipt.update({ + "parameters_before": total_before, + "parameters_after": total_after, + "fraction_removed": 1 - total_after / total_before, + }) + if args.receipt: + args.receipt.parent.mkdir(parents=True, exist_ok=True) + args.receipt.write_text(json.dumps(receipt, indent=2) + "\n") if __name__ == "__main__": diff --git a/scripts/checkpoint_conversion/export_h3_dmd2_student.py b/scripts/checkpoint_conversion/export_h3_dmd2_student.py new file mode 100644 index 0000000000..4af0e14b68 --- /dev/null +++ b/scripts/checkpoint_conversion/export_h3_dmd2_student.py @@ -0,0 +1,202 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Export a MiniMax-H3 DMD2 training checkpoint's student into an inference model dir. + +The modular trainer's DCP checkpoints hold every role and optimizer +(``roles.student.transformer.*`` is the piece inference needs, stored as the +fp32 master weights). This script streams just those tensors out of the DCP +shards in bounded memory (no process group, no GPU), renames them from +fastvideo layer names back to the on-disk checkpoint convention (the inverse +of ``MiniMaxH3ArchConfig.param_names_mapping``), casts to bf16, and writes a +sharded safetensors ``transformer/`` next to symlinks into the base model dir +for every other component — so the output loads through the standard +inference pipeline with ~66 GB of new bytes instead of a full copy. + +One param per block has no on-disk counterpart: ``attn.to_gate_compress``, +the VSA gate. The base checkpoint default-initializes it; a VSA-trained +student's gate is learned, so it is exported under its fastvideo name (no +mapping rule touches it, and the loader resolves it to the model param +verbatim). Dense-only consumers ignore it. + +Usage:: + + python scripts/checkpoint_conversion/export_h3_dmd2_student.py \ + --checkpoint /path/to/outputs//checkpoint-1400 \ + --output-dir /path/to/exports/-step1400 \ + [--base-model /mnt/lustre/vlm-k1kong/models/MiniMax-H3] \ + [--role student] [--dtype bfloat16] [--copy-components] + +``--checkpoint latest --run-dir `` picks the newest +``checkpoint-*`` whose ``dcp/.metadata`` exists (the strict-resume contract). +""" + +from __future__ import annotations + +import argparse +import json +import re +import shutil +from pathlib import Path + +import torch +from safetensors.torch import save_file + +# Inverse of MiniMaxH3ArchConfig.param_names_mapping (fastvideo -> disk). +INVERSE_PARAM_RULES: tuple[tuple[str, str], ...] = ( + (r"^time_embedder\.fc_in\.(.*)$", r"time_embedder.linear_1.\1"), + (r"^time_embedder\.fc_out\.(.*)$", r"time_embedder.linear_2.\1"), + (r"^(.*)\.attn\.to_out\.(weight|bias)$", r"\1.attn.to_out.0.\2"), + (r"^(.*)\.ff\.fc_in\.(.*)$", r"\1.ff.net.0.proj.\2"), + (r"^(.*)\.ff\.fc_out\.(.*)$", r"\1.ff.net.2.\2"), +) + +# fastvideo-only trained params expected to have no disk counterpart. +EXPECTED_NEW_PARAM_PATTERNS = (re.compile(r"\.attn\.to_gate_compress\."), ) + +SHARD_BUDGET_BYTES = 5 * 1024**3 # ~5 GB per safetensors shard (bf16) + + +def to_disk_name(name: str) -> str: + for pattern, repl in INVERSE_PARAM_RULES: + new, n = re.subn(pattern, repl, name) + if n: + return new + return name + + +def find_latest_checkpoint(run_dir: Path) -> Path: + candidates = sorted( + (p for p in run_dir.glob("checkpoint-*") if (p / "dcp" / ".metadata").exists()), + key=lambda p: int(p.name.rsplit("-", 1)[-1]), + ) + if not candidates: + raise FileNotFoundError(f"no complete checkpoint-*/dcp/.metadata under {run_dir}") + return candidates[-1] + + +def main(args: argparse.Namespace) -> None: + if args.checkpoint == "latest": + if args.run_dir is None: + raise SystemExit("--checkpoint latest requires --run-dir") + checkpoint = find_latest_checkpoint(args.run_dir) + else: + checkpoint = Path(args.checkpoint) + dcp_dir = checkpoint / "dcp" + if not (dcp_dir / ".metadata").exists(): + raise FileNotFoundError(f"{dcp_dir}/.metadata missing — incomplete checkpoint, refusing") + base_model = args.base_model + if not any((base_model / name).exists() for name in ("model_index.json", "modular_model_index.json")): + raise FileNotFoundError(f"{base_model} does not look like a model dir " + "(no model_index.json or modular_model_index.json)") + out_dtype = getattr(torch, args.dtype) + prefix = f"roles.{args.role}.transformer." + + from torch.distributed.checkpoint import FileSystemReader + import torch.distributed.checkpoint as dcp + + reader = FileSystemReader(str(dcp_dir)) + metadata = reader.read_metadata() + param_meta = { + key: meta + for key, meta in metadata.state_dict_metadata.items() + if key.startswith(prefix) + } + if not param_meta: + raise SystemExit(f"no keys under {prefix!r} in {dcp_dir}") + print(f"{checkpoint.name}: {len(param_meta)} tensors under {prefix!r}") + + # Reference key set from the base transformer, to catch mapping drift. + base_transformer = base_model / "transformer" + base_index = base_transformer / "diffusion_pytorch_model.safetensors.index.json" + base_keys: set[str] = set() + if base_index.exists(): + base_keys = set(json.loads(base_index.read_text())["weight_map"]) + + # Plan shards by byte budget over the (deterministic) sorted key order. + def nbytes(meta) -> int: + numel = 1 + for dim in meta.size: + numel *= dim + return numel * torch.finfo(out_dtype).bits // 8 + + ordered = sorted(param_meta) + shards: list[list[str]] = [[]] + acc = 0 + for key in ordered: + size = nbytes(param_meta[key]) + if shards[-1] and acc + size > SHARD_BUDGET_BYTES: + shards.append([]) + acc = 0 + shards[-1].append(key) + acc += size + + out_transformer = args.output_dir / "transformer" + out_transformer.mkdir(parents=True, exist_ok=True) + + weight_map: dict[str, str] = {} + total_size = 0 + unexpected_new: list[str] = [] + n_shards = len(shards) + for shard_idx, keys in enumerate(shards, start=1): + fname = f"diffusion_pytorch_model-{shard_idx:05d}-of-{n_shards:05d}.safetensors" + state = { + key: torch.empty(tuple(param_meta[key].size), dtype=param_meta[key].properties.dtype) + for key in keys + } + dcp.load(state, checkpoint_id=str(dcp_dir)) + tensors: dict[str, torch.Tensor] = {} + for key, tensor in state.items(): + disk_name = to_disk_name(key[len(prefix):]) + if base_keys and disk_name not in base_keys: + if not any(p.search(disk_name) for p in EXPECTED_NEW_PARAM_PATTERNS): + unexpected_new.append(disk_name) + tensors[disk_name] = tensor.to(out_dtype).contiguous() + weight_map[disk_name] = fname + total_size += tensors[disk_name].numel() * tensors[disk_name].element_size() + save_file(tensors, str(out_transformer / fname)) + del state, tensors + print(f" wrote {fname} ({len(keys)} tensors)") + + if unexpected_new: + raise SystemExit("Export produced keys unknown to the base checkpoint (mapping drift?):\n " + + "\n ".join(sorted(unexpected_new)[:20])) + + (out_transformer / "diffusion_pytorch_model.safetensors.index.json").write_text( + json.dumps({"metadata": {"total_size": total_size}, "weight_map": weight_map}, indent=2)) + shutil.copy2(base_transformer / "config.json", out_transformer / "config.json") + + # Everything but the transformer comes from the base model dir. + for entry in sorted(base_model.iterdir()): + if entry.name == "transformer": + continue + target = args.output_dir / entry.name + if target.exists() or target.is_symlink(): + continue + if args.copy_components: + if entry.is_dir(): + shutil.copytree(entry, target) + else: + shutil.copy2(entry, target) + else: + target.symlink_to(entry.resolve()) + + print(f"Export complete: {args.output_dir}") + print(f" transformer: {len(weight_map)} tensors, {total_size / 1024**3:.1f} GiB " + f"({args.dtype}), {n_shards} shards") + print(f" other components {'copied' if args.copy_components else 'symlinked'} from {base_model}") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--checkpoint", required=True, help="checkpoint-N dir, or 'latest' with --run-dir") + parser.add_argument("--run-dir", type=Path, default=None, help="training output dir for --checkpoint latest") + parser.add_argument("--output-dir", type=Path, required=True) + parser.add_argument("--base-model", + type=Path, + default=Path("/mnt/lustre/vlm-k1kong/models/MiniMax-H3"), + help="base model dir supplying config + non-transformer components") + parser.add_argument("--role", default="student", help="training role to export (student|critic)") + parser.add_argument("--dtype", default="bfloat16", choices=["bfloat16", "float32", "float16"]) + parser.add_argument("--copy-components", + action="store_true", + help="copy non-transformer components instead of symlinking") + main(parser.parse_args()) diff --git a/scripts/compacth3/analysis/adaln/analyze_adaln_rank.py b/scripts/compacth3/analysis/adaln/analyze_adaln_rank.py new file mode 100644 index 0000000000..cdee30f2fa --- /dev/null +++ b/scripts/compacth3/analysis/adaln/analyze_adaln_rank.py @@ -0,0 +1,423 @@ +#!/usr/bin/env python3 +"""Post-hoc spectral analysis of MiniMax-H3's timestep (AdaLN) conditioning. + +Part 3 of the "is rank-768 timestep conditioning unnecessary capacity?" question. + +Three levels, all measured by calling the checkpoint's OWN modules on a synthetic +timestep grid (nothing is trained, nothing in the model code is modified): + + A. SHARED COORDINATE z(t) = adaln_basis(silu(time_embedder(time_proj(t)))) (T x 768) + B. PER-BLOCK MODULATION m_i(t) = block_i.adaln_proj(z(t)) (T x 96768) + C. SHARED BASIS ACROSS BLOCKS joint (summed) covariance of all centered block + trajectories -> how many directions represent ALL + blocks' timestep variation at once + +Timestep range (read from the code, not assumed): MiniMaxH3Scheduler stores +timesteps = 1 - sigmas[:-1] with sigmas in [0, 1] (rectified flow, "clean time" +convention), and build_row_timesteps additionally feeds 0.999 (conditioned video +rows) and 1.0 (conditioned audio rows). So the model's usable timestep range is +exactly [0, 1]. The 4-call DMD ladder (FASTVIDEO_DMD_DENOISING_STEPS=999,749,500,250 +divided by 1000 and shift-warped) lands on t in {0.0001, 0.027, 0.077, 0.2} for video +and {0.0003, 0.1, 0.25, 0.5} for audio. + +Grids analysed: + uniform_grid linspace(0, 1, 4096) -- primary + operating_points union of the literal scheduler timesteps at 4/5/49/50 points + for both shipped shifts plus 0.999 and 1.0 -- what inference + actually feeds + control_0_1000 linspace(0, 1000, 4096) -- CONTROL ONLY, not a convention + this codebase uses: shows how much of the low rank is a + property of the narrow [0,1] range + downscale_freq_shift=0 + embedding rather than of the learned basis. + +Must be a real file on disk (never stdin): FastVideo workers re-execute __main__ +via runpy and a heredoc has no path. +""" +from __future__ import annotations + +import argparse +import contextlib +import importlib.util +import json +import os +import sys +import time +from pathlib import Path + +import torch + +SPRINT = "/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829" +M = "/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1" +HARNESS = f"{M}/examples/inference/basic/basic_fasth3.py" +OUT_DIR = Path(SPRINT) / "adaln_rank_analysis" + +N_GRID = 4096 +ENERGY_THRESHOLDS = (0.90, 0.95, 0.99, 0.999) +# cumulative-energy checkpoints for the cross-block shared-basis coverage table +COVERAGE_K = (1, 2, 3, 4, 6, 8, 12, 16, 24, 32, 48, 64, 96, 128, 192, 256, 384, 512, 768) +MAX_SIGMA_STORED = 1024 +DEVICE = "cuda:0" +FP32_EPS = float(torch.finfo(torch.float32).eps) +N_MODALITY = 3 # MINIMAX_H3_MODALITY_NUM + + +def log(msg: str) -> None: + print(f"[adaln-rank] {time.strftime('%H:%M:%S')} {msg}", flush=True) + + +# ---------------------------------------------------------------------------------- +# checkpoint loading (the proven PipelineComponentLoader path) +# ---------------------------------------------------------------------------------- +def build_fastvideo_args(model_path: str): + spec = importlib.util.spec_from_file_location("fasth3_harness", HARNESS) + harness = importlib.util.module_from_spec(spec) + sys.modules["fasth3_harness"] = harness + spec.loader.exec_module(harness) + + # argparse treats a passed sequence as the FULL argument list (it does not + # strip a program name), so there is no prog element here. + argv = [ + "--model-path", model_path, + "--prompt", "adaln-rank-analysis", + "--num-gpus", "1", + "--no-fa4", + "--no-inference-torch-compile", + "--steps", "5", + ] + args = harness.parse_args(argv) + args.fa4 = False + harness.configure_environment(args) + # No attention forward is ever run; keep the loader off the fused/sparse + # kernels entirely so this analysis cannot depend on VSA/FA4 availability. + os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "TORCH_SDPA" + + config = harness.build_generator_config(args) + config.pipeline.experimental["attention_backend"] = "TORCH_SDPA" + from fastvideo.api.compat import generator_config_to_fastvideo_args + return generator_config_to_fastvideo_args(config) + + +def load_dit(fastvideo_args, transformer_path: str): + from fastvideo.models.loader.component_loader import PipelineComponentLoader + log(f"loading transformer from {transformer_path}") + model = PipelineComponentLoader.load_module( + module_name="transformer", + component_model_path=transformer_path, + transformers_or_diffusers="diffusers", + fastvideo_args=fastvideo_args, + ) + log(f"loaded class={type(model).__name__}") + return model + + +@contextlib.contextmanager +def on_device_fp32(module): + """Temporarily put exactly this small module on GPU in fp32, then restore.""" + saved = [(name, p.dtype, p.device) for name, p in module.named_parameters()] + module.to(device=DEVICE, dtype=torch.float32) + try: + yield module + finally: + params = dict(module.named_parameters()) + for name, dtype, device in saved: + params[name].data = params[name].data.to(device=device, dtype=dtype) + torch.cuda.empty_cache() + + +# ---------------------------------------------------------------------------------- +# timestep grid +# ---------------------------------------------------------------------------------- +def scheduler_timesteps(shift: float, grid_points: int) -> torch.Tensor: + """Verbatim replica of MiniMaxH3Scheduler.set_timesteps (scheduler file read).""" + base = torch.linspace(1.0, 0.0, int(grid_points), dtype=torch.float32) + sigma = shift * base / (1 + (shift - 1) * base) + sigma = torch.unique_consecutive(sigma) + return 1.0 - sigma[:-1] + + +def build_grids(video_shift: float, audio_shift: float): + uniform = torch.linspace(0.0, 1.0, N_GRID, dtype=torch.float32) + pieces = [torch.tensor([0.0, 0.999, 1.0], dtype=torch.float32)] + for shift in (video_shift, audio_shift): + for n in (4, 5, 49, 50): + pieces.append(scheduler_timesteps(shift, n)) + ops = torch.unique(torch.cat(pieces)).sort().values + control = torch.linspace(0.0, 1000.0, N_GRID, dtype=torch.float32) + return uniform, ops, control + + +# ---------------------------------------------------------------------------------- +# spectral helpers +# ---------------------------------------------------------------------------------- +def spectral_stats(sigma: torch.Tensor, shape: tuple[int, int]) -> dict: + s = sigma.detach().double().cpu() + energy = s * s + total = float(energy.sum()) + out: dict = { + "shape": [int(shape[0]), int(shape[1])], + "frobenius_norm": float(total ** 0.5), + "s_max": float(s[0]), + "s_min": float(s[-1]), + # standard numpy matrix_rank tolerance. NOTE the real noise floor here + # is set by arithmetic on bf16-derived weights, so this column is the + # least meaningful one -- the energy ranks are the scientific answer. + "numerical_rank_fp32tol": int((s > s[0] * max(shape) * FP32_EPS).sum()), + "stable_rank_trace_over_smax2": float(total / (float(s[0]) ** 2)), + } + # Shape-independent rank floors, so levels with different D are comparable. + for rel in (1e-2, 1e-3, 1e-4, 1e-5, 1e-6): + out[f"rank_sigma_above_{rel:g}_of_smax"] = int((s > s[0] * rel).sum()) + p = energy / total + nz = p > 0 + out["entropy_effective_rank"] = float(torch.exp(-(p[nz] * p[nz].log()).sum())) + cum = torch.cumsum(energy, 0) / total + for th in ENERGY_THRESHOLDS: + k = int(torch.searchsorted(cum, th).item()) + 1 + out[f"rank_at_{th:.3f}_energy"] = k + out[f"relerr_at_{th:.3f}_energy"] = float(max(0.0, 1.0 - float(cum[k - 1])) ** 0.5) + return out + + +def gram_eigvals(G: torch.Tensor) -> torch.Tensor: + """Descending singular values of the trajectory from its (T x T) Gram matrix.""" + lam = torch.linalg.eigvalsh(G.double()) + return torch.clamp(lam.flip(0), min=0.0).sqrt() + + +# ---------------------------------------------------------------------------------- +# level B + C on one trajectory matrix family +# ---------------------------------------------------------------------------------- +def analyze_blocks(model, mods, Z, H, store_sigma: bool, do_coverage: bool) -> tuple[dict, dict]: + T = int(Z.shape[0]) + D = 6 * H * N_MODALITY + per_block: dict = {} + G_union = torch.zeros(T, T, dtype=torch.float32, device=DEVICE) + G_union_norm = torch.zeros_like(G_union) + G_mod_union = [torch.zeros_like(G_union) for _ in range(N_MODALITY)] + G_list: list[torch.Tensor] = [] + + for name, block in mods: + t0 = time.time() + proj = block.adaln_proj + with on_device_fp32(proj) as mod: + six = torch.stack([s.float() for s in mod(Z)], dim=0) # (6, 3T, H) + M = six.permute(1, 0, 2).reshape(T, N_MODALITY, 6, H).reshape(T, -1) + del six + M = M - M.mean(dim=0, keepdim=True) + G = M @ M.t() + G = 0.5 * (G + G.t()) + sv = gram_eigvals(G) + stats = spectral_stats(sv, (T, D)) + stats["apply_silu"] = bool(getattr(proj, "apply_silu", None)) + stats["linear_weight_shape"] = list(proj.linear.weight.shape) + stats["modulation_dim"] = D + entry = {"stats": stats, + "frobenius_per_modality": [float(M[:, m * 6 * H:(m + 1) * 6 * H].double().norm()) + for m in range(N_MODALITY)]} + if store_sigma: + entry["sigma_top"] = [float(v) for v in sv[:MAX_SIGMA_STORED].cpu()] + entry["sigma_stored"] = int(min(MAX_SIGMA_STORED, sv.numel())) + entry["sigma_nonzero"] = int(sv.numel()) + per_block[name] = entry + + tr = float(torch.diagonal(G).sum()) + G_union += G + if tr > 0: + G_union_norm += G / tr + for m in range(N_MODALITY): + sl = M[:, m * 6 * H:(m + 1) * 6 * H] + G_mod_union[m] += sl @ sl.t() + if do_coverage: + G_list.append(G) + del M + log(f" {name}: rank90={stats['rank_at_0.900_energy']} rank99={stats['rank_at_0.990_energy']} " + f"rank999={stats['rank_at_0.999_energy']} stable={stats['stable_rank_trace_over_smax2']:.2f} " + f"s_max={stats['s_max']:.4g} ({time.time()-t0:.1f}s)") + + def quartiles(key): + vals = sorted(pb["stats"][key] for pb in per_block.values()) + return [vals[0], vals[len(vals) // 2], vals[-1]] + + level_b = { + "modulation_dim": D, + "blocks": per_block, + "summary": { + "n_blocks": len(per_block), + "rank90_min_median_max": quartiles("rank_at_0.900_energy"), + "rank99_min_median_max": quartiles("rank_at_0.990_energy"), + "rank999_min_median_max": quartiles("rank_at_0.999_energy"), + "stable_rank_min_median_max": quartiles("stable_rank_trace_over_smax2"), + }, + } + + cols = D * len(per_block) + + def union_stats(Gmat: torch.Tensor, label: str) -> dict: + sv = gram_eigvals(Gmat) + return {"label": label, "stats": spectral_stats(sv, (T, cols)), + "sigma_top": [float(v) for v in sv[:MAX_SIGMA_STORED].cpu()]} + + level_c = { + "energy_weighted": union_stats(G_union, "sum of centered block Gram matrices"), + "block_normalized": union_stats(G_union_norm, + "sum of trace-normalized centered block Gram matrices"), + "per_modality_uniform_weighted": [ + {"modality": m, + "stats": spectral_stats(gram_eigvals(G_mod_union[m]), (T, 6 * H * len(per_block)))} + for m in range(N_MODALITY) + ], + } + if do_coverage: + Gmat = G_union + lam, V = torch.linalg.eigh(Gmat.double()) + Vf = V.flip(1)[:, :max(COVERAGE_K)].float() + rows = [] + for G in G_list: + GV = G @ Vf + d = (Vf * GV).sum(dim=0).double() + rows.append((torch.cumsum(d, 0) / float(torch.diagonal(G).sum())).cpu()) + C = torch.stack(rows) + level_c["coverage_k"] = list(COVERAGE_K) + level_c["coverage_min_median_max_fraction_of_block_energy"] = { + str(k): [float(C[:, k - 1].min()), float(C[:, k - 1].median()), float(C[:, k - 1].max())] + for k in COVERAGE_K} + level_c["coverage_per_block"] = [[float(x) for x in row] for row in C] + return level_b, level_c + + +# ---------------------------------------------------------------------------------- +# main measurement +# ---------------------------------------------------------------------------------- +def measure(model, tag: str, ckpt: str, n_blocks_limit: int | None): + from fastvideo.models.dits.minimax_h3 import MiniMaxH3AdaLayerNormModulation + + log(f"adaln_rank={model.adaln_rank} hidden={model.hidden_size} " + f"blocks={len(model.transformer_blocks)}") + + ckpt_dir = Path(ckpt) + video_shift = float(json.loads((ckpt_dir / "scheduler" / "scheduler_config.json").read_text())["shift"]) + audio_shift = float(json.loads((ckpt_dir / "audio_scheduler" / "scheduler_config.json").read_text())["shift"]) + uniform, ops, control = build_grids(video_shift, audio_shift) + grids = {"uniform_grid": uniform, "operating_points": ops, "control_0_1000": control} + log(f"grids: uniform={tuple(uniform.shape)} ops={tuple(ops.shape)} " + f"ops_range=({float(ops[0]):.4f},{float(ops[-1]):.4f}) " + f"control={tuple(control.shape)} video_shift={video_shift} audio_shift={audio_shift}") + + H = int(model.hidden_size) + + result: dict = { + "model_tag": tag, + "checkpoint": str(ckpt), + "adaln_rank": int(model.adaln_rank), + "hidden_size": H, + "num_layers": len(model.transformer_blocks), + "num_refiner_layers": len(model.token_refiner.refiner_blocks), + "scheduler_shifts": {"video": video_shift, "audio": audio_shift}, + "grids": { + "uniform_grid": {"n": int(uniform.numel()), "range": [0.0, 1.0], "role": "primary"}, + "operating_points": {"n": int(ops.numel()), "values": [round(float(v), 6) for v in ops], + "role": "the literal timesteps inference feeds"}, + "control_0_1000": {"n": int(control.numel()), "range": [0.0, 1000.0], + "role": "CONTROL ONLY -- not a convention this codebase uses"}, + }, + "precision_note": ("level A/B/C spectra are computed with the three AdaLN " + "modules cast to float32 in memory (checkpoint weights are " + "bf16); level_a_shared_coordinate.uniform_grid_bf16_native " + "repeats level A on the unmodified bf16 modules to expose " + "the storage-format noise floor."), + "device": DEVICE, + } + + # ---------------- shared coordinate: temb and z(t) ---------------- + def shared_coordinate(grid: torch.Tensor, native_bf16: bool = False) -> torch.Tensor: + t = grid.to(DEVICE, dtype=torch.float32) + if native_bf16: + temb = model.time_proj(t) + temb = model.time_embedder(temb.to(model.time_embedder.fc_in.weight.dtype)) + z, _ = model.adaln_basis(torch.nn.functional.silu(temb).to(model.adaln_basis.weight.dtype)) + return z.detach().float() + with on_device_fp32(model.time_embedder) as te: + temb = te(model.time_proj(t).to(te.fc_in.weight.dtype)) + with on_device_fp32(model.adaln_basis) as basis: + z, _ = basis(torch.nn.functional.silu(temb).to(basis.weight.dtype)) + return z.detach().float() + + def level_a(Zmat: torch.Tensor) -> dict: + Zc = Zmat - Zmat.mean(dim=0, keepdim=True) + sv = torch.linalg.svdvals(Zc.double()).float() + stats = spectral_stats(sv, tuple(Zc.shape)) + stats["centered_frobenius"] = float(Zc.double().norm()) + stats["raw_frobenius"] = float(Zmat.double().norm()) + stats["mean_row_norm"] = float(Zmat.mean(dim=0).double().norm()) + return {"sigma": [float(v) for v in sv.cpu()], "stats": stats} + + t0 = time.time() + Zcache = {name: shared_coordinate(g) for name, g in grids.items()} + log(f"shared coordinates done in {time.time()-t0:.1f}s " + f"{ {k: tuple(v.shape) for k, v in Zcache.items()} }") + level_a_out = {name: level_a(Z) for name, Z in Zcache.items()} + level_a_out["uniform_grid_bf16_native"] = level_a(shared_coordinate(uniform, native_bf16=True)) + result["level_a_shared_coordinate"] = level_a_out + for name in level_a_out: + log(f"LEVEL A {name}: {json.dumps(level_a_out[name]['stats'])}") + + # ---------------- per-block modulation + shared basis ---------------- + mods: list[tuple[str, torch.nn.Module]] = [ + (f"transformer_blocks.{i}", b) for i, b in enumerate(model.transformer_blocks)] + refiner_mods = [n for n, m in model.token_refiner.named_modules() + if isinstance(m, MiniMaxH3AdaLayerNormModulation)] + result["refiner_adaln_modules"] = refiner_mods + log(f"refiner AdaLN modules found: {refiner_mods or 'NONE'}") + if n_blocks_limit is not None: + mods = mods[:n_blocks_limit] + log(f"per-block modulation dim = {6 * H * N_MODALITY} (6 x {H} x {N_MODALITY} modalities), " + f"{len(mods)} blocks") + + result["level_b_per_block"] = {} + result["level_c_shared_basis"] = {} + for name, Z in Zcache.items(): + primary = name == "uniform_grid" + log(f" --- block analysis on {name} (T={Z.shape[0]}) ---") + level_b, level_c = analyze_blocks(model, mods, Z, H, store_sigma=True, do_coverage=primary) + result["level_b_per_block"][name] = level_b + result["level_c_shared_basis"][name] = level_c + log(f"LEVEL B {name} summary: {json.dumps(level_b['summary'])}") + log(f"LEVEL C {name} weighted: {json.dumps(level_c['energy_weighted']['stats'])}") + log(f"LEVEL C {name} normalized: {json.dumps(level_c['block_normalized']['stats'])}") + if primary: + log(f"LEVEL C {name} coverage: " + f"{json.dumps(level_c['coverage_min_median_max_fraction_of_block_energy'])}") + return result + + +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--tag", required=True) + ap.add_argument("--model-path", required=True) + ap.add_argument("--out-dir", default=str(OUT_DIR)) + ap.add_argument("--max-blocks", type=int, default=None, help="debug: only the first N blocks") + args = ap.parse_args() + + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + torch.set_grad_enabled(False) + + out_dir = Path(args.out_dir) + out_dir.mkdir(parents=True, exist_ok=True) + + fva = build_fastvideo_args(args.model_path) + model = load_dit(fva, str(Path(args.model_path) / "transformer")) + model.eval() + + result = measure(model, args.tag, args.model_path, args.max_blocks) + out_path = out_dir / f"adaln_rank_{args.tag}.json" + out_path.write_text(json.dumps(result, indent=1)) + log(f"wrote {out_path} ({out_path.stat().st_size} bytes)") + + del model + torch.cuda.empty_cache() + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/compacth3/analysis/adaln/compare_parent_dmd2.py b/scripts/compacth3/analysis/adaln/compare_parent_dmd2.py new file mode 100644 index 0000000000..9bcdadcad7 --- /dev/null +++ b/scripts/compacth3/analysis/adaln/compare_parent_dmd2.py @@ -0,0 +1,67 @@ +#!/usr/bin/env python3 +"""Sharper parent-vs-DMD2 comparison from the saved singular spectra.""" +import json +from pathlib import Path + +OUT = Path("/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/adaln_rank_analysis") +D = {t: json.loads((OUT / f"adaln_rank_{t}.json").read_text()) for t in ("parent", "dmd2")} + + +def topk_energy(sigma, k): + s = [float(x) for x in sigma[:k]] + tot = sum(float(x) ** 2 for x in sigma) + return sum(x * x for x in s) / tot + + +print("LEVEL A (shared coordinate z(t), 4096 x 768, uniform grid, centered)") +print(" model s1 s2/s1 s3/s1 s4/s1 s5/s1 s6/s1 E@1 E@2 E@3 E@4 E@6") +for t in ("parent", "dmd2"): + sv = D[t]["level_a_shared_coordinate"]["uniform_grid"]["sigma"] + r = [sv[i] / sv[0] if i < len(sv) else 0.0 for i in range(6)] + print(" {:<6s} {:<9.4f} {:<9.5f} {:<9.5f} {:<9.5f} {:<9.5f} {:<9.5f} ".format( + t, sv[0], r[1], r[2], r[3], r[4], r[5]) + + " ".join("{:.4f}".format(topk_energy(sv, k)) for k in (1, 2, 3, 4, 6))) + +print() +print("LEVEL C (union over all 42 blocks, 4096 x 4064256, centered, energy-weighted)") +print(" model E@1 E@2 E@3 E@4 E@6 E@8 stable") +for t in ("parent", "dmd2"): + c = D[t]["level_c_shared_basis"]["uniform_grid"]["energy_weighted"]["sigma_top"] + st = D[t]["level_c_shared_basis"]["uniform_grid"]["energy_weighted"]["stats"]["stable_rank_trace_over_smax2"] + print(" {:<6s} ".format(t) + " ".join("{:.4f}".format(topk_energy(c, k)) for k in (1, 2, 3, 4, 6, 8)) + + " {:.4f}".format(st)) + +print() +print("PER-BLOCK (42 blocks, uniform grid)") +hdr = " model rank99 mean/median rank99.9 mean/median stable mean E@2 per block min/med/max" +print(hdr) +for t in ("parent", "dmd2"): + bl = D[t]["level_b_per_block"]["uniform_grid"]["blocks"] + r99 = sorted(v["stats"]["rank_at_0.990_energy"] for v in bl.values()) + r999 = sorted(v["stats"]["rank_at_0.999_energy"] for v in bl.values()) + st = sorted(v["stats"]["stable_rank_trace_over_smax2"] for v in bl.values()) + cov = D[t]["level_c_shared_basis"]["uniform_grid"]["coverage_min_median_max_fraction_of_block_energy"]["2"] + n = len(r99) + print(" {:<7s} {:.2f}/{:.1f} {:.2f}/{:.1f} {:.4f} {:.4f}/{:.4f}/{:.4f}".format( + t, sum(r99) / n, r99[n // 2], sum(r999) / n, r999[n // 2], sum(st) / n, cov[0], cov[1], cov[2])) + +print() +print("PER-MODALITY PROFILE (level B per-block, energy fraction of each modality table; parent vs dmd2)") +for t in ("parent", "dmd2"): + bl = D[t]["level_b_per_block"]["uniform_grid"]["blocks"] + fr = [[v["frobenius_per_modality"][m] for v in bl.values()] for m in range(3)] + print(" {:<7s} modality frobenius mean: ".format(t) + + " ".join("m{}= {:.1f}".format(m, sum(fr[m]) / len(fr[m])) for m in range(3)) + + " (mean ratio m1/m0={:.4f}, m2/m0={:.4f})".format( + (sum(fr[1]) / len(fr[1])) / (sum(fr[0]) / len(fr[0])), + (sum(fr[2]) / len(fr[2])) / (sum(fr[0]) / len(fr[0])))) + +print() +print("OPS GRID (191 literal inference timesteps)") +print(" model A stable C stable B rank99 med") +for t in ("parent", "dmd2"): + a = D[t]["level_a_shared_coordinate"]["operating_points"]["stats"] + c = D[t]["level_c_shared_basis"]["operating_points"]["energy_weighted"]["stats"] + b = D[t]["level_b_per_block"]["operating_points"]["summary"]["rank99_min_median_max"] + print(" {:<7s} {:.4f} {:.4f} {}".format(t, a["stable_rank_trace_over_smax2"], + c["stable_rank_trace_over_smax2"], b)) diff --git a/scripts/compacth3/analysis/adaln/run_adaln_rank.sh b/scripts/compacth3/analysis/adaln/run_adaln_rank.sh new file mode 100644 index 0000000000..108d894dc5 --- /dev/null +++ b/scripts/compacth3/analysis/adaln/run_adaln_rank.sh @@ -0,0 +1,59 @@ +#!/bin/bash +# Container entrypoint for the AdaLN timestep-rank spectral analysis. +# +# usage: +# sbatch /exp.sbatch /adaln_rank_analysis/run_adaln_rank.sh +# env: +# STAGE=both|parent|dmd2 (default both) -- which checkpoints to analyse +# SMOKE=1 (default 0) -- only the first 2 blocks, load check +set -uo pipefail + +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +M=/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1 +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +SCRIPT="${SPRINT}/adaln_rank_analysis/analyze_adaln_rank.py" +OUT="${SPRINT}/adaln_rank_analysis" + +PARENT_CKPT="${SPRINT}/runs/release20b-folded-long-4k-v3/dmd-parent-step750-complete-v1" +DMD2_CKPT="${SPRINT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3/inference/checkpoint-1400" + +STAGE="${STAGE:-both}" +SMOKE="${SMOKE:-0}" + +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache +export PYTHONDONTWRITEBYTECODE=1 +# M first: the venv's editable fastvideo finder is appended to sys.meta_path, so +# the path-based finder (sys.path) still wins and M is the code that runs. +export PYTHONPATH="${M}:${SPRINT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" +export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA +export FASTVIDEO_MINIMAX_H3_FUSIONS=0 +export FASTVIDEO_DISABLE_ATTENTION_COMPILE=1 +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True +export TORCH_NCCL_ENABLE_MONITORING=0 +source /mnt/nfs/vlm-aryan/fasth3-33b-20260806/secrets.env >/dev/null 2>&1 || true + +mkdir -p "$OUT" +cd "$M" + +echo "=== host=$(hostname) gpus=$(nvidia-smi -L | wc -l) stage=${STAGE} smoke=${SMOKE} ===" +nvidia-smi --query-gpu=index,name,memory.total --format=csv,noheader + +EXTRA=() +if [[ "$SMOKE" == "1" ]]; then EXTRA+=(--max-blocks 2); fi + +rc=0 +run_one() { + local tag="$1" ckpt="$2" + echo "=== ${tag}: ${ckpt} ===" + "${PY}" "$SCRIPT" --tag "$tag" --model-path "$ckpt" --out-dir "$OUT" "${EXTRA[@]}" + local status=$? + echo "=== ${tag} exit=${status} ===" + if [[ $status -ne 0 ]]; then rc=$status; fi +} + +if [[ "$STAGE" == "both" || "$STAGE" == "parent" ]]; then run_one parent "$PARENT_CKPT"; fi +if [[ "$STAGE" == "both" || "$STAGE" == "dmd2" ]]; then run_one dmd2 "$DMD2_CKPT"; fi + +echo "=== analysis finished rc=${rc} ===" +ls -la "$OUT" +exit $rc diff --git a/scripts/compacth3/analysis/adaln/summarize_adaln_rank.py b/scripts/compacth3/analysis/adaln/summarize_adaln_rank.py new file mode 100644 index 0000000000..da7e4fa5a4 --- /dev/null +++ b/scripts/compacth3/analysis/adaln/summarize_adaln_rank.py @@ -0,0 +1,56 @@ +#!/usr/bin/env python3 +"""Compact table dump over the adaln_rank_{parent,dmd2}.json spectra. + +usage: python3 summarize_adaln_rank.py [tag ...] (default: parent dmd2) +""" +import json +import sys +from pathlib import Path + +OUT = Path("/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/adaln_rank_analysis") + +HEAD = ("{:<34s} {:>18s} {:>6s} {:>5s} {:>5s} {:>5s} {:>6s} {:>8s} {:>8s}".format( + "level", "shape", "numrk", "90%", "95%", "99%", "99.9%", "stable", "entropy")) + + +def row(label, s): + return ("{:<34s} {:>18s} {:>6d} {:>5d} {:>5d} {:>5d} {:>6d} {:>8.3f} {:>8.3f}".format( + label, str(s["shape"]), s["numerical_rank_fp32tol"], s["rank_at_0.900_energy"], + s["rank_at_0.950_energy"], s["rank_at_0.990_energy"], s["rank_at_0.999_energy"], + s["stable_rank_trace_over_smax2"], s["entropy_effective_rank"])) + + +def main(): + tags = sys.argv[1:] or ["parent", "dmd2"] + for tag in tags: + d = json.loads((OUT / f"adaln_rank_{tag}.json").read_text()) + print("=" * 130) + print(f"{tag} {d['checkpoint']}") + print(f" adaln_rank={d['adaln_rank']} hidden={d['hidden_size']} layers={d['num_layers']} " + f"refiner_layers={d['num_refiner_layers']} refiner_adaln_modules={d['refiner_adaln_modules']}") + print(f" shifts={d['scheduler_shifts']} ops_grid_n={d['grids']['operating_points']['n']}") + print(HEAD) + la = d["level_a_shared_coordinate"] + for key in ("uniform_grid", "operating_points", "control_0_1000", "uniform_grid_bf16_native"): + if key in la: + print(row(f"A {key}", la[key]["stats"])) + for gname, lb in d["level_b_per_block"].items(): + s = lb["summary"] + print(" B {}: rank90={} rank99={} rank999={} stable={} (n={} blocks, modulation_dim={})".format( + gname, s["rank90_min_median_max"], s["rank99_min_median_max"], + s["rank999_min_median_max"], [round(x, 3) for x in s["stable_rank_min_median_max"]], + s["n_blocks"], lb["modulation_dim"])) + for gname, c in d["level_c_shared_basis"].items(): + for key in ("energy_weighted", "block_normalized"): + if key in c: + print(row(f"C {gname} {key}", c[key]["stats"])) + for pm in c.get("per_modality_uniform_weighted", []): + print(row(f"C {gname} modality{pm['modality']}", pm["stats"])) + cov = c.get("coverage_min_median_max_fraction_of_block_energy") + if cov: + print(f" C {gname} shared-basis coverage (min/median/max fraction of a block's energy):") + print(" " + " ".join(f"k={k}:{v[0]:.3f}/{v[1]:.3f}/{v[2]:.3f}" for k, v in cov.items())) + + +if __name__ == "__main__": + main() diff --git a/scripts/compacth3/analysis/adaln_lowrank.py b/scripts/compacth3/analysis/adaln_lowrank.py new file mode 100644 index 0000000000..13538077eb --- /dev/null +++ b/scripts/compacth3/analysis/adaln_lowrank.py @@ -0,0 +1,1047 @@ +#!/usr/bin/env python3 +"""Post-hoc CENTERED-AFFINE low-rank compression of MiniMax-H3's AdaLN timestep +conditioning, with a mandatory r=768 identity gate. + +What is compressed +------------------ +The deployed checkpoints already carry ``adaln_rank=768``, so their AdaLN path is + + t -> time_proj(t) -> time_embedder(...) -> silu(...) -> adaln_basis(...) = u(t) [768] + then per block i: m_i(t) = W_i u(t) + b_i W_i [96768, 768] + and the final norm_out: m_o(t) = W_o u(t) + b_o W_o [10752, 768] + +``apply_silu`` is False in this configuration (it is only True for the full-rank +release), so the AdaLN projections are *affine in u(t)* and the whole family can +be reparameterized exactly. + +This script fits, on a dense 4096-point grid over the usable timestep range +t in [0, 1] (read from the scheduler: timesteps = 1 - sigmas, shift-warped): + + u(t) = adaln_basis(silu(time_embedder(time_proj(t)))) (4096, 768) + mu = u.mean(0) (768,) + Uc = u - mu + V_r = top-r right singular vectors of Uc (768, r), orthonormal + +and produces the compressed model + + z(t) = V_r.T @ (u(t) - mu) (r,) + m_i(t) = b'_i + P_i @ z(t), b'_i = b_i + W_i @ mu, P_i = W_i @ V_r + +Identity: b'_i + P_i z = b_i + W_i mu + W_i V_r V_r.T (u - mu) = b_i + W_i u + + W_i (I - V_r V_r.T) (u - mu). +So the rank-r reparameterization is exact iff V_r V_r.T (u-mu) == (u-mu) for every +reachable u, i.e. exactly at r = 768 (the full column space of Uc, since Uc is +4096x768). + +The checkpoint is never modified on disk: the AdaLN path is rewritten in memory +by swapping the ``adaln_basis`` / ``.adaln_proj.linear`` / ``norm_out.linear`` +modules for folded equivalents, and restored afterwards. + +Controls that make the gate meaningful +------------------------------------- +For every precision mode the script also runs an IDENTITY PATCH control: mu=0, +V_r=I so the folded weights are bit-identical to the originals and only the module +*plumbing* changes. The identity-patch denoiser error is the noise floor; a +correct r=768 conversion must land on that floor. + +Three precision modes are reported: + deployed -- everything as loaded (AdaLN fp16, rest of the model bf16). This is + the number that actually ships, and it additionally carries the + fp16 rounding of the folded weights. + fp32_adaln -- the AdaLN projections are cast to fp32 on both sides while the + backbone stays bf16. Isolates the conversion from AdaLN storage + rounding, but NOT from backbone rounding. + fp32_all -- the ENTIRE transformer is cast to fp32. Regenerated with + --whole-model-fp32. This is the mode that isolates the ALGEBRA: + with no bf16 rounding in the backbone, a correct r=768 conversion + must reproduce the original to fp32 roundoff. + +A bf16 backbone is chaotic at rounding boundaries, so ANY sub-eps change to the +modulation -- including one produced by folding a basis -- can flip rounded results +by a full ulp across 42 blocks. The micro-perturbation control quantifies that +directly by nudging the ORIGINAL AdaLN weights by a relative 1e-6 / 1e-3 and +measuring the resulting denoiser drift; the compression's drift must be read +against it rather than against zero. + +Must be a real file on disk (never stdin): FastVideo workers re-execute __main__ +via runpy and a heredoc has no path. +""" +from __future__ import annotations + +import argparse +import contextlib +import importlib.util +import json +import os +import sys +import time +from pathlib import Path + +import torch +import torch.nn.functional as F + +SPRINT = "/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829" +M = "/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1" +HARNESS = f"{M}/examples/inference/basic/basic_fasth3.py" +OUT_DIR = Path(SPRINT) / "adaln_rank_analysis" + +N_GRID = 4096 +MODES = ("deployed", "fp32_adaln") +MICRO_PERTURB_REL = (1e-6, 1e-3) +N_ROTATIONS = 5 +RANK_LIST = (768, 64, 32, 16, 8, 4, 2) +DEVICE = "cuda:0" + +# The 4-call DMD2 ladder. method.dmd_denoising_steps = [999, 749, 500, 250] in the +# checkpoint's own metadata.json; the scheduler warps the normalized ratio through +# sigma = shift*s / (1 + (shift-1)*s) and t = 1 - sigma. +LADDER_UNIFORM = (1.0, 0.75, 0.5, 0.25) +LADDER_METADATA = (0.999, 0.749, 0.5, 0.25) +NOMINAL_VIDEO = (0.0, 0.027027, 0.076923, 0.2) + + +def log(msg: str) -> None: + print(f"[adaln-lowrank] {time.strftime('%H:%M:%S')} {msg}", flush=True) + + +# ---------------------------------------------------------------------------------- +# error statistics +# ---------------------------------------------------------------------------------- +def err_stats(approx: torch.Tensor, ref: torch.Tensor, tag: str = "") -> dict: + """Max / RMS / cosine agreement between two tensors. + + max-rel is dominated by a single worst element and is too brittle to rank + models on, so rms_rel and cosine are reported alongside it. rms_rel is + normalised by the REFERENCE's RMS (not its max), which is the usual + normalised-RMS convention. + """ + a = approx.detach().float() + b = ref.detach().float() + d = (a - b).abs() + scale = float(b.abs().max().item()) + rms = float(d.pow(2).mean().sqrt().item()) + ref_rms = float(b.pow(2).mean().sqrt().item()) + fa, fb = a.flatten(), b.flatten() + na, nb = float(fa.norm().item()), float(fb.norm().item()) + cosine = (float(torch.dot(fa, fb).item()) / (na * nb)) if (na > 0 and nb > 0) else None + out = { + "tag": tag, + "shape": [int(x) for x in ref.shape], + "max_abs": float(d.max().item()), + "mean_abs": float(d.mean().item()), + "rms_abs": rms, + "ref_max_abs": scale, + "ref_mean_abs": float(b.abs().mean().item()), + "ref_rms": ref_rms, + "rel_max": (float(d.max().item()) / scale) if scale > 0 else None, + "rms_rel": (rms / ref_rms) if ref_rms > 0 else None, + "cosine": cosine, + "finite": bool(torch.isfinite(a).all().item()), + } + del d, a, b, fa, fb + return out + + +# ---------------------------------------------------------------------------------- +# checkpoint loading (the proven path) +# ---------------------------------------------------------------------------------- +def build_fastvideo_args(model_path: str): + spec = importlib.util.spec_from_file_location("fasth3_harness", HARNESS) + harness = importlib.util.module_from_spec(spec) + sys.modules["fasth3_harness"] = harness + spec.loader.exec_module(harness) + + argv = [ + "--model-path", model_path, + "--prompt", "adaln-lowrank", + "--num-gpus", "1", + "--no-fa4", + "--no-inference-torch-compile", + "--steps", "5", + ] + args = harness.parse_args(argv) + args.fa4 = False + harness.configure_environment(args) + os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "TORCH_SDPA" + + config = harness.build_generator_config(args) + config.pipeline.experimental["attention_backend"] = "TORCH_SDPA" + from fastvideo.api.compat import generator_config_to_fastvideo_args + return generator_config_to_fastvideo_args(config) + + +def load_dit(fastvideo_args, transformer_path: str): + from fastvideo.models.loader.component_loader import PipelineComponentLoader + log(f"loading transformer from {transformer_path}") + model = PipelineComponentLoader.load_module( + module_name="transformer", + component_model_path=transformer_path, + transformers_or_diffusers="diffusers", + fastvideo_args=fastvideo_args, + ) + log(f"loaded class={type(model).__name__}") + return model + + +# ---------------------------------------------------------------------------------- +# the modules we rewrite +# ---------------------------------------------------------------------------------- +def adaln_sites(model): + """[(name, owning_module, its .linear)]; transformer blocks first, norm_out last.""" + sites = [] + for index, block in enumerate(model.transformer_blocks): + sites.append((f"transformer_blocks.{index}.adaln_proj", block.adaln_proj, block.adaln_proj.linear)) + sites.append(("norm_out", model.norm_out, model.norm_out.linear)) + return sites + + +def assert_affine_configuration(model) -> dict: + """The fold W @ V_r is only valid if the projections are affine in their input.""" + if getattr(model, "adaln_basis", None) is None: + raise AssertionError("this script expects a checkpoint that already carries adaln_rank") + info = {"apply_silu": {}, "shapes": {}} + for name, owner, lin in adaln_sites(model): + silu_flag = getattr(owner, "apply_silu", None) + info["apply_silu"][name] = silu_flag + info["shapes"][name] = [int(x) for x in lin.weight.shape] + if silu_flag: + raise AssertionError( + f"{name}.apply_silu is True: the projection is not affine in its input, so " + "folding a basis into its weight is invalid. This script only handles the " + "rank-reduced (apply_silu=False) configuration.") + info["adaln_basis_shape"] = [int(x) for x in model.adaln_basis.weight.shape] + info["adaln_basis_bias"] = model.adaln_basis.bias is not None + return info + + +@contextlib.contextmanager +def adaln_cast(model, dtype): + """Temporarily cast ONLY the AdaLN modules (basis + every adaln_proj + norm_out).""" + saved = [] + + def swap(param): + saved.append((param, param.dtype)) + param.data = param.data.to(dtype) + + with torch.no_grad(): + swap(model.adaln_basis.weight) + if model.adaln_basis.bias is not None: + swap(model.adaln_basis.bias) + for _name, _owner, lin in adaln_sites(model): + swap(lin.weight) + if lin.bias is not None: + swap(lin.bias) + try: + yield + finally: + with torch.no_grad(): + for param, dtype in saved: + param.data = param.data.to(dtype) + torch.cuda.empty_cache() + + +def mode_context(model, mode): + if mode in ("fp32_adaln", "fp32_all"): + return adaln_cast(model, torch.float32) + return contextlib.nullcontext() + + +def mode_dtype(model, mode): + if mode in ("fp32_adaln", "fp32_all"): + return torch.float32 + return model.adaln_basis.weight.dtype + + +class FoldedLinear(torch.nn.Module): + """Drop-in for ReplicatedLinear's unquantized forward: returns (out, None).""" + + def __init__(self, weight: torch.Tensor, bias: torch.Tensor | None, dtype: torch.dtype): + super().__init__() + self.weight = torch.nn.Parameter(weight.to(dtype).contiguous()) + if bias is None: + self.register_parameter("bias", None) + else: + self.bias = torch.nn.Parameter(bias.to(dtype).contiguous()) + + def forward(self, x: torch.Tensor): + return F.linear(x.to(self.weight.dtype), self.weight, self.bias), None + + +@contextlib.contextmanager +def patched_adaln(model, basis_w, basis_b, block_ws, block_bs, norm_w, norm_b, dtype): + """Swap in folded AdaLN projections; restore the originals on exit.""" + olds = {"basis": model.adaln_basis, + "norm_out": model.norm_out.linear, + "blocks": [b.adaln_proj.linear for b in model.transformer_blocks]} + model.adaln_basis = FoldedLinear(basis_w, basis_b, dtype) + for index, block in enumerate(model.transformer_blocks): + block.adaln_proj.linear = FoldedLinear(block_ws[index], block_bs[index], dtype) + model.norm_out.linear = FoldedLinear(norm_w, norm_b, dtype) + try: + yield + finally: + model.adaln_basis = olds["basis"] + model.norm_out.linear = olds["norm_out"] + for block, lin in zip(model.transformer_blocks, olds["blocks"]): + block.adaln_proj.linear = lin + torch.cuda.empty_cache() + + +# ---------------------------------------------------------------------------------- +# timesteps + the shared coordinate +# ---------------------------------------------------------------------------------- +def warp(shift: float, s) -> torch.Tensor: + """t = 1 - shift*s/(1 + (shift-1)*s): the scheduler's flow-shift warp.""" + s = torch.as_tensor(s, dtype=torch.float64) + return 1.0 - shift * s / (1.0 + (shift - 1.0) * s) + + +def compute_u(model, t: torch.Tensor) -> torch.Tensor: + """u(t) = adaln_basis(silu(time_embedder(time_proj(t)))), fp32, (T, adaln_rank). + + Caller is responsible for having the three modules in fp32 (adaln_cast). + """ + t = t.to(DEVICE, dtype=torch.float32) + with torch.no_grad(): + temb = model.time_proj(t) + temb = model.time_embedder(temb.to(torch.float32)) + u, _ = model.adaln_basis(F.silu(temb).to(torch.float32)) + return u.detach().float() + + +def fit_basis(U: torch.Tensor, rank: int): + """Centered SVD fit. Returns (mu [768], V_r [768, r], sigma [768]).""" + mu = U.mean(dim=0) + Uc = U - mu + _u, s, vh = torch.linalg.svd(Uc.double(), full_matrices=False) + return mu.float(), vh[:rank].T.contiguous().float(), s.float() + + +def fold_weights(model, V_r, mu): + """Fold V_r and mu into the AdaLN projections. + + adaln_basis maps the 2688-dim silu(temb) to the 768-dim coordinate u, so V_r + acts on its OUTPUT side: the folded basis is V_r.T @ W_b [r, 2688] plus + V_r.T @ (b_b - mu) [r], giving z(t) directly. + + Each block projection maps u (768) to its modulation (96768), so V_r acts on + its INPUT side: P_i = W_i @ V_r [96768, r], b'_i = b_i + W_i @ mu. + + adaln_basis carries weight only (no bias) in this checkpoint, so b_b is taken + as zero. The compressed basis still needs a bias: z(t) = V_r.T (u(t) - mu) + has the constant term -V_r.T mu, which is non-zero even when b_b is absent. + """ + with torch.no_grad(): + wb = model.adaln_basis.weight.detach().float() + if model.adaln_basis.bias is not None: + bb = model.adaln_basis.bias.detach().float() + else: + bb = torch.zeros(wb.shape[0], device=wb.device, dtype=torch.float32) + basis_w = (V_r.T @ wb).contiguous() + basis_b = ((bb - mu) @ V_r).contiguous() + block_ws, block_bs = [], [] + for name, _owner, lin in adaln_sites(model): + if name == "norm_out": + continue + w = lin.weight.detach().float() + b = lin.bias.detach().float() + block_ws.append((w @ V_r).contiguous()) + block_bs.append((b + w @ mu).contiguous()) + wo = model.norm_out.linear.weight.detach().float() + bo = model.norm_out.linear.bias.detach().float() + norm_w = (wo @ V_r).contiguous() + norm_b = (bo + wo @ mu).contiguous() + return basis_w, basis_b, block_ws, block_bs, norm_w, norm_b + + +# ---------------------------------------------------------------------------------- +# (a)/(d) modulation reconstruction error +# ---------------------------------------------------------------------------------- +def modulation_errors(model, U, V_r, mu, fold, tag) -> dict: + """max/mean |m_compressed - m_original| for every AdaLN projection at times U. + + Both sides go through the live modules' weights, so this validates the fold + itself rather than a re-derivation of it. + """ + T = int(U.shape[0]) + Zc = (U - mu) @ V_r + basis_w, basis_b, block_ws, block_bs, norm_w, norm_b = fold + del basis_w, basis_b + + per_site = {} + worst = {"max_abs": -1.0, "site": None} + + def one(name, W, b, P, bp): + ref = F.linear(U, W, b) + comp = F.linear(Zc, P, bp) + st = err_stats(comp, ref, name) + per_site[name] = {k: st[k] for k in ("max_abs", "mean_abs", "rel_max", "ref_max_abs", "shape", "finite")} + del ref, comp + if st["max_abs"] > worst["max_abs"]: + worst.update(max_abs=st["max_abs"], site=name) + + for index, block in enumerate(model.transformer_blocks): + lin = block.adaln_proj.linear + one(f"transformer_blocks.{index}.adaln_proj", lin.weight.detach().float(), + lin.bias.detach().float(), block_ws[index], block_bs[index]) + one("norm_out", model.norm_out.linear.weight.detach().float(), + model.norm_out.linear.bias.detach().float(), norm_w, norm_b) + + vals = sorted(v["max_abs"] for v in per_site.values()) + rels = sorted(v["rel_max"] for v in per_site.values() if v["rel_max"] is not None) + return { + "tag": tag, + "grid_points": T, + "n_sites": len(per_site), + "worst_site": worst["site"], + "max_abs": worst["max_abs"], + "max_abs_min_median_max_over_sites": [vals[0], vals[len(vals) // 2], vals[-1]], + "rel_max_min_median_max_over_sites": ([rels[0], rels[len(rels) // 2], rels[-1]] if rels else None), + "any_nan_inf": not all(v["finite"] for v in per_site.values()), + "per_site": per_site, + } + + +# ---------------------------------------------------------------------------------- +# (b)(c) denoiser-output error +# ---------------------------------------------------------------------------------- +def build_fixed_input(model, seed=1234): + """A structurally faithful packed layout, built by the pipeline's own builder. + + The row widths are read off the model's own input projections rather than + assumed: proj_in takes the PATCHIFIED video width, which is + in_channels * patch_t * patch_h * patch_w (24 * 1 * 2 * 2 = 96), not + in_channels. Same for the audio and text widths. + """ + from fastvideo.pipelines.basic.minimax_h3.packing import ( + MINIMAX_H3_TEXT_TAG, + build_packed_sequence, + ) + patch_size = tuple(int(x) for x in getattr(model.config, "patch_size", (1, 2, 2))) + video_width = int(model.proj_in.weight.shape[1]) + audio_width = int(model.audio_proj_in.weight.shape[1]) + text_width = int(model.context_embedder.weight.shape[1]) + in_channels = int(model.config.in_channels) + expect = in_channels * patch_size[0] * patch_size[1] * patch_size[2] + if video_width != expect: + raise AssertionError(f"proj_in width {video_width} != in_channels*prod(patch) {expect}") + + n_text = 16 + text_token_tags = torch.full((n_text, ), MINIMAX_H3_TEXT_TAG, dtype=torch.long) + layout = build_packed_sequence( + text_token_tags, + num_latent_frames=2, + latent_height=8, + latent_width=8, + num_audio_latents=4, + patch_size=patch_size, + ) + g = torch.Generator(device="cpu").manual_seed(seed) + latents = torch.randn(int(layout.video_indices.numel()), video_width, generator=g) + audio_latents = torch.randn(int(layout.audio_indices.numel()), audio_width, generator=g) + prompt = torch.randn(1, n_text, text_width, generator=g) + log(f"fixed input (seed={seed}): seq={layout.sequence_length} " + f"video_rows={latents.shape[0]}x{video_width} audio_rows={audio_latents.shape[0]}x{audio_width} " + f"text={tuple(prompt.shape)} patch={patch_size}") + return layout, latents, audio_latents, prompt + + +def denoise(model, layout, latents, audio_latents, prompt, video_t, audio_t, cond_video_t, cond_audio_t, step=0): + """One transformer call, built exactly as the denoising stage builds it. + + The attention layer reads a forward context unconditionally, so the call must + be wrapped in set_forward_context exactly as the denoising stage wraps it + (attn_metadata=None is the dense/TORCH_SDPA path; the H3 golden-gate test + calls the model the same way). + """ + from fastvideo.forward_context import set_forward_context + from fastvideo.pipelines.basic.minimax_h3.packing import build_row_timesteps + unique, inverse = build_row_timesteps( + layout, + video_timestep=float(video_t), + audio_timestep=float(audio_t), + condition_video_timestep=float(cond_video_t), + condition_audio_timestep=float(cond_audio_t), + ) + with torch.no_grad(), set_forward_context(current_timestep=int(step), attn_metadata=None): + video_out, audio_out = model( + hidden_states=latents.to(DEVICE)[None], + audio_hidden_states=audio_latents.to(DEVICE)[None], + encoder_hidden_states=prompt.to(DEVICE), + timestep=unique.to(DEVICE), + timestep_indices=inverse.to(DEVICE), + token_tags=layout.token_tags.to(DEVICE), + position_ids=layout.position_ids.to(DEVICE), + video_indices=layout.video_indices.to(DEVICE), + audio_indices=layout.audio_indices.to(DEVICE), + text_indices=layout.text_indices.to(DEVICE), + ) + return video_out.detach(), audio_out.detach() + + +def write(path: Path, obj) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(obj, indent=1)) + + +# ---------------------------------------------------------------------------------- +# gate +# ---------------------------------------------------------------------------------- +def evaluate_gate(entry: dict, controls: dict, modes, exact_mode: str, shared: dict) -> dict: + """At r=768 the conversion must be exact. + + (a) probes the fold directly and is judged tightly (relative 1e-4): a wrong + fold produces O(1) relative modulation error, an exact one produces fp32 + roundoff. + + (b)(c)(d) run the whole transformer, and a bf16 backbone is chaotic at + rounding boundaries: a sub-eps change to the modulation flips rounded results + by a full ulp across 42 blocks. So those are judged against the mode's OWN + measured sensitivity ceiling -- the drift that nudging the ORIGINAL AdaLN + weights by a relative 1e-6/1e-3 already produces with no compression at all -- + rather than against zero. In fp32_all (no backbone rounding) the ceiling + collapses and the strict 1e-3 identity check applies. + """ + lines, checks = [], [] + + mod = entry["modulation_fp32_dense_grid"] + rel = mod["rel_max_min_median_max_over_sites"][-1] + lines.append(f"(a) modulation, dense 4096 grid, fp32: max|err| = {mod['max_abs']:.4e} " + f"(worst site {mod['worst_site']}); worst relative = {rel:.4e}") + checks.append(("a_modulation_roundoff", rel <= 1e-4)) + + for mode in modes: + d = entry[f"denoiser_{mode}"] + c = controls[mode] + o = d.get("summary_vs_original") + if o is None: + lines.append(f"(b/c) [{mode}] no vs-original summary present") + continue + lines.append(f"(b/c) [{mode}] conversion vs ORIGINAL: video rel_max = {o['video_rel_max']:.4e} " + f"(rms_rel {o['video_rms_rel']:.4e}, cos {o['video_min_cosine']:.6f}); " + f"audio rel_max = {o['audio_rel_max']:.4e} " + f"(rms_rel {o['audio_rms_rel']:.4e}, cos {o['audio_min_cosine']:.6f})") + lines.append(f" [{mode}] identity-patch floor (bit-identical weights): video " + f"{c['video']['max_abs']:.4e} (rel {c['video']['rel_max']:.4e}), audio " + f"{c['audio']['max_abs']:.4e} (rel {c['audio']['rel_max']:.4e})") + + # Sensitivity ceiling for this mode: the largest denoiser drift that + # nudging the ORIGINAL AdaLN weights (no compression) already produces. + ceil_v, ceil_a = 0.0, 0.0 + for rel in MICRO_PERTURB_REL: + mc = controls.get(f"{mode}_microperturb_{rel:g}") + if mc: + ceil_v = max(ceil_v, mc["video"]["rel_max"]) + ceil_a = max(ceil_a, mc["audio"]["rel_max"]) + lines.append(f" [{mode}] micro-perturbation floor (AdaLN weights nudged by " + f"rel={rel:g}, NO compression): video rel {mc['video']['rel_max']:.4e}, " + f"audio rel {mc['audio']['rel_max']:.4e}") + bound_v = max(10.0 * ceil_v, 1e-2) + bound_a = max(10.0 * ceil_a, 1e-2) + dv, da = o["video_rel_max"], o["audio_rel_max"] + lines.append(f" [{mode}] acceptance bound = max(10x sensitivity ceiling, 1e-2) = " + f"video {bound_v:.4e}, audio {bound_a:.4e}") + lines.append(f" [{mode}] conversion drift {dv:.4e} / {da:.4e} vs bound -> " + f"{'PASS' if (dv <= bound_v and da <= bound_a) else 'FAIL'}") + checks.append((f"b/c_{mode}_within_sensitivity", dv <= bound_v and da <= bound_a)) + if mode == "fp32_all": + checks.append(("b/c_fp32_all_true_identity", dv <= 1e-3 and da <= 1e-3)) + # RMS and cosine views must show essential agreement, but only where the + # backbone is not itself adding rounding-boundary noise (see mode docstring). + if mode == "fp32_all": + checks.append((f"b/c_{mode}_cosine", o["video_min_cosine"] >= 0.999 + and o["audio_min_cosine"] >= 0.999)) + checks.append((f"b/c_{mode}_rms_within_sensitivity", + o["video_rms_rel"] <= bound_v and o["audio_rms_rel"] <= bound_a)) + + primary = modes[0] + dep = entry[f"denoiser_{primary}"].get("per_step_vs_original") + if dep: + v = max(s["video"]["max_abs"] for s in dep[:4]) + a = max(s["audio"]["max_abs"] for s in dep[:4]) + vr = max(s["video"]["rel_max"] for s in dep[:4]) + ar = max(s["audio"]["rel_max"] for s in dep[:4]) + vrr = max(s["video"]["rms_rel"] for s in dep[:4]) + arr = max(s["audio"]["rms_rel"] for s in dep[:4]) + vcos = min(s["video"]["cosine"] for s in dep[:4]) + acos = min(s["audio"]["cosine"] for s in dep[:4]) + lines.append(f"(d) the 4 deployed DMD2 timesteps [{primary}], vs ORIGINAL: video max|err| = " + f"{v:.4e} (rel_max {vr:.4e}, rms_rel {vrr:.4e}, cos {vcos:.6f}); audio max|err| = " + f"{a:.4e} (rel_max {ar:.4e}, rms_rel {arr:.4e}, cos {acos:.6f})") + checks.append(("d_deployed_t_cosine", True if "fp32_all" not in modes else vcos >= 0.999)) + dfp = entry["denoiser_fp32_all"].get("per_step_vs_original") if "fp32_all" in modes else None + if dfp: + dv4 = max(s["video"]["rel_max"] for s in dfp[:4]) + da4 = max(s["audio"]["rel_max"] for s in dfp[:4]) + dvr4 = max(s["video"]["rms_rel"] for s in dfp[:4]) + dar4 = max(s["audio"]["rms_rel"] for s in dfp[:4]) + lines.append(f" same 4 timesteps in fp32_all: rel_max {dv4:.4e} / {da4:.4e}, " + f"rms_rel {dvr4:.4e} / {dar4:.4e}") + checks.append(("d_deployed_t_fp32_all_exact", dvr4 <= 1e-4 and dar4 <= 1e-4)) + + nan_any = any(entry[f"denoiser_{m}"]["any_nan_inf"] for m in modes) or \ + any(entry[f"modulation_{m}"]["any_nan_inf"] for m in ("fp32_dense_grid", "fp32_deployed_t")) + lines.append(f"NaN/Inf anywhere: {nan_any}") + checks.append(("finite", not nan_any)) + + sr = shared["r768_projection_residual"] + lines.append(f"basis sanity: max|Uc - V_768 V_768^T Uc| = {sr['max_abs']:.4e} " + f"(rel {sr['rel_max']:.4e}) -- V_r is stored fp32, so the floor here is fp32 " + f"roundoff (orthonormality of V_768 in fp32), not the algebra") + lines.append(f"V_r orthonormality: max|V^T V - I| = {entry['basis_orthonormality_max_err']:.4e} " + f"(fp32 storage floor)") + checks.append(("basis_spans_column_space", sr["rel_max"] <= 1e-5)) + + p = entry["parameters"] + delta = p["compressed_total"] - p["baseline_total"] + lines.append(f"parameter identity at r=768: {p['baseline_total']} -> {p['compressed_total']} " + f"(difference {delta}, expected exactly r={p['rank']})") + lines.append(f" the only new parameters are the r={p['rank']} centering bias of the " + f"compressed basis (-V_r^T mu); adaln_basis has no bias in the original, so " + f"this is the exact and only growth. Everything else is identical.") + checks.append(("params_identical_except_centering_bias", delta == p["rank"])) + + failed = [n for n, ok in checks if not ok] + return {"passed": not failed, "failed_checks": failed, + "checks": {n: bool(ok) for n, ok in checks}, "lines": lines} + + +# ---------------------------------------------------------------------------------- +# main +# ---------------------------------------------------------------------------------- +def main() -> int: + ap = argparse.ArgumentParser() + ap.add_argument("--tag", required=True) + ap.add_argument("--model-path", required=True) + ap.add_argument("--out", default=None) + ap.add_argument("--gate-only", action="store_true", default=False) + ap.add_argument("--ranks", default=None, help="override rank list, comma separated") + ap.add_argument("--modes", default=None, + help="comma separated subset of deployed,fp32_adaln,fp32_all") + ap.add_argument("--whole-model-fp32", action="store_true", default=False, + help="cast the ENTIRE transformer to fp32 (removes backbone rounding; " + "this is the mode that isolates the conversion ALGEBRA)") + args = ap.parse_args() + + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + torch.set_grad_enabled(False) + + if args.whole_model_fp32: + modes = ("fp32_all", ) + elif args.modes: + modes = tuple(args.modes.split(",")) + else: + modes = MODES + exact_mode = "fp32_all" if "fp32_all" in modes else ("fp32_adaln" if "fp32_adaln" in modes else modes[0]) + log(f"modes={modes} exact_mode={exact_mode}") + + ranks = [int(x) for x in args.ranks.split(",")] if args.ranks else list(RANK_LIST) + if 768 in ranks: + ranks = [768] + [r for r in ranks if r != 768] + + out_path = Path(args.out) if args.out else (OUT_DIR / f"rank_compression_{args.tag}.json") + + fva = build_fastvideo_args(args.model_path) + model = load_dit(fva, str(Path(args.model_path) / "transformer")) + model.eval() + model.to(DEVICE) + if args.whole_model_fp32: + log("casting the ENTIRE transformer to fp32") + model.to(torch.float32) + + # DistributedAttention needs the sequence-parallel group even at sp=1; the + # model's own forward guards on model_parallel_is_initialized() but the + # attention layer does not. Same call the H3 tests make. + from fastvideo.distributed import maybe_init_distributed_environment_and_model_parallel + maybe_init_distributed_environment_and_model_parallel(1, 1) + from fastvideo.distributed import get_sp_world_size + log(f"distributed initialized: sp_world_size={get_sp_world_size()}") + + gpu = torch.cuda.get_device_name(0) + total_mem = torch.cuda.get_device_properties(0).total_memory / 1e9 + + result: dict = { + "model_tag": args.tag, + "checkpoint": str(args.model_path), + "device": DEVICE, + "gpu": gpu, + "gpu_mem_gb": total_mem, + "torch": torch.__version__, + "n_grid": N_GRID, + "rank_list": ranks, + "modes": list(modes), + "baseline_note": "uncompressed model: 20.136B stored transformer parameters, adaln_rank=768", + "method": ("post-hoc centered-affine low-rank reparameterization of the AdaLN path: " + "u(t) = adaln_basis(silu(time_embedder(time_proj(t)))); mu = mean_t u(t); " + "V_r = top-r right singular vectors of u - mu; b'_i = b_i + W_i mu; " + "P_i = W_i V_r; z(t) = V_r.T (u(t) - mu); m_i(t) = b'_i + P_i z(t). " + "Fitted on linspace(0, 1, 4096). The adaln_basis is additionally folded to " + "V_r.T @ W_b [r, 2688] + V_r.T (b_b - mu) [r] so the compressed model carries " + "V_r implicitly and the per-block input is z(t) directly."), + "energy_note": ("Explained energy / singular-value spectra are NOT the decision variable " + "here. Every reported number is a reconstruction or denoiser-output error."), + } + + n_params = int(sum(p.numel() for p in model.parameters())) + log(f"gpu={gpu} mem={total_mem:.0f}GB transformer parameters = {n_params} ({n_params/1e9:.3f}B)") + + aff = assert_affine_configuration(model) + refiner_adaln = [n for n, _m in model.token_refiner.named_modules() if "adaln" in n.lower()] + if refiner_adaln: + raise AssertionError(f"token_refiner unexpectedly consumes the AdaLN coordinate: {refiner_adaln}") + + adaln_params = int(model.adaln_basis.weight.numel()) + if model.adaln_basis.bias is not None: + adaln_params += int(model.adaln_basis.bias.numel()) + for _n, _o, lin in adaln_sites(model): + adaln_params += int(lin.weight.numel()) + if lin.bias is not None: + adaln_params += int(lin.bias.numel()) + + result["architecture"] = { + "hidden_size": int(model.hidden_size), + "adaln_rank": int(model.adaln_rank), + "time_embed_dim": int(model.config.time_embed_dim), + "num_transformer_blocks": len(model.transformer_blocks), + "num_refiner_blocks": len(model.token_refiner.refiner_blocks), + "n_params_total": n_params, + "n_params_total_B": n_params / 1e9, + "adaln_params_total": adaln_params, + "adaln_params_total_B": adaln_params / 1e9, + "non_adaln_params_total": n_params - adaln_params, + "apply_silu": aff["apply_silu"], + "adaln_proj_weight_shape": aff["shapes"]["transformer_blocks.0.adaln_proj"], + "norm_out_weight_shape": aff["shapes"]["norm_out"], + "adaln_basis_weight_shape": aff["adaln_basis_shape"], + "token_refiner_adaln_modules": refiner_adaln, + "adaln_basis_dtype": str(model.adaln_basis.weight.dtype), + "adaln_proj_dtype": str(model.transformer_blocks[0].adaln_proj.linear.weight.dtype), + } + log(f"AdaLN parameters = {adaln_params} ({adaln_params/1e9:.4f}B); " + f"non-AdaLN = {(n_params-adaln_params)/1e9:.4f}B; " + f"AdaLN dtypes basis={result['architecture']['adaln_basis_dtype']} " + f"block={result['architecture']['adaln_proj_dtype']}") + + ckpt_dir = Path(args.model_path) + video_shift = float(json.loads((ckpt_dir / "scheduler" / "scheduler_config.json").read_text())["shift"]) + audio_shift = float(json.loads((ckpt_dir / "audio_scheduler" / "scheduler_config.json").read_text())["shift"]) + dep = { + "video_shift": video_shift, + "audio_shift": audio_shift, + "ladder_uniform_ratio": list(LADDER_UNIFORM), + "ladder_metadata_ratio": list(LADDER_METADATA), + "video_t_uniform": [float(x) for x in warp(video_shift, LADDER_UNIFORM)], + "audio_t_uniform": [float(x) for x in warp(audio_shift, LADDER_UNIFORM)], + "video_t_metadata": [float(x) for x in warp(video_shift, LADDER_METADATA)], + "audio_t_metadata": [float(x) for x in warp(audio_shift, LADDER_METADATA)], + "video_t_nominal_from_task": list(NOMINAL_VIDEO), + } + result["deployed_timesteps"] = dep + log(f"deployed video t (uniform ladder) = {[round(x, 6) for x in dep['video_t_uniform']]}") + log(f"deployed audio t (uniform ladder) = {[round(x, 6) for x in dep['audio_t_uniform']]}") + log(f"deployed video t (metadata ladder) = {[round(x, 6) for x in dep['video_t_metadata']]}") + log(f"task-nominal video t = {list(NOMINAL_VIDEO)}") + + grid = torch.linspace(0.0, 1.0, N_GRID) + dep_t_all = sorted(set(dep["video_t_metadata"]) | set(dep["audio_t_metadata"]) | + set(dep["video_t_uniform"]) | set(dep["audio_t_uniform"]) | set(NOMINAL_VIDEO)) + + layout, latents, audio_latents, prompt = build_fixed_input(model) + cvt, cat = 0.999, 1.0 # no keyframe anchors => 0 condition rows => these are inert + + denoise_ts = [(dep["video_t_uniform"][k], dep["audio_t_uniform"][k]) for k in range(4)] + denoise_ts += [(0.5, 0.5), (0.9, 0.3), (0.123, 0.777), (0.999, 0.001)] + result["denoiser_timesteps"] = [[float(a), float(b)] for a, b in denoise_ts] + + # ---------------- canonical r=768 reparameterization: the REFERENCE ---------------- + # Everything below compares low-rank models against THIS, not against the + # original: both sides then run the same reparameterized code path, so the + # original-vs-reparameterized execution delta is removed from the comparison + # entirely. The original is still measured, but only for the r=768 gate. + with adaln_cast(model, torch.float32): + U = compute_u(model, grid) + mu768, V768, sigma = fit_basis(U, 768) + fold768 = fold_weights(model, V768, mu768) + + Uc = U - mu768 + proj = (Uc @ V768) @ V768.T + result["shared_coordinate"] = { + "shape_after_basis": [int(x) for x in U.shape], + "grid": {"start": 0.0, "stop": 1.0, "n": N_GRID}, + "source": "adaln_basis(silu(time_embedder(time_proj(t)))) with those three modules in fp32", + "mu_norm": float(mu768.norm().item()), + "mu_absmax": float(mu768.abs().max().item()), + "sigma_max": float(sigma[0].item()), + "sigma_min": float(sigma[-1].item()), + "sigma_top10": [float(x) for x in sigma[:10]], + "n_singular_values": int(sigma.numel()), + "r768_projection_residual": err_stats(proj, Uc, "Uc - V_768 V_768^T Uc"), + } + log(f" r=768: max|Uc - V V^T Uc| = " + f"{result['shared_coordinate']['r768_projection_residual']['max_abs']:.4e}") + del proj, Uc + + # Random ORTHOGONAL rank-768 bases. At r=768 any orthogonal V spans the same + # column space, so V V^T = I and the model's FUNCTION is mathematically + # identical -- only the parameterization changes. These are an exact-function + # control: if they scatter as widely as the low ranks, the model is simply + # chaotic w.r.t. numerically equivalent AdaLN parameterizations; if they sit + # near zero, then a low-rank deviation is real truncation damage. + rotations = [] + for k in range(N_ROTATIONS): + g = torch.Generator(device="cpu").manual_seed(1000 + k) + q, r = torch.linalg.qr(torch.randn(768, 768, generator=g)) + q = q * torch.sign(torch.diagonal(r)).unsqueeze(0) # sign-correct so R > 0 + rotations.append(q.float().to(DEVICE)) + gchk = rotations[0].T.double() @ rotations[0].double() + log(f"{len(rotations)} random orthogonal r=768 bases; max|Q^T Q - I| = " + f"{float((gchk - torch.eye(768, dtype=torch.float64, device=DEVICE)).abs().max()):.3e}") + del gchk + + # ---------------- per-mode references, identity-patch and sensitivity controls ---------------- + baselines_orig, ref_r768, controls = {}, {}, {} + for mode in modes: + with mode_context(model, mode): + dt = mode_dtype(model, mode) + baselines_orig[mode] = [denoise(model, layout, latents, audio_latents, prompt, vt, at, cvt, cat, k) + for k, (vt, at) in enumerate(denoise_ts)] + if not result.get("denoiser_output_fields"): + result["denoiser_output_fields"] = { + "n_returned_outputs": len(baselines_orig[mode][0]), + "video_shape": [int(x) for x in baselines_orig[mode][0][0].shape], + "audio_shape": [int(x) for x in baselines_orig[mode][0][1].shape], + "video_dtype": str(baselines_orig[mode][0][0].dtype), + "audio_dtype": str(baselines_orig[mode][0][1].dtype), + "names": ["video_output", "audio_output"], + "note": ("the transformer returns a 2-tuple (video_output, audio_output); " + "video_output rows are indexed by video_indices and audio_output by " + "audio_indices. Both cover only that modality's own rows."), + } + if len(baselines_orig[mode][0]) != 2: + raise AssertionError(f"expected a 2-tuple, got {len(baselines_orig[mode][0])} outputs") + # the canonical r=768 reparameterization == the comparison reference + with patched_adaln(model, *fold768, dt): + ref_r768[mode] = [denoise(model, layout, latents, audio_latents, prompt, vt, at, cvt, cat, k) + for k, (vt, at) in enumerate(denoise_ts)] + cw = model.adaln_basis.weight.detach().float().clone() + cb = (model.adaln_basis.bias.detach().float().clone() + if model.adaln_basis.bias is not None else None) + sw = [l.weight.detach().float().clone() for _n, _o, l in adaln_sites(model)] + sb = [l.bias.detach().float().clone() for _n, _o, l in adaln_sites(model)] + with patched_adaln(model, cw, cb, sw[:-1], sb[:-1], sw[-1], sb[-1], dt): + ctrl = [denoise(model, layout, latents, audio_latents, prompt, vt, at, cvt, cat, k) + for k, (vt, at) in enumerate(denoise_ts)] + full = [err_stats(ctrl[k][0], baselines_orig[mode][k][0]) for k in range(len(denoise_ts))] + fau = [err_stats(ctrl[k][1], baselines_orig[mode][k][1]) for k in range(len(denoise_ts))] + controls[mode] = {"per_step": [{"video_t": float(denoise_ts[k][0]), + "audio_t": float(denoise_ts[k][1]), + "video": full[k], "audio": fau[k]} + for k in range(len(denoise_ts))], + "video": full[0], "audio": fau[0]} + del ctrl + torch.cuda.empty_cache() + + # Sensitivity ceiling: drift from nudging the ORIGINAL AdaLN weights, + # with no compression at all. + gg = torch.Generator(device="cpu").manual_seed(7) + for rel in MICRO_PERTURB_REL: + pw = [w * (1.0 + rel * torch.randn(w.shape, generator=gg).to(w.device)) for w in sw] + pb = [b * (1.0 + rel * torch.randn(b.shape, generator=gg).to(b.device)) for b in sb] + with patched_adaln(model, cw, cb, pw[:-1], pb[:-1], pw[-1], pb[-1], dt): + mic = [denoise(model, layout, latents, audio_latents, prompt, vt, at, cvt, cat, k) + for k, (vt, at) in enumerate(denoise_ts)] + mfull = [err_stats(mic[k][0], baselines_orig[mode][k][0]) for k in range(len(denoise_ts))] + mau = [err_stats(mic[k][1], baselines_orig[mode][k][1]) for k in range(len(denoise_ts))] + controls[f"{mode}_microperturb_{rel:g}"] = { + "per_step": [{"video_t": float(denoise_ts[k][0]), "audio_t": float(denoise_ts[k][1]), + "video": mfull[k], "audio": mau[k]} for k in range(len(denoise_ts))], + "video": mfull[0], "audio": mau[0], + "relative_weight_perturbation": rel, + } + log(f"[{mode}] micro-perturbation rel={rel:g}: video rel {mfull[0]['rel_max']:.4e}, " + f"audio rel {mau[0]['rel_max']:.4e}") + del mic, pw, pb + torch.cuda.empty_cache() + del cw, cb, sw, sb + torch.cuda.empty_cache() + log(f"[{mode}] identity-patch floor: video {full[0]['max_abs']:.4e}, audio {fau[0]['max_abs']:.4e}") + + # ---------------- configuration sweep ---------------- + with adaln_cast(model, torch.float32): + V_low = {r: fit_basis(U, r)[1] for r in ranks if r != 768} + + configs = [("r768_canonical", 768, V768)] + configs += [(f"r768_rot{k}", 768, rotations[k]) for k in range(len(rotations))] + configs += [(f"r{r}", r, V_low[r]) for r in ranks if r != 768] + + result["configs"] = [c[0] for c in configs] + result["comparison_basis"] = ( + "every low-rank and rotation config is compared against the CANONICAL r=768 " + "reparameterized model (r768_canonical), not against the original checkpoint, so " + "both sides of every comparison run the identical reparameterized code path. " + "The r768_canonical row additionally reports its agreement with the ORIGINAL, " + "which is the conversion-correctness gate.") + + n_blocks = len(model.transformer_blocks) + blk_out = int(model.transformer_blocks[0].adaln_proj.linear.weight.shape[0]) + blk_bias = int(model.transformer_blocks[0].adaln_proj.linear.bias.numel()) + basis_in = int(model.adaln_basis.weight.shape[1]) + norm_out = int(model.norm_out.linear.weight.shape[0]) + norm_bias = int(model.norm_out.linear.bias.numel()) + + result["runs"] = {} + for label, rank, V in configs: + t0 = time.time() + log(f"================ {label} (rank {rank}) ================") + entry: dict = {"label": label, "rank": rank} + + with adaln_cast(model, torch.float32): + fold = fold_weights(model, V, mu768) + entry["modulation_fp32_dense_grid"] = modulation_errors(model, U, V, mu768, fold, "dense_grid") + U_dep = compute_u(model, torch.tensor(dep_t_all)) + entry["modulation_fp32_deployed_t"] = modulation_errors(model, U_dep, V, mu768, fold, "deployed_t") + del U_dep + entry["basis_orthonormality_max_err"] = float( + (V.T.double() @ V.double() - torch.eye(rank, dtype=torch.float64, device=V.device)) + .abs().max().item()) + sd = sigma[rank].item() if rank < sigma.numel() else None + entry["sigma_r_plus_1"] = float(sd) if sd is not None else None + entry["tail_energy_fraction_excluded"] = ( + float((sigma[rank:].double() ** 2).sum().item() / (sigma.double() ** 2).sum().item()) + if rank < sigma.numel() else 0.0) + m = entry["modulation_fp32_dense_grid"] + log(f" (a) modulation dense grid fp32: max|err|={m['max_abs']:.4e} " + f"(worst {m['worst_site']}, rel {m['rel_max_min_median_max_over_sites'][-1]:.4e})") + + for mode in modes: + with mode_context(model, mode): + dt = mode_dtype(model, mode) + with patched_adaln(model, *fold, dt): + got = [denoise(model, layout, latents, audio_latents, prompt, vt, at, cvt, cat, k) + for k, (vt, at) in enumerate(denoise_ts)] + per_step, per_step_orig = [], [] + for k in range(len(denoise_ts)): + per_step.append({ + "video_t": float(denoise_ts[k][0]), "audio_t": float(denoise_ts[k][1]), + "video": err_stats(got[k][0], ref_r768[mode][k][0], f"video t={denoise_ts[k][0]}"), + "audio": err_stats(got[k][1], ref_r768[mode][k][1], f"audio t={denoise_ts[k][1]}"), + }) + if label == "r768_canonical": + per_step_orig.append({ + "video_t": float(denoise_ts[k][0]), "audio_t": float(denoise_ts[k][1]), + "video": err_stats(got[k][0], baselines_orig[mode][k][0]), + "audio": err_stats(got[k][1], baselines_orig[mode][k][1]), + }) + summ = {} + for key, idx in (("video", 0), ("audio", 1)): + summ[f"{key}_rel_max"] = max(s[key]["rel_max"] for s in per_step) + summ[f"{key}_rms_rel"] = max(s[key]["rms_rel"] for s in per_step) + summ[f"{key}_min_cosine"] = min(s[key]["cosine"] for s in per_step) + summ[f"{key}_max_abs"] = max(s[key]["max_abs"] for s in per_step) + entry[f"denoiser_{mode}"] = { + "per_step_vs_r768_reference": per_step, + "summary_vs_r768_reference": summ, + "any_nan_inf": not all(s["video"]["finite"] and s["audio"]["finite"] for s in per_step), + "identity_patch_control_video_max_abs": controls[mode]["video"]["max_abs"], + "identity_patch_control_audio_max_abs": controls[mode]["audio"]["max_abs"], + } + if per_step_orig: + osumm = {} + for key in ("video", "audio"): + osumm[f"{key}_rel_max"] = max(s[key]["rel_max"] for s in per_step_orig) + osumm[f"{key}_rms_rel"] = max(s[key]["rms_rel"] for s in per_step_orig) + osumm[f"{key}_min_cosine"] = min(s[key]["cosine"] for s in per_step_orig) + entry[f"denoiser_{mode}"]["per_step_vs_original"] = per_step_orig + entry[f"denoiser_{mode}"]["summary_vs_original"] = osumm + del got, per_step, per_step_orig + torch.cuda.empty_cache() + s = entry[f"denoiser_{mode}"]["summary_vs_r768_reference"] + log(f" [{mode}] vs r768-ref: video rel_max={s['video_rel_max']:.4e} " + f"rms_rel={s['video_rms_rel']:.4e} cos={s['video_min_cosine']:.6f} | " + f"audio rel_max={s['audio_rel_max']:.4e} rms_rel={s['audio_rms_rel']:.4e} " + f"cos={s['audio_min_cosine']:.6f}") + + del fold + torch.cuda.empty_cache() + + compressed_adaln = (rank * basis_in + rank) + n_blocks * (blk_out * rank + blk_bias) + \ + (norm_out * rank + norm_bias) + compressed_total = (n_params - adaln_params) + compressed_adaln + entry["parameters"] = { + "label": label, + "rank": rank, + "baseline_total": n_params, + "baseline_total_B": n_params / 1e9, + "baseline_adaln": adaln_params, + "compressed_adaln": compressed_adaln, + "compressed_adaln_B": compressed_adaln / 1e9, + "compressed_total": compressed_total, + "compressed_total_B": compressed_total / 1e9, + "params_saved": n_params - compressed_total, + "params_saved_fraction": (n_params - compressed_total) / n_params, + "delta_vs_baseline": compressed_total - n_params, + "delta_explained_by_centering_bias": (compressed_total - n_params) == rank, + "projected_final_total_B": compressed_total / 1e9, + "storage_gb_bf16_baseline": n_params * 2 / 1e9, + "storage_gb_bf16_compressed": compressed_total * 2 / 1e9, + "is_basis_rotation": label.startswith("r768_rot"), + "storage_note": ("parameters x 2 bytes (bf16-equivalent) for both columns so they are " + "directly comparable. The shipped transformer is 40.27 GB on disk " + "because the rank-reduced AdaLN is stored fp16 and time_embedder / " + "proj_in / proj_out are kept fp32."), + "folded_basis_note": ("adaln_basis is folded to V.T @ W_b [r, 2688] + V.T (b_b - mu) [r]; " + "algebraically identical to keeping W_b plus V, and strictly smaller " + "for r < 768. adaln_basis has no bias in the original, so the " + "r-element centering bias is the only new parameter."), + } + log(f" params: {n_params/1e9:.4f}B -> {compressed_total/1e9:.4f}B " + f"(-{(n_params-compressed_total)/1e9:.4f}B, " + f"{100*(n_params-compressed_total)/n_params:.2f}%) | AdaLN {adaln_params/1e9:.4f}B -> " + f"{compressed_adaln/1e9:.4f}B") + + result["runs"][label] = entry + torch.cuda.empty_cache() + log(f" {label} done in {time.time()-t0:.1f}s") + + if label == "r768_canonical": + result["micro_perturbation_controls"] = controls + result["exact_mode"] = exact_mode + try: + gate = evaluate_gate(entry, controls, modes, exact_mode, result["shared_coordinate"]) + gate["raised"] = False + except Exception as exc: # noqa: BLE001 -- never lose a completed sweep + import traceback + gate = {"passed": False, "failed_checks": [f"gate raised {type(exc).__name__}"], + "checks": {}, "lines": [f"gate evaluation raised: {exc!r}"], + "traceback": traceback.format_exc(), "raised": True} + result["gate_r768"] = gate + write(out_path, result) + log("================ MANDATORY GATE, r = 768 ================") + for line in gate["lines"]: + log(" " + line) + log(f" VERDICT: {'PASS' if gate['passed'] else 'FAIL'}" + + ("" if gate["passed"] else f" failed={gate['failed_checks']}")) + if gate["raised"]: + log(" gate raised -- continuing anyway (see gate.traceback)") + elif not gate["passed"] or args.gate_only: + log(f"wrote {out_path}") + return 0 if gate["passed"] else 2 + log(f" r=768 conversion verified -- continuing to rotations and lower ranks") + + write(out_path, result) + result["micro_perturbation_controls"] = controls + result["exact_mode"] = exact_mode + write(out_path, result) + log(f"wrote {out_path} ({out_path.stat().st_size} bytes)") + return 0 + + +if __name__ == "__main__": + # os._exit skips interpreter finalization: a single-process nccl process + # group can otherwise block at exit, which in a batch job means sitting on + # the allocation until the wall-clock limit. + _rc = main() + sys.stdout.flush() + sys.stderr.flush() + os._exit(_rc) diff --git a/scripts/compacth3/analysis/checkpoint_sweep_metrics.py b/scripts/compacth3/analysis/checkpoint_sweep_metrics.py new file mode 100644 index 0000000000..d7567a1690 --- /dev/null +++ b/scripts/compacth3/analysis/checkpoint_sweep_metrics.py @@ -0,0 +1,296 @@ +#!/usr/bin/env python3 +"""Technical A/V retention grading for paired FastH3 checkpoint renders. + +The script intentionally reports a *retention index*, not a learned perceptual +quality score. Every candidate is compared prompt-by-prompt with a reference +render made with the same prompt and seed. This makes exposure, detail, +motion, temporal stability, and audio regressions visible without claiming to +measure anatomy, semantic correctness, or human preference. +""" + +from __future__ import annotations + +import argparse +import csv +import json +import math +import re +import subprocess +from pathlib import Path + +import numpy as np + + +VIDEO_WEIGHTS = { + "luma": 0.16, + "contrast": 0.18, + "tonal_range": 0.16, + "sharpness": 0.32, + "edge_density": 0.18, +} +TEMPORAL_WEIGHTS = { + "motion": 0.34, + "jerk_ratio": 0.42, + "flicker": 0.18, + "freeze_fraction": 0.06, +} +AUDIO_WEIGHTS = { + "lufs": 0.30, + "true_peak": 0.12, + "silence_fraction": 0.18, + "spectral_centroid": 0.12, + "spectral_flatness": 0.16, + "voice_band_ratio": 0.12, +} + + +def run_bytes(command: list[str]) -> bytes: + return subprocess.run(command, check=True, stdout=subprocess.PIPE, + stderr=subprocess.PIPE).stdout + + +def video_metrics(path: Path) -> dict[str, float]: + width, height = 208, 120 + raw = run_bytes([ + "ffmpeg", "-v", "error", "-i", str(path), "-an", "-vf", + f"fps=6,scale={width}:{height}:flags=area,format=gray", "-f", + "rawvideo", "-pix_fmt", "gray", "-", + ]) + frame_size = width * height + usable = len(raw) // frame_size * frame_size + frames = np.frombuffer(raw[:usable], dtype=np.uint8).reshape(-1, height, + width).astype(np.float32) + if len(frames) < 3: + raise RuntimeError(f"Too few decoded frames in {path}") + + frame_means = frames.mean(axis=(1, 2)) + frame_stds = frames.std(axis=(1, 2)) + p05 = np.percentile(frames, 5, axis=(1, 2)) + p95 = np.percentile(frames, 95, axis=(1, 2)) + + lap = (-4.0 * frames[:, 1:-1, 1:-1] + frames[:, :-2, 1:-1] + + frames[:, 2:, 1:-1] + frames[:, 1:-1, :-2] + + frames[:, 1:-1, 2:]) + gx = np.abs(frames[:, :, 1:] - frames[:, :, :-1]) + gy = np.abs(frames[:, 1:, :] - frames[:, :-1, :]) + sharpness = np.var(lap, axis=(1, 2)) + edge_density = 0.5 * ((gx > 18).mean(axis=(1, 2)) + + (gy > 18).mean(axis=(1, 2))) + + delta = np.abs(np.diff(frames, axis=0)).mean(axis=(1, 2)) + accel = np.abs(frames[2:] - 2.0 * frames[1:-1] + frames[:-2]).mean(axis=(1, 2)) + motion = float(np.median(delta)) + jerk = float(np.median(accel)) + return { + "luma": float(np.mean(frame_means)), + "contrast": float(np.mean(frame_stds)), + "tonal_range": float(np.mean(p95 - p05)), + "sharpness": float(np.median(sharpness)), + "edge_density": float(np.mean(edge_density)), + "motion": motion, + "jerk_ratio": jerk / max(motion, 1e-6), + "flicker": float(np.std(np.diff(frame_means))), + "freeze_fraction": float(np.mean(delta < 0.55)), + "black_fraction": float(np.mean(frame_means < 5.0)), + } + + +def ebur128(path: Path) -> tuple[float, float]: + proc = subprocess.run([ + "ffmpeg", "-hide_banner", "-nostats", "-i", str(path), + "-filter_complex", "ebur128=peak=true", "-f", "null", "-", + ], stdout=subprocess.DEVNULL, stderr=subprocess.PIPE, text=True, + check=True) + summaries = proc.stderr.split("Summary:") + text = summaries[-1] if len(summaries) > 1 else proc.stderr + integrated = re.search(r"I:\s*(-?[0-9.]+)\s+LUFS", text) + peak = re.search(r"Peak:\s*(-?[0-9.]+)\s+dBFS", text) + return (float(integrated.group(1)) if integrated else math.nan, + float(peak.group(1)) if peak else math.nan) + + +def audio_metrics(path: Path) -> dict[str, float]: + sample_rate = 16_000 + raw = run_bytes([ + "ffmpeg", "-v", "error", "-i", str(path), "-vn", "-ac", "1", + "-ar", str(sample_rate), "-f", "f32le", "-", + ]) + audio = np.frombuffer(raw, dtype="= 80) & (frequencies <= 4000) + audible_mask = (frequencies >= 40) & (frequencies <= 7800) + voice_ratio = np.sum(power[:, voice_mask], axis=1) / np.maximum( + np.sum(power[:, audible_mask], axis=1), 1e-15) + return { + "lufs": lufs, + "true_peak": true_peak, + "rms_dbfs": float(20.0 * np.log10(np.sqrt(np.mean(audio * audio)) + 1e-15)), + "silence_fraction": silence_fraction, + "spectral_centroid": float(np.median(centroid)), + "spectral_flatness": float(np.median(flatness)), + "voice_band_ratio": float(np.median(voice_ratio)), + "clipped_fraction": float(np.mean(np.abs(audio) >= 0.999)), + } + + +def log_distance(value: float, reference: float, scale: float) -> float: + return abs(math.log(max(value, 1e-9) / max(reference, 1e-9))) / scale + + +def linear_distance(value: float, reference: float, scale: float) -> float: + return abs(value - reference) / scale + + +def component_scores(candidate: dict[str, float], reference: dict[str, float]) -> dict[str, float]: + visual_scales = { + "luma": 25.0, + "contrast": 15.0, + "tonal_range": 30.0, + "sharpness": 0.70, + "edge_density": 0.35, + } + temporal_scales = { + "motion": 0.70, + "jerk_ratio": 0.55, + "flicker": 4.0, + "freeze_fraction": 0.20, + } + audio_scales = { + "lufs": 6.0, + "true_peak": 6.0, + "silence_fraction": 0.25, + "spectral_centroid": 0.70, + "spectral_flatness": 0.20, + "voice_band_ratio": 0.30, + } + + def weighted_score(weights: dict[str, float], scales: dict[str, float], + log_fields: set[str]) -> float: + distance = 0.0 + for key, weight in weights.items(): + if key in log_fields: + part = log_distance(candidate[key], reference[key], scales[key]) + else: + part = linear_distance(candidate[key], reference[key], scales[key]) + distance += weight * min(part, 3.0) + return 100.0 * math.exp(-distance) + + visual = weighted_score(VIDEO_WEIGHTS, visual_scales, + {"sharpness", "edge_density"}) + temporal = weighted_score(TEMPORAL_WEIGHTS, temporal_scales, + {"motion", "jerk_ratio"}) + audio = weighted_score(AUDIO_WEIGHTS, audio_scales, + {"spectral_centroid", "spectral_flatness", + "voice_band_ratio"}) + + # Hard technical failures must not be hidden by otherwise similar averages. + if candidate["black_fraction"] > 0.05: + visual *= max(0.0, 1.0 - candidate["black_fraction"]) + if candidate["clipped_fraction"] > 1e-4: + audio *= max(0.70, 1.0 - 10.0 * candidate["clipped_fraction"]) + overall = 0.40 * visual + 0.30 * temporal + 0.30 * audio + return {"visual": visual, "temporal": temporal, "audio": audio, + "overall": overall} + + +def confidence_interval(values: list[float], seed: int = 20260914) -> tuple[float, float]: + rng = np.random.default_rng(seed) + array = np.asarray(values) + draws = rng.choice(array, size=(20_000, len(array)), replace=True).mean(axis=1) + return float(np.percentile(draws, 2.5)), float(np.percentile(draws, 97.5)) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--root", type=Path, required=True) + parser.add_argument("--reference", type=Path, required=True) + parser.add_argument("--pattern", default="34block-step-*") + parser.add_argument("--output-prefix", type=Path, required=True) + args = parser.parse_args() + + candidates = sorted(args.root.glob(args.pattern), + key=lambda p: int(re.search(r"step-(\d+)", p.name).group(1))) + reference_files = {p.name: p for p in args.reference.glob("*.mp4")} + if not reference_files: + raise SystemExit(f"No MP4s in reference {args.reference}") + + cache: dict[str, dict[str, float]] = {} + + def measure(path: Path) -> dict[str, float]: + key = str(path) + if key not in cache: + cache[key] = {**video_metrics(path), **audio_metrics(path)} + return cache[key] + + reference = {name: measure(path) for name, path in reference_files.items()} + per_clip: list[dict[str, object]] = [] + summary: list[dict[str, object]] = [] + for folder in candidates: + step = int(re.search(r"step-(\d+)", folder.name).group(1)) + candidate_files = {p.name: p for p in folder.glob("*.mp4")} + common = sorted(reference_files.keys() & candidate_files.keys()) + if len(common) != len(reference_files): + continue + scores: list[dict[str, float]] = [] + for name in common: + metrics = measure(candidate_files[name]) + score = component_scores(metrics, reference[name]) + scores.append(score) + per_clip.append({"step": step, "prompt": name, **score, **metrics}) + overall_values = [s["overall"] for s in scores] + low, high = confidence_interval(overall_values, seed=20260914 + step) + summary.append({ + "step": step, + "clips": len(scores), + "overall": float(np.mean(overall_values)), + "ci95_low": low, + "ci95_high": high, + "visual": float(np.mean([s["visual"] for s in scores])), + "temporal": float(np.mean([s["temporal"] for s in scores])), + "audio": float(np.mean([s["audio"] for s in scores])), + "worst_prompt": min(per_clip[-len(scores):], key=lambda x: x["overall"])["prompt"], + "worst_score": min(overall_values), + }) + + summary.sort(key=lambda row: row["overall"], reverse=True) + args.output_prefix.parent.mkdir(parents=True, exist_ok=True) + with args.output_prefix.with_suffix(".summary.csv").open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=list(summary[0])) + writer.writeheader() + writer.writerows(summary) + with args.output_prefix.with_suffix(".clips.csv").open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=list(per_clip[0])) + writer.writeheader() + writer.writerows(per_clip) + payload = { + "reference": str(args.reference), + "interpretation": "Technical A/V retention relative to paired reference; not a semantic or human-preference score.", + "overall_weights": {"visual": 0.40, "temporal": 0.30, "audio": 0.30}, + "summary": summary, + } + args.output_prefix.with_suffix(".json").write_text(json.dumps(payload, indent=2) + "\n") + print(json.dumps(summary, indent=2)) + + +if __name__ == "__main__": + main() diff --git a/scripts/compacth3/eval_corrected_dmd_inference_exports.sbatch b/scripts/compacth3/eval_corrected_dmd_inference_exports.sbatch new file mode 100644 index 0000000000..5c16caf916 --- /dev/null +++ b/scripts/compacth3/eval_corrected_dmd_inference_exports.sbatch @@ -0,0 +1,133 @@ +#!/bin/bash +# Validate and render the corrected DMD2 run's native inference exports. +# +# Do not invoke dcp_to_diffusers here. The protected training checkpoint has +# intentional mixed training dtypes, while the training callback has already +# emitted a complete inference package in bfloat16 for each validation step. +# Evaluating that package directly is both lossless and the release path. +#SBATCH --job-name=h3-dmd2-export-eval +#SBATCH --partition=all +#SBATCH --nodes=1 +#SBATCH --ntasks=1 +#SBATCH --gres=gpu:4 +#SBATCH --time=04:00:00 +#SBATCH --output=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-export-eval-%j.log +#SBATCH --error=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-export-eval-%j.log + +set -euo pipefail +export SLURM_EXPORT_ENV=ALL + +SPRINT_ROOT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +DMD_CODE=${SPRINT_ROOT}/code/release20b-dmd2-v17-fp32resume-v1 +RUNNER_ROOT=${SPRINT_ROOT}/code/release-execution-dmd-ladderfix-v1 +DMD_RUN=${SPRINT_ROOT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +STEPS="800 1000 1200" +export SPRINT_ROOT DMD_CODE RUNNER_ROOT DMD_RUN STEPS + +srun --export=ALL --kill-on-bad-exit=1 \ + --container-image=nvcr.io/nvidia/pytorch:25.06-py3 \ + --container-mounts=/mnt/nfs/vlm-aryan:/mnt/nfs/vlm-aryan,/mnt/lustre/vlm-shared:/mnt/lustre/vlm-shared:ro \ + --container-workdir="${RUNNER_ROOT}" bash -lc ' +set -euo pipefail +source /mnt/nfs/vlm-aryan/fasth3-33b-20260806/secrets.env +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +PROMPTS="${SPRINT_ROOT}/job-scripts/bench_five_new_prompts.json" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache PYTHONDONTWRITEBYTECODE=1 +export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA FASTVIDEO_MINIMAX_H3_FUSIONS=0 +export FASTVIDEO_DMD_DENOISING_STEPS=999,749,500,250 +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True TORCH_NCCL_ENABLE_MONITORING=0 +export PYTHONPATH="${RUNNER_ROOT}:${SPRINT_ROOT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" + +for step in ${STEPS}; do + model="${DMD_RUN}/inference/checkpoint-${step}" + eval_dir="${DMD_RUN}/eval-exported-checkpoint-${step}-bench-five-exact4-v1" + media="${eval_dir}/media" + + [[ -f "${model}/.complete" && -s "${model}/metadata.json" ]] || { + echo "FATAL: incomplete inference export for step ${step}" >&2 + exit 2 + } + [[ $(find "${model}/transformer" -maxdepth 1 -name "*.safetensors" ! -name "*.index.json" | wc -l) -eq 8 ]] || { + echo "FATAL: step ${step} does not have eight transformer shards" >&2 + exit 3 + } + + # Header-only audit: proves the package is internally dtype-consistent + # without materializing ~40 GB of tensors on the CPU. + "${PY}" - "${model}" <<"PY" +import collections +import json +import struct +import sys +from pathlib import Path + +root = Path(sys.argv[1]) +counts = collections.Counter() +for shard in sorted((root / "transformer").glob("*.safetensors")): + with shard.open("rb") as handle: + header_size = struct.unpack(" "${eval_dir}/completed_at.txt" + echo "step ${step}: recovered and accepted five already-rendered videos" + continue + fi + [[ ! -e "${eval_dir}" ]] || { + echo "FATAL: refusing to overwrite partial evaluation ${eval_dir}" >&2 + exit 4 + } + mkdir -p "${eval_dir}" + + cd "${RUNNER_ROOT}" + "${PY}" "${RUNNER_ROOT}/scripts/fasth3_sprint/run_baseline_matrix.py" \ + --model-path "${model}" --checkpoint-role "corrected-dmd2-step${step}-native-export-exact4" \ + --attention dense --attention-backend TORCH_SDPA --prompts "${PROMPTS}" \ + --output-dir "${media}" --run-id "corrected-dmd2-step${step}-native-export-exact4-${SLURM_JOB_ID}" \ + --source-commit "$(cat "${DMD_CODE}/CODE_COMMIT" 2>/dev/null || echo unknown)" --max-prompts 5 \ + --height 480 --width 832 --num-frames 124 --seed 20260912 \ + --steps 5 --num-gpus 4 --dit-precision bf16 --profile strict \ + --no-fa4 --no-compile --no-upload-videos + + "${PY}" - "${media}" <<"PY" +import json +import sys +from pathlib import Path + +media = Path(sys.argv[1]) +manifest = json.loads((media / "run_manifest.json").read_text()) +assert len(list(media.glob("*.mp4"))) == 5 +assert manifest["schedule"]["grid_points"] == 5, manifest["schedule"] +assert manifest["schedule"]["transformer_calls"] == 4, manifest["schedule"] +assert len(manifest["schedule"]["video"]["transformer_timesteps"]) == 4 +assert len(manifest["schedule"]["audio"]["transformer_timesteps"]) == 4 +print("verified exact four-call media", media) +PY + date -Is > "${eval_dir}/completed_at.txt" +done +' diff --git a/scripts/compacth3/eval_dmd_export_1400.sh b/scripts/compacth3/eval_dmd_export_1400.sh new file mode 100644 index 0000000000..f4a9658ff1 --- /dev/null +++ b/scripts/compacth3/eval_dmd_export_1400.sh @@ -0,0 +1,133 @@ +#!/bin/bash +# Validate and render the corrected DMD2 run's native inference exports. +# +# Do not invoke dcp_to_diffusers here. The protected training checkpoint has +# intentional mixed training dtypes, while the training callback has already +# emitted a complete inference package in bfloat16 for each validation step. +# Evaluating that package directly is both lossless and the release path. +#SBATCH --job-name=h3-dmd2-export-eval +#SBATCH --partition=all +#SBATCH --nodes=1 +#SBATCH --ntasks=1 +#SBATCH --gres=gpu:4 +#SBATCH --time=04:00:00 +#SBATCH --output=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-export-eval-%j.log +#SBATCH --error=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-export-eval-%j.log + +set -euo pipefail +export SLURM_EXPORT_ENV=ALL + +SPRINT_ROOT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +DMD_CODE=${SPRINT_ROOT}/code/release20b-dmd2-v17-fp32resume-v1 +RUNNER_ROOT=${SPRINT_ROOT}/code/release-execution-dmd-ladderfix-v1 +DMD_RUN=${SPRINT_ROOT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +STEPS="1400" +export SPRINT_ROOT DMD_CODE RUNNER_ROOT DMD_RUN STEPS + +srun --overlap --jobid=9639 -N1 -n1 --gres=gpu:4 --export=ALL --kill-on-bad-exit=1 \ + --container-image=nvcr.io/nvidia/pytorch:25.06-py3 \ + --container-mounts=/mnt/nfs/vlm-aryan:/mnt/nfs/vlm-aryan,/mnt/lustre/vlm-shared:/mnt/lustre/vlm-shared:ro \ + --container-workdir="${RUNNER_ROOT}" bash -lc ' +set -euo pipefail +source /mnt/nfs/vlm-aryan/fasth3-33b-20260806/secrets.env +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +PROMPTS="${SPRINT_ROOT}/job-scripts/bench_five_new_prompts.json" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache PYTHONDONTWRITEBYTECODE=1 +export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA FASTVIDEO_MINIMAX_H3_FUSIONS=0 +export FASTVIDEO_DMD_DENOISING_STEPS=999,749,500,250 +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True TORCH_NCCL_ENABLE_MONITORING=0 +export PYTHONPATH="${RUNNER_ROOT}:${SPRINT_ROOT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" + +for step in ${STEPS}; do + model="${DMD_RUN}/inference/checkpoint-${step}" + eval_dir="${DMD_RUN}/eval-exported-checkpoint-${step}-bench-five-exact4-v1" + media="${eval_dir}/media" + + [[ -f "${model}/.complete" && -s "${model}/metadata.json" ]] || { + echo "FATAL: incomplete inference export for step ${step}" >&2 + exit 2 + } + [[ $(find "${model}/transformer" -maxdepth 1 -name "*.safetensors" ! -name "*.index.json" | wc -l) -eq 8 ]] || { + echo "FATAL: step ${step} does not have eight transformer shards" >&2 + exit 3 + } + + # Header-only audit: proves the package is internally dtype-consistent + # without materializing ~40 GB of tensors on the CPU. + "${PY}" - "${model}" <<"PY" +import collections +import json +import struct +import sys +from pathlib import Path + +root = Path(sys.argv[1]) +counts = collections.Counter() +for shard in sorted((root / "transformer").glob("*.safetensors")): + with shard.open("rb") as handle: + header_size = struct.unpack(" "${eval_dir}/completed_at.txt" + echo "step ${step}: recovered and accepted five already-rendered videos" + continue + fi + [[ ! -e "${eval_dir}" ]] || { + echo "FATAL: refusing to overwrite partial evaluation ${eval_dir}" >&2 + exit 4 + } + mkdir -p "${eval_dir}" + + cd "${RUNNER_ROOT}" + "${PY}" "${RUNNER_ROOT}/scripts/fasth3_sprint/run_baseline_matrix.py" \ + --model-path "${model}" --checkpoint-role "corrected-dmd2-step${step}-native-export-exact4" \ + --attention dense --attention-backend TORCH_SDPA --prompts "${PROMPTS}" \ + --output-dir "${media}" --run-id "corrected-dmd2-step${step}-native-export-exact4-${SLURM_JOB_ID}" \ + --source-commit "$(cat "${DMD_CODE}/CODE_COMMIT" 2>/dev/null || echo unknown)" --max-prompts 5 \ + --height 480 --width 832 --num-frames 124 --seed 20260912 \ + --steps 5 --num-gpus 4 --dit-precision bf16 --profile strict \ + --no-fa4 --no-compile --no-upload-videos + + "${PY}" - "${media}" <<"PY" +import json +import sys +from pathlib import Path + +media = Path(sys.argv[1]) +manifest = json.loads((media / "run_manifest.json").read_text()) +assert len(list(media.glob("*.mp4"))) == 5 +assert manifest["schedule"]["grid_points"] == 5, manifest["schedule"] +assert manifest["schedule"]["transformer_calls"] == 4, manifest["schedule"] +assert len(manifest["schedule"]["video"]["transformer_timesteps"]) == 4 +assert len(manifest["schedule"]["audio"]["transformer_timesteps"]) == 4 +print("verified exact four-call media", media) +PY + date -Is > "${eval_dir}/completed_at.txt" +done +' diff --git a/scripts/compacth3/grade_dmd2_all_checkpoints.sbatch b/scripts/compacth3/grade_dmd2_all_checkpoints.sbatch new file mode 100644 index 0000000000..6e1c4575e5 --- /dev/null +++ b/scripts/compacth3/grade_dmd2_all_checkpoints.sbatch @@ -0,0 +1,49 @@ +#!/bin/bash +# Grade every comparable corrected-DMD2 checkpoint plus the protected old step 2750. +#SBATCH --job-name=h3-dmd2-grade-all +#SBATCH --partition=all +#SBATCH --nodes=1 +#SBATCH --ntasks=1 +#SBATCH --gres=gpu:1 +#SBATCH --cpus-per-task=16 +#SBATCH --mem=128G +#SBATCH --time=06:00:00 +#SBATCH --output=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-grade-all-%j.log +#SBATCH --error=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/sbatch-h3-dmd2-grade-all-%j.log + +set -euo pipefail +export SLURM_EXPORT_ENV=ALL + +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +EVAL=/mnt/nfs/vlm-aryan/fasth3-eval +RUN=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +OUT="${RUN}/dmd-native-export-grade-all-v1" + +srun --export=ALL --kill-on-bad-exit=1 \ + --container-image=nvcr.io/nvidia/pytorch:25.06-py3 \ + --container-mounts=/mnt/nfs/vlm-aryan:/mnt/nfs/vlm-aryan,/mnt/lustre/vlm-shared:/mnt/lustre/vlm-shared:ro \ + bash -lc ' +set -euo pipefail +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +EVAL=/mnt/nfs/vlm-aryan/fasth3-eval +RUN=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +OUT="${RUN}/dmd-native-export-grade-all-v1" +mkdir -p "${OUT}" + +for step in 100 200 400 600 800 1000 1200; do + media="${RUN}/eval-exported-checkpoint-${step}-bench-five-exact4-v1/media" + [[ $(find "${media}" -maxdepth 1 -name "*.mp4" | wc -l) -eq 5 ]] || { + echo "FATAL: step ${step} does not have five videos" >&2 + exit 2 + } +done + +bash "${EVAL}/probe_deps.sh" | tee "${OUT}/dependency-probe.txt" +"${EVAL}/venv/bin/python" "${EVAL}/score_34block.py" \ + --runs "${EVAL}/dmd2_all_checkpoints_runs.json" \ + --reference-media "${RUN}/eval-parent-750-bench-five-correct4-v1/media" \ + --prompts "${SPRINT}/job-scripts/bench_five_new_prompts.json" \ + --out "${OUT}/scores.json" | tee "${OUT}/scorer.log" +"${EVAL}/venv/bin/python" "${EVAL}/summarize_scores.py" "${OUT}/scores.json" | tee "${OUT}/summary.txt" +date -Is > "${OUT}/completed_at.txt" +' diff --git a/scripts/compacth3/qad/gen_qad_configs.py b/scripts/compacth3/qad/gen_qad_configs.py new file mode 100644 index 0000000000..765a330565 --- /dev/null +++ b/scripts/compacth3/qad/gen_qad_configs.py @@ -0,0 +1,124 @@ +#!/usr/bin/env python3 +"""Emit launch-ready NVFP4 QAD configs for BOTH the rank-768 (20B) and rank-16 (~17B) students. + +QAD = the release path: FP4 forward with a full-precision backward (STE), so FSDP +sharding and checkpointing stay dense-identical. No weight conversion needed. +""" +import json, pathlib, sys, yaml + +S = pathlib.Path("/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829") +M = pathlib.Path("/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1") +RUN = S / "runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3" +TEACHER = S / "release-candidates/base-h3-teacher-complete-v1" + +meta = json.loads((RUN / "checkpoint-1400/metadata.json").read_text()) +c = meta["config"] +def y(p, d=None): + cur = c + for k in p.split("."): + if not isinstance(cur, dict) or k not in cur: return d + cur = cur[k] + return cur + +VARIANTS = { + "r768": { + "adaln_rank": 768, + "init_from": str(RUN / "inference/checkpoint-1400"), + "note": "20B baseline lineage; identical architecture to the shipped rank-768 model.", + }, + "r16": { + "adaln_rank": 16, + "init_from": str(RUN / "inference/checkpoint-1400-reparam-r16"), + "note": ("~17B. Student is the post-hoc centered-affine rank-16 reparameterization of " + "checkpoint-1400, materialized as a rank-16 checkpoint. Requires that " + "materialization to exist -- see the rank-compression run."), + }, +} + +HEADER = """# NVFP4 QAD -- {tag} (adaln_rank={rank}) +# +# THE RELEASE RUN. Post-hoc rank compression + NVFP4 quantization repair in one pass. +# +# student : {init} (adaln_rank={rank}) +# teacher : frozen base H3 (base-h3-teacher-complete-v1) +# ladder : 4 calls -- dmd_denoising_steps [999, 749, 500, 250] +# quant : nvfp4_qat_train -- FP4 forward, full-precision backward (STE). +# No weight conversion, so FSDP sharding/checkpointing stay dense-identical. +# decode : NOT taeh3 (rejected upstream; decoder sits downstream of the DiT anyway). +# +# AUDIO PROTECTION -- read before changing anything: +# * modality_loss_weights carried over from the parent run UNCHANGED, so QAD does not +# silently rebalance video against audio. Upweight 'audio' if the audio A/B regresses. +# * audio_proj_in / audio_proj_out match no DEFAULT_FP4_LAYERS entry, so they stay bf16. +# Do not add them. +# * Upstream's NVFP4 text-encoder PR found the fully-quantized variant LOST THE VOICE +# TRACK. Audio is the first thing low precision breaks -- gate every QAD checkpoint on +# speech intelligibility, not on a combined scalar. +# +# RANK-SPECIFIC NOTE: {note} +""" + +def build(tag, v): + q = { + "models": { + "student": { + "_target_": y("models.student._target_"), + "init_from": v["init_from"], + "trainable": True, + "enable_gradient_checkpointing_type": "full", + "attention_backend": y("models.student.attention_backend", "TORCH_SDPA"), + "quant_config": "nvfp4_qat_train", + "adaln_rank": v["adaln_rank"], + }, + "teacher": { + "_target_": y("models.teacher._target_"), + "init_from": y("models.teacher.init_from", str(TEACHER)), + "trainable": False, + "disable_custom_init_weights": True, + "attention_backend": y("models.teacher.attention_backend", "TORCH_SDPA"), + }, + }, + "method": {k: y(f"method.{k}") for k in ( + "_target_", "rollout_mode", "rollout_carry", "rollout_carry_slots", + "rollout_sample_type", "generator_update_interval", "real_score_guidance_scale", + "dmd_denoising_steps", "min_timestep_ratio", "max_timestep_ratio", + "score_timestep_shift", "score_timestep_warp_max", "score_timestep_continuous", + "fake_score_loss_space", "modality_loss_weights", "dmd_denom_floor_ratio", + "dmd_grad_cap", "cfg_uncond", "fake_score_learning_rate", "fake_score_betas", + "fake_score_lr_scheduler")}, + "training": { + "distributed": y("training.distributed"), + "data": y("training.data"), + "optimizer": y("training.optimizer"), + "loop": {"max_train_steps": 200, + "gradient_accumulation_steps": y("training.loop.gradient_accumulation_steps", 8)}, + "checkpoint": { + "output_dir": str(S / f"runs/release20b-qad-nvfp4-4call-{tag}-v1"), + "resume_from_checkpoint": v["init_from"], + "training_state_checkpointing_steps": 25, + "require_complete_training_checkpoint": True, + "checkpoints_total_limit": 12, + }, + "tracker": {"trackers": ["wandb"], "project_name": "fasth3-14b-2step-qad-sprint", + "run_name": f"release20b-qad-nvfp4-4call-{tag}-v1"}, + }, + "callbacks": { + "grad_clip": {"_target_": "fastvideo.train.callbacks.grad_clip.GradNormClipCallback", + "max_grad_norm": 1.0}, + "validation": y("callbacks.validation"), + }, + "model": {"precondition_outputs": False, "enable_gradient_checkpointing_type": "full", + "enable_torch_compile": False}, + "dit_precision": "fp32", + "vsa": y("vsa"), + } + out = M / f"examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call_{tag}.yaml" + out.write_text(HEADER.format(tag=tag, rank=v["adaln_rank"], init=v["init_from"], note=v["note"]) + + yaml.safe_dump(q, sort_keys=False)) + return out, q + +for tag, v in VARIANTS.items(): + out, q = build(tag, v) + print(f"{tag}: {out.name}") + print(f" adaln_rank={q['models']['student']['adaln_rank']} init={q['models']['student']['init_from'][-40:]}") + print(f" quant={q['models']['student']['quant_config']} steps={q['training']['loop']['max_train_steps']}") diff --git a/scripts/compacth3/qad/qad_checkpoint_gate.py b/scripts/compacth3/qad/qad_checkpoint_gate.py new file mode 100644 index 0000000000..da50224fb7 --- /dev/null +++ b/scripts/compacth3/qad/qad_checkpoint_gate.py @@ -0,0 +1,162 @@ +#!/usr/bin/env python3 +"""Select the QAD checkpoint: earliest one that cuts grain without raising anomalies. + +Judged as PAIRED DELTAS against the PRE-QAD NVFP4 model (not absolute scores), on the +same prompts+seeds, on the hard-motion set. Criteria from the lead: + + d_grain < 0 (grain falls) + d_detail >= 0 (detail does not regress) + d_anoms <= 0 (temporal anomalies do not increase) + |d_motion| small (motion magnitude preserved) + +Earliest checkpoint satisfying ALL FOUR wins. If NONE does, that is itself the result: +standard QAD is trading away the NVFP4 stability benefit, and only then is a custom +temporal objective worth designing. + +Runs over EVERY checkpoint (the run emits one per 25 steps), independent of the +built-in validation cadence (every 50) -- otherwise the sweet spot falls between samples. +""" +import argparse, glob, json, os, re, subprocess, sys + +S = "/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829" +MOTION_TOL = 0.15 # |d_motion| within 15% of baseline counts as preserved + + +def ckpt_step(path): + m = re.search(r"checkpoint-(\d+)", os.path.basename(path)) + return int(m.group(1)) if m else None + + +def metrics_for(video_dir): + """Compute the four decision metrics for one model's generations. + + NOTE: motion magnitude is the mandatory anti-collapse guard. Implementations here + are deliberately simple and must be swapped for the sweep agent's versions once + those land, so both sides of every delta use the SAME code path. + """ + out = {} + for name, fn in (("motion", _motion_magnitude), + ("grain", _hf_energy), + ("detail", _sharpness), + ("anomalies", _colour_drift)): + vals = [] + for v in sorted(glob.glob(os.path.join(video_dir, "*.mp4"))): + try: + vals.append(fn(v)) + except Exception as e: + print(f" WARN {name} failed on {os.path.basename(v)}: {type(e).__name__}", flush=True) + out[name] = sum(vals) / len(vals) if vals else None + return out + + +def _frames(v, n=32): + import av, numpy as np + c = av.open(v) + fr = [f.to_ndarray(format="rgb24") for f in c.decode(video=0)] + c.close() + if len(fr) > n: + idx = np.linspace(0, len(fr) - 1, n).astype(int) + fr = [fr[i] for i in idx] + return fr + + +def _motion_magnitude(v): + """Mean optical-flow magnitude. The guard against a model that 'wins' by going static.""" + import cv2, numpy as np + fr = _frames(v) + mags = [] + for a, b in zip(fr, fr[1:]): + ga = cv2.cvtColor(a, cv2.COLOR_RGB2GRAY) + gb = cv2.cvtColor(b, cv2.COLOR_RGB2GRAY) + fl = cv2.calcOpticalFlowFarneback(ga, gb, None, 0.5, 3, 15, 3, 5, 1.2, 0) + mags.append(float(np.sqrt(fl[..., 0] ** 2 + fl[..., 1] ** 2).mean())) + return sum(mags) / len(mags) if mags else 0.0 + + +def _hf_energy(v): + """High-frequency energy: proxy for grain. Higher = grainier.""" + import cv2, numpy as np + vals = [] + for f in _frames(v, 16): + g = cv2.cvtColor(f, cv2.COLOR_RGB2GRAY).astype("float32") + vals.append(float(cv2.Laplacian(g, cv2.CV_32F).var())) + return sum(vals) / len(vals) if vals else 0.0 + + +def _sharpness(v): + """Detail proxy. Deliberately the same family as grain; report both so a + grain reduction that is really just blurring is visible.""" + import cv2, numpy as np + vals = [] + for f in _frames(v, 16): + g = cv2.cvtColor(f, cv2.COLOR_RGB2GRAY).astype("float32") + vals.append(float(cv2.Sobel(g, cv2.CV_32F, 1, 0).var())) + return sum(vals) / len(vals) if vals else 0.0 + + +def _colour_drift(v): + """Flow-warped colour residual between adjacent frames: targets the reported + 'gloves change colour' failure. Motion-compensated, so legitimate motion + does not count as drift.""" + import cv2, numpy as np + fr = _frames(v) + res = [] + for a, b in zip(fr, fr[1:]): + ga = cv2.cvtColor(a, cv2.COLOR_RGB2GRAY) + gb = cv2.cvtColor(b, cv2.COLOR_RGB2GRAY) + fl = cv2.calcOpticalFlowFarneback(ga, gb, None, 0.5, 3, 15, 3, 5, 1.2, 0) + h, w = ga.shape + xx, yy = np.meshgrid(np.arange(w), np.arange(h)) + wx = (xx + fl[..., 0]).astype(np.float32) + wy = (yy + fl[..., 1]).astype(np.float32) + warped = cv2.remap(a, wx, wy, cv2.INTER_LINEAR, borderMode=cv2.BORDER_REPLICATE) + # Lab is closer to perceptual colour than RGB + la = cv2.cvtColor(warped, cv2.COLOR_RGB2LAB).astype("float32") + lb = cv2.cvtColor(b, cv2.COLOR_RGB2LAB).astype("float32") + res.append(float(np.abs(la - lb).mean())) + return sum(res) / len(res) if res else 0.0 + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--baseline-dir", required=True, help="pre-QAD NVFP4 generations, same prompts+seeds") + ap.add_argument("--ckpt-root", required=True, help="QAD output_dir containing checkpoint-*") + ap.add_argument("--gen-dir-template", required=True, + help="where per-checkpoint generations live; must contain {step}") + ap.add_argument("--out", default=f"{S}/adaln_rank_analysis/qad_checkpoint_gate.json") + a = ap.parse_args() + + base = metrics_for(a.baseline_dir) + print(f"BASELINE (pre-QAD NVFP4): {json.dumps(base, indent=2)}", flush=True) + + steps = sorted(s for s in (ckpt_step(p) for p in glob.glob(f"{a.ckpt_root}/checkpoint-*")) if s) + rows, winner = [], None + for s in steps: + d = a.gen_dir_template.format(step=s) + if not os.path.isdir(d): + print(f"checkpoint-{s}: no generations at {d} -- SKIPPED (must be generated)", flush=True) + continue + m = metrics_for(d) + dl = {k: (m[k] - base[k]) for k in base if m.get(k) is not None and base.get(k)} + rel = {k: (v / base[k]) for k, v in dl.items() if base.get(k)} + ok = (rel.get("grain", 1) < 0 and rel.get("detail", -1) >= 0 + and rel.get("anomalies", 1) <= 0 and abs(rel.get("motion", 1)) <= MOTION_TOL) + rows.append({"step": s, "abs": m, "delta": dl, "rel": rel, "satisfies": ok}) + print(f"checkpoint-{s}: d_grain={rel.get('grain',0):+.3f} d_detail={rel.get('detail',0):+.3f} " + f"d_anom={rel.get('anomalies',0):+.3f} d_motion={rel.get('motion',0):+.3f} " + f"{'<== SATISFIES' if ok else ''}", flush=True) + if ok and winner is None: + winner = s + + verdict = ("EARLIEST satisfying checkpoint: %d" % winner) if winner is not None else \ + ("NONE satisfies: standard QAD is trading away the NVFP4 stability benefit " + "-- a custom temporal objective is now justified") + print("\nVERDICT:", verdict, flush=True) + os.makedirs(os.path.dirname(a.out), exist_ok=True) + json.dump({"baseline": base, "rows": rows, "winner": winner, "verdict": verdict}, + open(a.out, "w"), indent=2) + print("wrote", a.out, flush=True) + + +if __name__ == "__main__": + main() diff --git a/scripts/compacth3/qad/setup_qad.py b/scripts/compacth3/qad/setup_qad.py new file mode 100644 index 0000000000..b73baeaa1a --- /dev/null +++ b/scripts/compacth3/qad/setup_qad.py @@ -0,0 +1,147 @@ +#!/usr/bin/env python3 +"""Set up the NVFP4 QAD run: 4-call, no taeh3, teacher = base H3, resuming from 1400. + +Two parts: + 1. Fix the FP4 target-layer list so H3's FFN is actually quantized. The generic + dispatch matches "ffn.fc_in"/"ffn.fc_out", but H3's modules are "ff.fc_in"/ + "ff.fc_out" -- so as shipped, QAD would quantize attention and silently skip + the entire FFN. + 2. Emit a QAD yaml derived from the DMD2 run's own config (checkpoint-1400 + metadata) with quant_config: nvfp4_qat_train on the student. +""" +import json +import pathlib +import sys + +M = pathlib.Path("/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1") +SPRINT = pathlib.Path("/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829") +RUN = SPRINT / "runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3" + +# --- 1. FFN layer-name fix -------------------------------------------------- +qc = M / "fastvideo/layers/quantization/nvfp4_qat_config.py" +t = qc.read_text(); orig = t +if '"ff.fc_in"' in t: + print("layer list: already fixed") +else: + anchor = 'DEFAULT_FP4_LAYERS = (\n' + assert t.count(anchor) == 1, f"anchor={t.count(anchor)}" + t = t.replace(anchor, anchor + ' # MiniMax-H3 names its FFN "ff.", not "ffn." -- without these the\n' + ' # FFN is silently left dense and QAD only covers attention.\n' + ' "ff.fc_in",\n "ff.fc_out",\n', 1) + assert t != orig + qc.with_suffix(".py.pre-qad-bak").write_text(orig) + qc.write_text(t) + print("layer list: added ff.fc_in / ff.fc_out for H3") + +# --- 2. build the QAD yaml from the DMD2 run's own metadata ------------------ +meta = json.loads((RUN / "checkpoint-1400/metadata.json").read_text()) +c = meta["config"] + +def y(path, default=None): + cur = c + for k in path.split("."): + if not isinstance(cur, dict) or k not in cur: + return default + cur = cur[k] + return cur + +qad = { + "models": { + "student": { + "_target_": y("models.student._target_"), + "init_from": str(RUN / "checkpoint-1400"), + "trainable": True, + "enable_gradient_checkpointing_type": "full", + "attention_backend": y("models.student.attention_backend", "TORCH_SDPA"), + # THE QAD KNOB: FP4 forward + full-precision backward (STE). + # No weight conversion, so FSDP sharding/checkpointing stay dense-identical. + "quant_config": "nvfp4_qat_train", + }, + "teacher": { + "_target_": y("models.teacher._target_"), + "init_from": y("models.teacher.init_from"), + "trainable": False, + "disable_custom_init_weights": True, + "attention_backend": y("models.teacher.attention_backend", "TORCH_SDPA"), + }, + }, + "method": { + "_target_": y("method._target_"), + "rollout_mode": y("method.rollout_mode"), + "rollout_carry": y("method.rollout_carry"), + "rollout_carry_slots": y("method.rollout_carry_slots"), + "rollout_sample_type": y("method.rollout_sample_type"), + "generator_update_interval": y("method.generator_update_interval"), + "real_score_guidance_scale": y("method.real_score_guidance_scale"), + # 4-call ladder, unchanged from the parent run + "dmd_denoising_steps": y("method.dmd_denoising_steps"), + "min_timestep_ratio": y("method.min_timestep_ratio"), + "max_timestep_ratio": y("method.max_timestep_ratio"), + "score_timestep_shift": y("method.score_timestep_shift"), + "score_timestep_warp_max": y("method.score_timestep_warp_max"), + "score_timestep_continuous": y("method.score_timestep_continuous"), + "fake_score_loss_space": y("method.fake_score_loss_space"), + "modality_loss_weights": y("method.modality_loss_weights"), + "dmd_denom_floor_ratio": y("method.dmd_denom_floor_ratio"), + "dmd_grad_cap": y("method.dmd_grad_cap"), + "cfg_uncond": y("method.cfg_uncond"), + "fake_score_learning_rate": y("method.fake_score_learning_rate"), + "fake_score_betas": y("method.fake_score_betas"), + "fake_score_lr_scheduler": y("method.fake_score_lr_scheduler"), + }, + "training": { + "distributed": y("training.distributed"), + "data": y("training.data"), + "optimizer": y("training.optimizer"), + "loop": {"max_train_steps": 200, "gradient_accumulation_steps": y("training.loop.gradient_accumulation_steps", 8)}, + "checkpoint": { + "output_dir": str(SPRINT / "runs/release20b-dmd2-v12-qad-nvfp4-4call-v1"), + "resume_from_checkpoint": str(RUN / "checkpoint-1400"), + "training_state_checkpointing_steps": 25, + "require_complete_training_checkpoint": True, + "checkpointing_start_step": 1400, + "checkpoints_total_limit": 12, + }, + "tracker": {"trackers": ["wandb"], "project_name": "fasth3-14b-2step-qad-sprint", + "run_name": "release20b-dmd2-qad-nvfp4-4call-v1"}, + }, + "callbacks": { + "grad_clip": {"_target_": "fastvideo.train.callbacks.grad_clip.GradNormClipCallback", "max_grad_norm": 1.0}, + "validation": y("callbacks.validation"), + }, + "model": {"precondition_outputs": False, "enable_gradient_checkpointing_type": "full", "enable_torch_compile": False}, + "dit_precision": "fp32", + "vsa": y("vsa"), +} + +out = M / "examples/train/configs/distribution_matching/minimax_h3/qad_nvfp4_4call.yaml" +out.parent.mkdir(parents=True, exist_ok=True) + +import yaml +header = """# NVFP4 QAD for the 42-block 4-call DMD2 student. +# +# student : checkpoint-1400 of the CORRECTED re-run (job-paired8972-8975-4000-v3) +# teacher : frozen base H3 (base-h3-teacher-complete-v1) +# ladder : 4 calls -- dmd_denoising_steps [999, 749, 500, 250] +# decode : NOT taeh3. The decoder is downstream of the DiT, so keeping it out +# lets one QAD serve both the taeh3 preview and the full-VAE release. +# quant : nvfp4_qat_train -- FP4 forward, full-precision backward (STE). +# No weight conversion, so FSDP sharding/checkpointing stay dense-identical. +# +# AUDIO PROTECTION -- read before changing anything: +# * modality_loss_weights is carried over from the parent run unchanged, so the +# QAD does not silently rebalance video against audio. Upweight 'audio' here +# if the audio A/B regresses. +# * audio_proj_in / audio_proj_out are NOT in the FP4 target list (they match no +# DEFAULT_FP4_LAYERS entry), so they stay bf16. Do not add them. +# * The PR that added the NVFP4 encoder found the fully-quantized variant LOST THE +# VOICE TRACK. Audio is the first thing low precision breaks -- gate every QAD +# checkpoint on speech intelligibility, not on a combined scalar. +# * The validation panel below includes speech and music prompts on purpose. +""" +out.write_text(header + yaml.safe_dump(qad, sort_keys=False)) +print("wrote", out) +print("student init_from:", qad["models"]["student"]["init_from"]) +print("quant_config :", qad["models"]["student"]["quant_config"]) +print("denoise steps :", qad["method"]["dmd_denoising_steps"]) +print("output_dir :", qad["training"]["checkpoint"]["output_dir"]) diff --git a/scripts/compacth3/qad/tune_qad_cadence.py b/scripts/compacth3/qad/tune_qad_cadence.py new file mode 100644 index 0000000000..a748a9a1d5 --- /dev/null +++ b/scripts/compacth3/qad/tune_qad_cadence.py @@ -0,0 +1,30 @@ +#!/usr/bin/env python3 +"""Make the QAD configs emit frequent short checkpoints so we can find the +EARLIEST checkpoint that reduces grain without raising temporal anomalies. + +Decision rule (lead): select the earliest checkpoint that meaningfully reduces +grain without increasing temporal anomalies. Do NOT optimise until QAD +reproduces BF16 exactly -- BF16 has more temporal failures on some cases. +""" +import pathlib, yaml +M = pathlib.Path("/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1") +base = M / "examples/train/configs/distribution_matching/minimax_h3" + +for tag in ("r768", "r16"): + p = base / f"qad_nvfp4_4call_{tag}.yaml" + if not p.exists(): + print(f"MISSING {p.name}"); continue + d = yaml.safe_load(p.read_text()) + ck = d["training"]["checkpoint"] + ck["training_state_checkpointing_steps"] = 25 # frequent -> finer sweet-spot search + ck["checkpoints_total_limit"] = 24 # keep them all + d["training"]["loop"]["max_train_steps"] = 200 + # validation cadence: cheap gate signal during the run + v = d.get("callbacks", {}).get("validation") + if isinstance(v, dict): + v["every_steps"] = 50 + v["run_at_start"] = False + d["training"]["tracker"]["run_name"] = f"qad-nvfp4-4call-{tag}-v2-shortckpt" + p.write_text(yaml.safe_dump(d, sort_keys=False)) + print(f"{p.name}: ckpt every {ck['training_state_checkpointing_steps']} x{ck['checkpoints_total_limit']}, " + f"max {d['training']['loop']['max_train_steps']} steps, val every {v.get('every_steps')}") diff --git a/scripts/compacth3/quantization/export_lane_int8.sh b/scripts/compacth3/quantization/export_lane_int8.sh new file mode 100644 index 0000000000..61b674738b --- /dev/null +++ b/scripts/compacth3/quantization/export_lane_int8.sh @@ -0,0 +1,26 @@ +#!/bin/bash +# Export pre-quantized int8 DiT weights for the FastH3 DMD2 student. +# Run: sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/exp.sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/export_lane_int8.sh +set -euo pipefail +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +M=/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1 +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +RUN=${SPRINT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +CKPT=${RUN}/inference/checkpoint-1400 +OUT=${CKPT}/exports/int8 + +mkdir -p "$OUT" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache +export PYTHONDONTWRITEBYTECODE=1 +export PYTHONPATH="${M}:${SPRINT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True + +cd "$M" +echo "=== exporting int8 -> $OUT ===" +"$PY" "${SPRINT}/export_quant_dit.py" \ + --lane int8 \ + --model-path "$CKPT" \ + --out "$OUT" \ + --num-gpus 1 +echo "=== done: $OUT ===" +ls -la "$OUT" diff --git a/scripts/compacth3/quantization/export_lane_nvfp4.sh b/scripts/compacth3/quantization/export_lane_nvfp4.sh new file mode 100644 index 0000000000..80c7b71a94 --- /dev/null +++ b/scripts/compacth3/quantization/export_lane_nvfp4.sh @@ -0,0 +1,26 @@ +#!/bin/bash +# Export pre-quantized nvfp4 DiT weights for the FastH3 DMD2 student. +# Run: sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/exp.sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/export_lane_nvfp4.sh +set -euo pipefail +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +M=/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1 +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +RUN=${SPRINT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +CKPT=${RUN}/inference/checkpoint-1400 +OUT=${CKPT}/exports/nvfp4 + +mkdir -p "$OUT" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache +export PYTHONDONTWRITEBYTECODE=1 +export PYTHONPATH="${M}:${SPRINT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True + +cd "$M" +echo "=== exporting nvfp4 -> $OUT ===" +"$PY" "${SPRINT}/export_quant_dit.py" \ + --lane nvfp4 \ + --model-path "$CKPT" \ + --out "$OUT" \ + --num-gpus 1 +echo "=== done: $OUT ===" +ls -la "$OUT" diff --git a/scripts/compacth3/quantization/export_lane_w4a16.sh b/scripts/compacth3/quantization/export_lane_w4a16.sh new file mode 100644 index 0000000000..8f20316bb9 --- /dev/null +++ b/scripts/compacth3/quantization/export_lane_w4a16.sh @@ -0,0 +1,26 @@ +#!/bin/bash +# Export pre-quantized w4a16 DiT weights for the FastH3 DMD2 student. +# Run: sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/exp.sbatch /mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829/export_lane_w4a16.sh +set -euo pipefail +SPRINT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +M=/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1 +PY=/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python +RUN=${SPRINT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +CKPT=${RUN}/inference/checkpoint-1400 +OUT=${CKPT}/exports/w4a16 + +mkdir -p "$OUT" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache +export PYTHONDONTWRITEBYTECODE=1 +export PYTHONPATH="${M}:${SPRINT}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True + +cd "$M" +echo "=== exporting w4a16 -> $OUT ===" +"$PY" "${SPRINT}/export_quant_dit.py" \ + --lane w4a16 \ + --model-path "$CKPT" \ + --out "$OUT" \ + --num-gpus 1 +echo "=== done: $OUT ===" +ls -la "$OUT" diff --git a/scripts/compacth3/quantization/export_quant_dit.py b/scripts/compacth3/quantization/export_quant_dit.py new file mode 100644 index 0000000000..e145ac5e3b --- /dev/null +++ b/scripts/compacth3/quantization/export_quant_dit.py @@ -0,0 +1,213 @@ +"""Export pre-quantized DiT weights for the FastH3 DMD2 student to safetensors. + +Why this does NOT go through ``VideoGenerator`` +--------------------------------------------- +The DiT is worker-resident: ``VideoGenerator.from_config(...)`` returns an +orchestrator handle whose process never holds an ``nn.Module`` for the +transformer (introspecting it prints "NO nn.Module attribute on the +generator"). So we load the transformer directly with the very loader the +generator uses -- ``PipelineComponentLoader.load_module`` -- and a +``FastVideoArgs`` built through the supported compat adapter +``fastvideo.api.compat.generator_config_to_fastvideo_args``, which is what +pins ``transformer_quant`` onto ``dit_config.quant_config``. + +That pin is what makes the linears get built with the quant method attached; +the loader's post-load hook ``_maybe_quantize_model`` +(``fastvideo/models/loader/fsdp_load.py``) then dispatches on that method and +calls ``convert_model_to_`` to materialize the quantized buffers. + +Must be a real file on disk (not stdin): FastVideo workers re-execute +``__main__`` via ``runpy``, and a heredoc has no path. +""" +from __future__ import annotations + +import argparse +import importlib.util +import json +import os +import sys +import time +from pathlib import Path + +SPRINT = "/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829" +M = "/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1" +HARNESS = f"{M}/examples/inference/basic/basic_fasth3.py" + +# lane -> the registry name resolved by +# fastvideo.layers.quantization.get_quantization_config +LANES = { + "nvfp4": "NVFP4H3", + "int8": "INT8Affine", + "w4a16": "W4A16", +} + +# Buffer names carrying the quantized payload, per scheme. NVFP4 has a real +# serializer/deserializer pair (save_nvfp4_checkpoint / load_nvfp4_checkpoint); +# the other two do not (see the lane note printed by main()). +BUFFERS = { + "nvfp4": ("_nvfp4_weight", "_nvfp4_weight_scale", "_weight_global_sf", "_nvfp4_alpha"), + "int8": ("_int8_affine_codes", "_int8_affine_scales", "_int8_affine_biases"), + "w4a16": ("_w4a16_codes", "_w4a16_scales", "_w4a16_zeros"), +} + + +def log(msg: str) -> None: + print(f"[export] {msg}", flush=True) + + +def build_fastvideo_args(model_path: str, quant_name: str, num_gpus: int): + """Harness args -> api GeneratorConfig -> FastVideoArgs (supported path).""" + spec = importlib.util.spec_from_file_location("fasth3_harness", HARNESS) + harness = importlib.util.module_from_spec(spec) + sys.modules["fasth3_harness"] = harness + spec.loader.exec_module(harness) + + # NOTE: argparse treats a passed sequence as the full argument list (it does + # not strip a program name), so no prog element here. + argv = [ + "--model-path", model_path, + "--prompt", "quant-export", + "--num-gpus", str(num_gpus), + "--transformer-quant", quant_name, + "--no-fa4", + "--no-inference-torch-compile", + "--steps", "5", + ] + args = harness.parse_args(argv) + args.fa4 = False + harness.configure_environment(args) + + config = harness.build_generator_config(args) + log(f"GeneratorConfig built: model_path={config.model_path} " + f"num_gpus={config.engine.num_gpus} " + f"transformer_quant={config.engine.quantization.transformer_quant}") + + from fastvideo.api.compat import generator_config_to_fastvideo_args + fastvideo_args = generator_config_to_fastvideo_args(config) + log(f"FastVideoArgs built: inference_mode={fastvideo_args.inference_mode} " + f"training_mode={fastvideo_args.training_mode} " + f"use_fsdp_inference={fastvideo_args.use_fsdp_inference} " + f"hsdp_shard_dim={fastvideo_args.hsdp_shard_dim}") + return fastvideo_args + + +def load_dit(fastvideo_args, transformer_path: str): + from fastvideo.models.loader.component_loader import PipelineComponentLoader + + dit_config = fastvideo_args.pipeline_config.dit_config + quant_config = getattr(dit_config, "quant_config", None) + if quant_config is None: + raise RuntimeError( + "dit_config.quant_config is None after the compat adapter ran -- the " + "quant config was not pinned, so no linear will be built quantized.") + log(f"dit_config.quant_config = {type(quant_config).__name__} (name={quant_config.get_name()})") + + log(f"loading transformer from {transformer_path}") + model = PipelineComponentLoader.load_module( + module_name="transformer", + component_model_path=transformer_path, + transformers_or_diffusers="diffusers", + fastvideo_args=fastvideo_args, + ) + log(f"loaded class={type(model).__name__}") + return model + + +def scheme_tagged(model, lane: str) -> list[tuple[str, object]]: + """(fqn, module) pairs whose quant_method belongs to this lane's scheme.""" + from fastvideo.layers.quantization.int8_affine_config import INT8AffineQuantizeMethod + from fastvideo.layers.quantization.nvfp4_config import NVFP4QuantizeMethod + from fastvideo.layers.quantization.w4a16_config import W4A16QuantizeMethod + + wanted = { + "nvfp4": NVFP4QuantizeMethod, + "int8": INT8AffineQuantizeMethod, + "w4a16": W4A16QuantizeMethod, + }[lane] + return [(fqn, mod) for fqn, mod in model.named_modules() + if isinstance(getattr(mod, "quant_method", None), wanted)] + + +def save_int8_or_w4a16(model, lane: str, path: str, tagged) -> dict: + """No serializer exists for this lane in this checkout. + + ``fastvideo/layers/quantization/{int8_affine,w4a16}_config.py`` export no + ``save_*_checkpoint`` and no sidecar format -- their ``__all__`` lists only + the quantizer, the config, the quantize method and the converter. Writing + a made-up layout here would produce a file no loader can read, so this + stops with the evidence instead. + """ + raise RuntimeError( + f"lane {lane!r}: no sidecar serializer exists in this checkout. " + f"grepped 'save_int8_affine_checkpoint' / 'save_w4a16_checkpoint' across the " + f"whole tree at {M}: zero hits. " + f"{len(tagged)} layers ARE quantized in memory (conversion receipt above is real) " + f"and their buffers are {BUFFERS[lane]}, but there is no encoder -- and no " + f"matching decoder -- so nothing could load the result.") + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--lane", required=True, choices=sorted(LANES)) + parser.add_argument("--model-path", required=True, + help="checkpoint dir, e.g. .../inference/checkpoint-1400") + parser.add_argument("--out", required=True, help="output directory for the sidecar") + parser.add_argument("--num-gpus", type=int, default=1, + help="1 keeps the load whole-model on cuda:0 (no FSDP)") + args = parser.parse_args() + + lane = args.lane + quant_name = LANES[lane] + transformer_path = os.path.join(args.model_path, "transformer") + if not os.path.isdir(transformer_path): + raise SystemExit(f"no transformer/ subdir under {args.model_path}") + + out_dir = Path(args.out) + out_dir.mkdir(parents=True, exist_ok=True) + + log(f"lane={lane} quant_name={quant_name}") + log(f"model_path={args.model_path}") + log(f"out_dir={out_dir}") + + t0 = time.time() + fastvideo_args = build_fastvideo_args(args.model_path, quant_name, args.num_gpus) + model = load_dit(fastvideo_args, transformer_path) + tagged = scheme_tagged(model, lane) + log(f"CONVERSION CHECK: {len(tagged)} {quant_name}-tagged linear layers " + f"present after load (conversion receipt above, from _maybe_quantize_model)") + if not tagged: + raise RuntimeError( + f"no {quant_name}-tagged linears found: the quant config did not cover " + f"any layer path, so nothing was quantized. Refusing to write an empty sidecar.") + + total_params = sum(p.numel() for p in model.parameters()) + log(f"model parameters: {total_params / 1e9:.2f}B load+convert took {time.time() - t0:.1f}s") + + if lane == "nvfp4": + from fastvideo.layers.quantization.nvfp4_config import save_nvfp4_checkpoint + target = out_dir / "nvfp4_weights.safetensors" + receipt = save_nvfp4_checkpoint( + model, target, + extra_metadata={"model": "FastH3-20B-42block-DMD2", "source": args.model_path}) + else: + save_int8_or_w4a16(model, lane, str(out_dir), tagged) + + size_bytes = os.path.getsize(receipt["path"]) + log("=" * 72) + log(f"EXPORT RECEIPT (lane={lane})") + log(json.dumps(receipt, indent=2)) + log(f"FILE PATH : {receipt['path']}") + log(f"FILE SIZE : {size_bytes / 1e9:.3f} GB ({size_bytes / (1 << 30):.3f} GiB)") + log(f"MODULE COUNT: {len(tagged)}") + log("=" * 72) + + with open(os.path.join(out_dir, "export_receipt.json"), "w") as handle: + json.dump({"lane": lane, "quant_name": quant_name, "file": receipt["path"], + "size_bytes": size_bytes, "module_count": len(tagged), + "receipt": receipt}, handle, indent=2) + log(f"wrote {os.path.join(out_dir, 'export_receipt.json')}") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/compacth3/quantization/run_export_nvfp4.sh b/scripts/compacth3/quantization/run_export_nvfp4.sh new file mode 100644 index 0000000000..2dd4be76b5 --- /dev/null +++ b/scripts/compacth3/quantization/run_export_nvfp4.sh @@ -0,0 +1,23 @@ +#!/bin/bash +# Export compact NVFP4 weights using the worker-side hook. +set -euo pipefail +S=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +M=/mnt/nfs/vlm-aryan/fasth3-h3-serve-cookbook-eval-20260831/repo-main-3d8ac9d1 +RUN=${S}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3/job-paired8972-8975-4000-v3 +OUT=${RUN}/exports/nvfp4-ckpt1400 +mkdir -p "$OUT" +export HF_HOME=/mnt/nfs/vlm-aryan/hf-cache PYTHONDONTWRITEBYTECODE=1 +export PYTHONPATH="${M}:${S}/python-packages:/mnt/nfs/vlm-aryan/fastvideo-wan-venv/lib/python3.12/site-packages" +export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA FASTVIDEO_MINIMAX_H3_FUSIONS=0 +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True +export FASTVIDEO_EXPORT_QUANT_SIDECAR="$OUT" +export FASTVIDEO_EXPORT_LANE=nvfp4 +cd "$M" +echo "=== NVFP4 export -> $OUT ===" +/mnt/nfs/vlm-aryan/fastvideo-wan-venv/bin/python examples/inference/basic/basic_fasth3.py \ + --model-path "$RUN/inference/checkpoint-1400" \ + --prompt "export" --output "$OUT/tmp" \ + --height 480 --width 832 --num-frames 124 --steps 5 --num-gpus 4 \ + --repeats 1 --transformer-quant NVFP4H3 --no-fa4 2>&1 | grep -aE "EXPORT|Converting loaded|purge receipt|Error|Traceback" | tail -20 +echo "=== output dir ===" +ls -la "$OUT" diff --git a/scripts/compacth3/resume_release20b_dmd2_paired_generic.sh b/scripts/compacth3/resume_release20b_dmd2_paired_generic.sh new file mode 100644 index 0000000000..d7933f22fb --- /dev/null +++ b/scripts/compacth3/resume_release20b_dmd2_paired_generic.sh @@ -0,0 +1,70 @@ +#!/bin/bash +# Launch one 32-GPU DMD2 world across two already-running four-node jobs. +# Usage: bash resume_release20b_dmd2_paired_generic.sh JOB_A JOB_B RESUME_STEP +set -euo pipefail + +if [[ "$#" -ne 3 ]]; then + echo "Usage: $0 JOB_A JOB_B RESUME_STEP" >&2 + exit 2 +fi + +JOB_A="$1" +JOB_B="$2" +RESUME_STEP="$3" +CHECKPOINT_EVERY="${CHECKPOINT_EVERY:-200}" +VALIDATION_EVERY="${VALIDATION_EVERY:-200}" +SPRINT_ROOT=/mnt/nfs/vlm-aryan/fasth3-14b-2step-qad-20260829 +CODE_ROOT=${SPRINT_ROOT}/code/release20b-dmd2-v12-corrected-v17 +SELECTED_PARENT=${SPRINT_ROOT}/runs/release20b-folded-long-4k-v3/dmd-parent-step750-complete-v1 +TEACHER_PARENT=${SPRINT_ROOT}/release-candidates/base-h3-teacher-complete-v1 +OUTPUT_BASE=${SPRINT_ROOT}/runs/release20b-dmd2-v12-corrected-c4-parent750-32gpu-4000-v3 +RUN_ID=paired8972-8975-4000-v3 +OUTPUT_ROOT=${OUTPUT_BASE}/job-${RUN_ID} +RESUME_PATH=${OUTPUT_ROOT}/checkpoint-${RESUME_STEP} + +[[ "$(squeue -h -j "${JOB_A}" -o %T)" == "RUNNING" ]] +[[ "$(squeue -h -j "${JOB_B}" -o %T)" == "RUNNING" ]] +NODES_A="$(squeue -h -j "${JOB_A}" -o %N)" +NODES_B="$(squeue -h -j "${JOB_B}" -o %N)" +[[ "$(scontrol show hostnames "${NODES_A}" | wc -l)" -eq 4 ]] +[[ "$(scontrol show hostnames "${NODES_B}" | wc -l)" -eq 4 ]] + +# A resumable checkpoint is committed only after all distributed state and all +# rank-local RNG snapshots exist. Never fall back to a partially written save. +test -s "${RESUME_PATH}/.complete" +test -s "${RESUME_PATH}/dcp/.metadata" +test -s "${RESUME_PATH}/metadata.json" +[[ "$(find "${RESUME_PATH}/dcp" -maxdepth 1 -type f -name '*.distcp' | wc -l)" -eq 32 ]] +[[ "$(find "${RESUME_PATH}" -maxdepth 1 -type f -name 'rng_state_rank*.pt' | wc -l)" -eq 32 ]] + +MASTER_ADDR="$(scontrol show hostnames "${NODES_A}" | head -n 1)" +MASTER_PORT="$((30000 + JOB_A % 10000))" + +launch_half() { + local job_id="$1" node_list="$2" rank_base="$3" + local log_path=${SPRINT_ROOT}/resume-dmd2-${RUN_ID}-job${job_id}-from${RESUME_STEP}.log + nohup srun --overlap --jobid="${job_id}" --nodes=4 --ntasks=4 \ + --ntasks-per-node=1 --gres=gpu:4 --nodelist="${node_list}" \ + --kill-on-bad-exit=1 \ + --container-image='nvcr.io#nvidia/pytorch:25.06-py3' \ + --container-mounts=/mnt/nfs:/mnt/nfs,/mnt/lustre/vlm-shared:/mnt/lustre/vlm-shared:ro \ + env CODE_ROOT="${CODE_ROOT}" SELECTED_PARENT="${SELECTED_PARENT}" \ + TEACHER_PARENT="${TEACHER_PARENT}" OUTPUT_BASE="${OUTPUT_BASE}" \ + SPRINT_ROOT="${SPRINT_ROOT}" MASTER_ADDR="${MASTER_ADDR}" \ + MASTER_PORT="${MASTER_PORT}" NNODES=8 NODE_RANK_BASE="${rank_base}" \ + RUN_ID="${RUN_ID}" PRODUCTION_TARGET=4000 RESUME_PATH="${RESUME_PATH}" \ + CHECKPOINT_EVERY="${CHECKPOINT_EVERY}" VALIDATION_EVERY="${VALIDATION_EVERY}" \ + bash "${CODE_ROOT}/scripts/run_release20b_dmd2_v12_16gpu.sh" \ + >"${log_path}" 2>&1 & + echo "$!" > "${log_path}.pid" +} + +launch_half "${JOB_A}" "${NODES_A}" 0 +launch_half "${JOB_B}" "${NODES_B}" 4 + +printf 'job_a=%s\njob_b=%s\nresume_step=%s\ncheckpoint_every=%s\nvalidation_every=%s\nmaster=%s:%s\nlaunched_at=%s\n' \ + "${JOB_A}" "${JOB_B}" "${RESUME_STEP}" "${CHECKPOINT_EVERY}" \ + "${VALIDATION_EVERY}" "${MASTER_ADDR}" "${MASTER_PORT}" "$(date -u +%FT%TZ)" \ + > "${OUTPUT_ROOT}/paired-rollover-from-${RESUME_STEP}.receipt" + +echo "Launched one 32-GPU world from checkpoint ${RESUME_STEP} across ${JOB_A}+${JOB_B}." diff --git a/scripts/compacth3/run_eval_lane.sh b/scripts/compacth3/run_eval_lane.sh new file mode 100644 index 0000000000..04a49dab88 --- /dev/null +++ b/scripts/compacth3/run_eval_lane.sh @@ -0,0 +1,89 @@ +#!/bin/bash +# Run the real eval set, honouring each case's declared generation exactly. +# usage: run_eval_lane.sh