diff --git a/docs/guides/configuration.mdx b/docs/guides/configuration.mdx index dbe47dfdb4..5e47d36f87 100644 --- a/docs/guides/configuration.mdx +++ b/docs/guides/configuration.mdx @@ -63,6 +63,34 @@ Only load configuration files and `_target_` values from trusted sources. A YAML The `distributed:` section is **not** instantiated using `_target_`. Recipes parse it with a fixed schema. Use `strategy: fsdp2`, `strategy: ddp`, or `strategy: megatron_fsdp`. You can also configure parallelism sizes, such as `dp_size`, `tp_size`, and `pp_size`, and strategy-specific options. When pipeline parallelism is enabled (`pp_size > 1`), add a `pipeline:` subsection with options such as `pp_schedule`, `pp_microbatch_size`, and `layers_per_stage`. For examples, see the [Pipeline Parallelism with AutoPipeline](/development/pipeline-parallelism) guide and the recipe configs. +## Select Trainable Modules + +The LLM and VLM fine-tuning recipes accept a `freeze_config:` section. Use `freeze_modules` and `unfreeze_modules` with typed selectors to control which modules are trainable: + +```yaml +freeze_config: + freeze_modules: + - path: vision_tower # exact canonical module path + - path: audio_tower # repeat the key to select more exact module paths + - glob: "*.speech_encoder" # case-sensitive shell-style glob on the full module path + unfreeze_modules: + - path: multi_modal_projector +``` + +Each selector is a mapping with exactly one key — `path` matches one module by its exact fully qualified name, while `glob` matches module names with `fnmatch`-style wildcards (`*` crosses `.` separators, so `*_proj` matches projection modules at any depth). List entries are additive: repeat `path` or `glob` entries to select several modules, and a single entry combining both keys is rejected. Matching is recursive: a selected module's entire subtree is frozen or unfrozen. Bare strings are rejected; unknown options and selectors that match no parameters raise an error before training starts. + +Trainability is resolved in a fixed order. Full fine-tuning preserves the model's existing `requires_grad` state, including intentionally frozen model-owned parameters; PEFT establishes a LoRA-trainable, base-frozen baseline. `freeze_modules` selectors then freeze their modules, and `unfreeze_modules` selectors unfreeze theirs, winning on overlap — this makes combinations such as LoRA plus a fully trainable multimodal projector a two-line configuration. Freezing only controls `requires_grad`; it never changes a parameter's storage dtype. The policy is validated on the complete model, re-resolved after tensor/expert parallel and activation-checkpointing surgery immediately before DDP/FSDP construction, and resolved once more after checkpoint loading before optimizer construction. + +The modality-specific booleans `freeze_vision_tower`, `freeze_audio_tower`, `freeze_language_model`, and `freeze_video_embedder` remain supported and keep their established attribute and substring matching, applied before `unfreeze_modules`: + +```yaml +freeze_config: + freeze_vision_tower: true + freeze_audio_tower: true +``` + +Legacy-only configurations retain the implicit `freeze_vision_tower: true` default. A configuration that declares `freeze_modules` or `unfreeze_modules`, even as an empty list, uses explicit-selector semantics and does not implicitly freeze vision modules; add a vision selector (for example, `glob: "*vision*"`) when such a configuration should also freeze the vision tower. Omit both selector fields to retain legacy behavior; `null`, `false`, and numeric values are rejected. + ## Prewarm One-Time CUDA Initialization {/* docs-review-start: mamba-ssd-prewarm */} diff --git a/docs/guides/llm/sequence-classification.mdx b/docs/guides/llm/sequence-classification.mdx index 82cf3b54e4..419bde7a1d 100644 --- a/docs/guides/llm/sequence-classification.mdx +++ b/docs/guides/llm/sequence-classification.mdx @@ -1,17 +1,17 @@ --- title: "Sequence Classification (SFT/PEFT) with NeMo AutoModel" -description: "" +description: "Train a sequence classification model with NeMo AutoModel using GLUE MRPC, RoBERTa, and optional LoRA." position: 10 --- ## Introduction -Sequence classification tasks (e.g., sentiment analysis, topic classification, GLUE tasks) map input text to a discrete label. NeMo AutoModel provides a lightweight recipe specialized for this setting that integrates with popular pretrained model formats and dataset sources. Integration with Hugging Face is supported. +Sequence classification tasks (for example, sentiment analysis, topic classification, and GLUE tasks) map input text to a discrete label. NeMo AutoModel provides a lightweight recipe specialized for this setting that integrates with popular pretrained model formats and dataset sources. Integration with Hugging Face is supported. -This guide shows how to train a sequence classification model using the `TrainFinetuneRecipeForSequenceClassification` recipe, including optional Parameter-Efficient Fine-Tuning (LoRA). +This guide shows how to train a sequence classification model using the `TrainFinetuneRecipeForSequenceClassification` recipe, including optional Parameter-Efficient Fine-Tuning (PEFT) with LoRA. ## Quickstart -Use the example config for GLUE MRPC with RoBERTa-large + LoRA: +Use the example config for GLUE MRPC with RoBERTa-large and LoRA: ```bash uv run automodel examples/llm_seq_cls/glue/mrpc_roberta_lora.yaml @@ -66,6 +66,7 @@ distributed: tp_size: 1 cp_size: 1 sequence_parallel: false + autocast_dtype: bfloat16 peft: _target_: nemo_automodel.components._peft.lora.PeftConfig @@ -76,6 +77,10 @@ peft: alpha: 16 dropout: 0.1 +freeze_config: + unfreeze_modules: + - glob: "*classifier" + dataset: _target_: nemo_automodel.components.datasets.llm.seq_cls.GLUE_MRPC split: train @@ -108,8 +113,10 @@ optimizer: ## LoRA (PEFT) Settings -- `target_modules`: glob to select linear layers (e.g., `"*.proj"`). -- `dim` (rank), `alpha`, `dropout`: tune per model/compute budget. Values `dim=8, alpha=16, dropout=0.1` are a good starting point for RoBERTa. +- `target_modules`: Glob to select linear layers (for example, `"*.proj"`). +- `dim` (rank), `alpha`, `dropout`: Tune per model and compute budget. Values `dim=8, alpha=16, dropout=0.1` are a good starting point for RoBERTa. +- `freeze_config.unfreeze_modules`: Keeps the classification head fully trainable while PEFT freezes other non-LoRA parameters. +- `distributed.autocast_dtype`: Runs the forward pass in the selected compute dtype while trainable parameters retain their configured storage dtype, including when single-rank FSDP is skipped. - The recipe automatically applies the adapters; no additional code changes are required. ## Running on Multiple GPUs diff --git a/docs/guides/vlm/gemma4.md b/docs/guides/vlm/gemma4.md index 3f4a0e4638..76a36ca153 100644 --- a/docs/guides/vlm/gemma4.md +++ b/docs/guides/vlm/gemma4.md @@ -328,7 +328,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/docs/guides/vlm/gemma4.mdx b/docs/guides/vlm/gemma4.mdx index 4e655d0a7a..4f90e346e5 100644 --- a/docs/guides/vlm/gemma4.mdx +++ b/docs/guides/vlm/gemma4.mdx @@ -329,7 +329,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/docs/guides/vlm/nemotron-omni.md b/docs/guides/vlm/nemotron-omni.md index 75606ac1ce..43999be52e 100644 --- a/docs/guides/vlm/nemotron-omni.md +++ b/docs/guides/vlm/nemotron-omni.md @@ -152,7 +152,6 @@ distributed: ep_size: 8 # 128 MoE experts across 8 GPUs freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/docs/guides/vlm/nemotron-omni.mdx b/docs/guides/vlm/nemotron-omni.mdx index 4854eed3e0..8d7d6bc420 100644 --- a/docs/guides/vlm/nemotron-omni.mdx +++ b/docs/guides/vlm/nemotron-omni.mdx @@ -156,7 +156,6 @@ distributed: ep_size: 8 # 128 MoE experts across 8 GPUs freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/audio_finetune/qwen2_5_omni_asr/ami_sft_3b.yaml b/examples/audio_finetune/qwen2_5_omni_asr/ami_sft_3b.yaml index 45e6578163..7ecd1bb796 100644 --- a/examples/audio_finetune/qwen2_5_omni_asr/ami_sft_3b.yaml +++ b/examples/audio_finetune/qwen2_5_omni_asr/ami_sft_3b.yaml @@ -59,7 +59,6 @@ distributed: # Full fine-tune: audio tower + language model trainable; vision tower frozen # (we never feed images on AMI). freeze_config: - freeze_embeddings: false freeze_vision_tower: true freeze_audio_tower: false freeze_language_model: false diff --git a/examples/audio_finetune/qwen2_5_omni_asr/ami_sft_7b.yaml b/examples/audio_finetune/qwen2_5_omni_asr/ami_sft_7b.yaml index a262cc6b49..8a68afd6cf 100644 --- a/examples/audio_finetune/qwen2_5_omni_asr/ami_sft_7b.yaml +++ b/examples/audio_finetune/qwen2_5_omni_asr/ami_sft_7b.yaml @@ -62,7 +62,6 @@ distributed: # Full fine-tune: audio tower + language model trainable; vision tower frozen # (we never feed images on AMI). freeze_config: - freeze_embeddings: false freeze_vision_tower: true freeze_audio_tower: false freeze_language_model: false diff --git a/examples/audio_finetune/qwen2_5_omni_asr/multi_en_sft_3b.yaml b/examples/audio_finetune/qwen2_5_omni_asr/multi_en_sft_3b.yaml index edd4297785..0406e7349b 100644 --- a/examples/audio_finetune/qwen2_5_omni_asr/multi_en_sft_3b.yaml +++ b/examples/audio_finetune/qwen2_5_omni_asr/multi_en_sft_3b.yaml @@ -72,7 +72,6 @@ distributed: # Full fine-tune: audio tower + language model trainable; vision tower frozen # (we never feed images on these English ASR corpora). freeze_config: - freeze_embeddings: false freeze_vision_tower: true freeze_audio_tower: false freeze_language_model: false diff --git a/examples/audio_finetune/qwen3_omni_asr/ami_sft.yaml b/examples/audio_finetune/qwen3_omni_asr/ami_sft.yaml index 9b6142713c..3eb6ee819b 100644 --- a/examples/audio_finetune/qwen3_omni_asr/ami_sft.yaml +++ b/examples/audio_finetune/qwen3_omni_asr/ami_sft.yaml @@ -88,7 +88,6 @@ distributed: # Full fine-tune: audio tower + language model trainable; only the vision tower # stays frozen (we never feed images). No PEFT adapters. freeze_config: - freeze_embeddings: false freeze_vision_tower: true freeze_audio_tower: false freeze_language_model: false diff --git a/examples/audio_finetune/qwen3_omni_asr/multi_en_sft.yaml b/examples/audio_finetune/qwen3_omni_asr/multi_en_sft.yaml index 2cc9c94cae..b9a11babd4 100644 --- a/examples/audio_finetune/qwen3_omni_asr/multi_en_sft.yaml +++ b/examples/audio_finetune/qwen3_omni_asr/multi_en_sft.yaml @@ -82,7 +82,6 @@ distributed: # Full fine-tune: audio tower + language model trainable; only the vision tower # stays frozen (we never feed images on these English ASR corpora). No PEFT. freeze_config: - freeze_embeddings: false freeze_vision_tower: true freeze_audio_tower: false freeze_language_model: false diff --git a/examples/convergence/tulu3/models/gemma4-31b/gemma4_31b_tulu3_packed2k_cp1_gbs32_4node.yaml b/examples/convergence/tulu3/models/gemma4-31b/gemma4_31b_tulu3_packed2k_cp1_gbs32_4node.yaml index 6d705d0a1b..2729976608 100644 --- a/examples/convergence/tulu3/models/gemma4-31b/gemma4_31b_tulu3_packed2k_cp1_gbs32_4node.yaml +++ b/examples/convergence/tulu3/models/gemma4-31b/gemma4_31b_tulu3_packed2k_cp1_gbs32_4node.yaml @@ -100,7 +100,6 @@ lr_scheduler: min_lr: 1.0e-6 freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/llm_seq_cls/glue/mrpc_roberta_lora.yaml b/examples/llm_seq_cls/glue/mrpc_roberta_lora.yaml index 26e17395a1..8a0cbdd6f8 100644 --- a/examples/llm_seq_cls/glue/mrpc_roberta_lora.yaml +++ b/examples/llm_seq_cls/glue/mrpc_roberta_lora.yaml @@ -38,17 +38,21 @@ distributed: cp_size: 1 sequence_parallel: false + autocast_dtype: bfloat16 peft: _target_: nemo_automodel.components._peft.lora.PeftConfig target_modules: - "*.query" - "*.value" - # Note: classifier head is fully trained (not LoRA), unfrozen automatically in train_seq_cls.py dim: 8 alpha: 16 dropout: 0.1 +freeze_config: + unfreeze_modules: + - glob: "*classifier" + dataset: _target_: nemo_automodel.components.datasets.llm.seq_cls.GLUE_MRPC split: train @@ -71,4 +75,3 @@ optimizer: eps: 1e-8 lr: 2.0e-5 # Standard learning rate for BERT/RoBERTa fine-tuning weight_decay: 0.01 # Crucial for stable training on small datasets - diff --git a/examples/long_context_validation/gemma4_31B/gemma4_31b_base_coderforge_cp8_64k_1e5_800steps.yaml b/examples/long_context_validation/gemma4_31B/gemma4_31b_base_coderforge_cp8_64k_1e5_800steps.yaml index 331e96792f..fb7ced1cec 100644 --- a/examples/long_context_validation/gemma4_31B/gemma4_31b_base_coderforge_cp8_64k_1e5_800steps.yaml +++ b/examples/long_context_validation/gemma4_31B/gemma4_31b_base_coderforge_cp8_64k_1e5_800steps.yaml @@ -108,7 +108,6 @@ lr_scheduler: min_lr: 5.0e-7 freeze_config: - freeze_embeddings: true # NO-OP in this path (apply_parameter_freezing ignores it); kept for parity freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false # embeddings + tied LM head trainable — required to learn tokens 48/49/52 diff --git a/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2.yaml b/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2.yaml index 9351ca15aa..e929517aee 100644 --- a/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2.yaml +++ b/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2.yaml @@ -88,7 +88,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2_megatron_fsdp.yaml b/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2_megatron_fsdp.yaml index 8ab0252f59..2477d5e14e 100644 --- a/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2_megatron_fsdp.yaml +++ b/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2_megatron_fsdp.yaml @@ -86,7 +86,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2_peft.yaml b/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2_peft.yaml index fde2ffe79a..9a90d6f642 100644 --- a/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2_peft.yaml +++ b/examples/vlm_finetune/gemma3/gemma3_vl_4b_cord_v2_peft.yaml @@ -99,7 +99,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma3/gemma3_vl_4b_medpix.yaml b/examples/vlm_finetune/gemma3/gemma3_vl_4b_medpix.yaml index bff427ca9b..087b7db824 100644 --- a/examples/vlm_finetune/gemma3/gemma3_vl_4b_medpix.yaml +++ b/examples/vlm_finetune/gemma3/gemma3_vl_4b_medpix.yaml @@ -89,7 +89,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma3/gemma3_vl_4b_medpix_peft.yaml b/examples/vlm_finetune/gemma3/gemma3_vl_4b_medpix_peft.yaml index 150adc79c5..428aac61d9 100644 --- a/examples/vlm_finetune/gemma3/gemma3_vl_4b_medpix_peft.yaml +++ b/examples/vlm_finetune/gemma3/gemma3_vl_4b_medpix_peft.yaml @@ -100,7 +100,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma3n/gemma3n_vl_4b_medpix.yaml b/examples/vlm_finetune/gemma3n/gemma3n_vl_4b_medpix.yaml index ef94a29a3a..c759a314a5 100644 --- a/examples/vlm_finetune/gemma3n/gemma3n_vl_4b_medpix.yaml +++ b/examples/vlm_finetune/gemma3n/gemma3n_vl_4b_medpix.yaml @@ -90,7 +90,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma3n/gemma3n_vl_4b_medpix_peft.yaml b/examples/vlm_finetune/gemma3n/gemma3n_vl_4b_medpix_peft.yaml index 27f838df8f..c9051a27cc 100644 --- a/examples/vlm_finetune/gemma3n/gemma3n_vl_4b_medpix_peft.yaml +++ b/examples/vlm_finetune/gemma3n/gemma3n_vl_4b_medpix_peft.yaml @@ -103,7 +103,6 @@ peft: use_triton: True freeze_config: - freeze_embeddings: false freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe.yaml b/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe.yaml index bab78d17a4..7c22af1155 100644 --- a/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe.yaml @@ -106,7 +106,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_medpix_ep8cp2_4k.yaml b/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_medpix_ep8cp2_4k.yaml index ba9e21e633..a295df2fef 100644 --- a/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_medpix_ep8cp2_4k.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_medpix_ep8cp2_4k.yaml @@ -96,7 +96,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_mock.yaml b/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_mock.yaml index 2271071b5d..4f5ab76a1c 100644 --- a/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_mock.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_mock.yaml @@ -92,7 +92,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_packing.yaml b/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_packing.yaml index 45b0827cbd..e72f07c53a 100644 --- a/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_packing.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_packing.yaml @@ -96,7 +96,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_peft.yaml b/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_peft.yaml index 44cfa85853..c3e1ebf7ba 100644 --- a/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_peft.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_26b_a4b_moe_peft.yaml @@ -118,7 +118,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_2b.yaml b/examples/vlm_finetune/gemma4/gemma4_2b.yaml index 35bacea221..4f0afcb5a6 100644 --- a/examples/vlm_finetune/gemma4/gemma4_2b.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_2b.yaml @@ -92,7 +92,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_2b_peft.yaml b/examples/vlm_finetune/gemma4/gemma4_2b_peft.yaml index d3780fa092..2d38d6fffe 100644 --- a/examples/vlm_finetune/gemma4/gemma4_2b_peft.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_2b_peft.yaml @@ -106,7 +106,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_31b.yaml b/examples/vlm_finetune/gemma4/gemma4_31b.yaml index 9000992078..4a873749b0 100644 --- a/examples/vlm_finetune/gemma4/gemma4_31b.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_31b.yaml @@ -99,7 +99,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_31b_ffpa_mock_8k.yaml b/examples/vlm_finetune/gemma4/gemma4_31b_ffpa_mock_8k.yaml index d73375acce..d8faacb48c 100644 --- a/examples/vlm_finetune/gemma4/gemma4_31b_ffpa_mock_8k.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_31b_ffpa_mock_8k.yaml @@ -125,7 +125,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_31b_ffpa_mock_packing_cp8_16k.yaml b/examples/vlm_finetune/gemma4/gemma4_31b_ffpa_mock_packing_cp8_16k.yaml index 2d5846638b..ff297f7d79 100644 --- a/examples/vlm_finetune/gemma4/gemma4_31b_ffpa_mock_packing_cp8_16k.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_31b_ffpa_mock_packing_cp8_16k.yaml @@ -115,7 +115,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_31b_peft.yaml b/examples/vlm_finetune/gemma4/gemma4_31b_peft.yaml index 7b19016cd5..cd37710c98 100644 --- a/examples/vlm_finetune/gemma4/gemma4_31b_peft.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_31b_peft.yaml @@ -109,7 +109,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_31b_te.yaml b/examples/vlm_finetune/gemma4/gemma4_31b_te.yaml index 2f1b6c732f..80419b4e39 100644 --- a/examples/vlm_finetune/gemma4/gemma4_31b_te.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_31b_te.yaml @@ -104,7 +104,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_31b_tp4.yaml b/examples/vlm_finetune/gemma4/gemma4_31b_tp4.yaml index 2380f7086c..2ce8d534bd 100644 --- a/examples/vlm_finetune/gemma4/gemma4_31b_tp4.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_31b_tp4.yaml @@ -97,7 +97,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_31b_tp4_pp2.yaml b/examples/vlm_finetune/gemma4/gemma4_31b_tp4_pp2.yaml index ff3bc0d4cf..11c05c9c84 100644 --- a/examples/vlm_finetune/gemma4/gemma4_31b_tp4_pp2.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_31b_tp4_pp2.yaml @@ -106,7 +106,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_31b_tp4_pp4.yaml b/examples/vlm_finetune/gemma4/gemma4_31b_tp4_pp4.yaml index 3848cdb375..8be44527a6 100644 --- a/examples/vlm_finetune/gemma4/gemma4_31b_tp4_pp4.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_31b_tp4_pp4.yaml @@ -111,7 +111,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_31b_tulu3_text_cp8_16k.yaml b/examples/vlm_finetune/gemma4/gemma4_31b_tulu3_text_cp8_16k.yaml index 6d2aa75b4f..14c8ff8199 100644 --- a/examples/vlm_finetune/gemma4/gemma4_31b_tulu3_text_cp8_16k.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_31b_tulu3_text_cp8_16k.yaml @@ -97,7 +97,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_4b.yaml b/examples/vlm_finetune/gemma4/gemma4_4b.yaml index f10a50e6c8..09a5a1d925 100644 --- a/examples/vlm_finetune/gemma4/gemma4_4b.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_4b.yaml @@ -93,7 +93,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_4b_mock.yaml b/examples/vlm_finetune/gemma4/gemma4_4b_mock.yaml index 0d1d00de8a..6b0b74ce72 100644 --- a/examples/vlm_finetune/gemma4/gemma4_4b_mock.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_4b_mock.yaml @@ -82,7 +82,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_4b_peft.yaml b/examples/vlm_finetune/gemma4/gemma4_4b_peft.yaml index 61ecea86e3..40039799c2 100644 --- a/examples/vlm_finetune/gemma4/gemma4_4b_peft.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_4b_peft.yaml @@ -105,7 +105,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_4b_te.yaml b/examples/vlm_finetune/gemma4/gemma4_4b_te.yaml index 27cf8e1d28..1007dbafbf 100644 --- a/examples/vlm_finetune/gemma4/gemma4_4b_te.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_4b_te.yaml @@ -100,7 +100,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_e4b_tulu3_text_cp16_64k.yaml b/examples/vlm_finetune/gemma4/gemma4_e4b_tulu3_text_cp16_64k.yaml index 0737534107..5933debf0d 100644 --- a/examples/vlm_finetune/gemma4/gemma4_e4b_tulu3_text_cp16_64k.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_e4b_tulu3_text_cp16_64k.yaml @@ -121,7 +121,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4/gemma4_e4b_tulu3_text_tp2_cp2.yaml b/examples/vlm_finetune/gemma4/gemma4_e4b_tulu3_text_tp2_cp2.yaml index 87bbbdba07..759823ba23 100644 --- a/examples/vlm_finetune/gemma4/gemma4_e4b_tulu3_text_tp2_cp2.yaml +++ b/examples/vlm_finetune/gemma4/gemma4_e4b_tulu3_text_tp2_cp2.yaml @@ -117,7 +117,6 @@ optimizer: exp_avg_sq_dtype: torch.float32 freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4_joint_drafter/gemma4_31b_joint_drafter_tulu_magicoder_mix.yaml b/examples/vlm_finetune/gemma4_joint_drafter/gemma4_31b_joint_drafter_tulu_magicoder_mix.yaml index 7088badca4..39df4db8ab 100644 --- a/examples/vlm_finetune/gemma4_joint_drafter/gemma4_31b_joint_drafter_tulu_magicoder_mix.yaml +++ b/examples/vlm_finetune/gemma4_joint_drafter/gemma4_31b_joint_drafter_tulu_magicoder_mix.yaml @@ -143,7 +143,6 @@ freeze_config: # The drafter consumes the base's token embeddings via its pre_projection # input at EVERY recurrent round, so embeddings must stay trainable for # gradient flow from the drafter loss back to the base. - freeze_embeddings: false freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4_joint_drafter/gemma4_4b_joint_drafter_medpix.yaml b/examples/vlm_finetune/gemma4_joint_drafter/gemma4_4b_joint_drafter_medpix.yaml index b6191eadf5..e5bc90737f 100644 --- a/examples/vlm_finetune/gemma4_joint_drafter/gemma4_4b_joint_drafter_medpix.yaml +++ b/examples/vlm_finetune/gemma4_joint_drafter/gemma4_4b_joint_drafter_medpix.yaml @@ -132,7 +132,6 @@ freeze_config: # Embeddings must stay trainable: the drafter consumes the base's token # embeddings via its pre_projection input at EVERY round, so freezing them # would block gradient flow from the drafter loss to the base. - freeze_embeddings: false freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/gemma4_joint_drafter/gemma4_4b_joint_drafter_tulu_magicoder_mix.yaml b/examples/vlm_finetune/gemma4_joint_drafter/gemma4_4b_joint_drafter_tulu_magicoder_mix.yaml index 49208d25be..c79ddd02d4 100644 --- a/examples/vlm_finetune/gemma4_joint_drafter/gemma4_4b_joint_drafter_tulu_magicoder_mix.yaml +++ b/examples/vlm_finetune/gemma4_joint_drafter/gemma4_4b_joint_drafter_tulu_magicoder_mix.yaml @@ -132,7 +132,6 @@ freeze_config: # The drafter consumes the base's token embeddings via its pre_projection # input at EVERY recurrent round, so embeddings must stay trainable for # gradient flow from the drafter loss back to the base. - freeze_embeddings: false freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/internvl/internvl_3_5_4b.yaml b/examples/vlm_finetune/internvl/internvl_3_5_4b.yaml index a1d3f67da9..5418fa18c3 100644 --- a/examples/vlm_finetune/internvl/internvl_3_5_4b.yaml +++ b/examples/vlm_finetune/internvl/internvl_3_5_4b.yaml @@ -97,7 +97,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/kimi/kimi25vl_medpix.yaml b/examples/vlm_finetune/kimi/kimi25vl_medpix.yaml index f3f63f5cbe..95b7cc217c 100644 --- a/examples/vlm_finetune/kimi/kimi25vl_medpix.yaml +++ b/examples/vlm_finetune/kimi/kimi25vl_medpix.yaml @@ -120,7 +120,6 @@ optimizer: - 0.95 freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/kimi/kimi2vl_cordv2.yaml b/examples/vlm_finetune/kimi/kimi2vl_cordv2.yaml index 51cff19323..c1203df06d 100644 --- a/examples/vlm_finetune/kimi/kimi2vl_cordv2.yaml +++ b/examples/vlm_finetune/kimi/kimi2vl_cordv2.yaml @@ -112,7 +112,6 @@ optimizer: - 0.95 freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/minimax_m3/minimax_m3_vl_lora_pp4ep8_8node.yaml b/examples/vlm_finetune/minimax_m3/minimax_m3_vl_lora_pp4ep8_8node.yaml index d3b7e5d698..d025ad572a 100644 --- a/examples/vlm_finetune/minimax_m3/minimax_m3_vl_lora_pp4ep8_8node.yaml +++ b/examples/vlm_finetune/minimax_m3/minimax_m3_vl_lora_pp4ep8_8node.yaml @@ -107,7 +107,6 @@ distributed: wrap_outer_model: true freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_cp2_medpix_2k.yaml b/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_cp2_medpix_2k.yaml index 3d54946f9c..c3386fc3ed 100644 --- a/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_cp2_medpix_2k.yaml +++ b/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_cp2_medpix_2k.yaml @@ -85,7 +85,6 @@ distributed: wrap_outer_model: true freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_ep32pp4.yaml b/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_ep32pp4.yaml index 2aa40543a1..b9e900e466 100644 --- a/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_ep32pp4.yaml +++ b/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_ep32pp4.yaml @@ -89,7 +89,6 @@ distributed: wrap_outer_model: true freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_tulu3_text_cp8_16k.yaml b/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_tulu3_text_cp8_16k.yaml index 5362085f38..b52790a38b 100644 --- a/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_tulu3_text_cp8_16k.yaml +++ b/examples/vlm_finetune/minimax_m3/minimax_m3_vl_sft_tulu3_text_cp8_16k.yaml @@ -97,7 +97,6 @@ distributed: wrap_outer_model: true freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/mistral/ministral3_14b_medpix.yaml b/examples/vlm_finetune/mistral/ministral3_14b_medpix.yaml index 81a431b7f6..a9b826cdfd 100644 --- a/examples/vlm_finetune/mistral/ministral3_14b_medpix.yaml +++ b/examples/vlm_finetune/mistral/ministral3_14b_medpix.yaml @@ -90,7 +90,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/mistral/ministral3_3b_medpix.yaml b/examples/vlm_finetune/mistral/ministral3_3b_medpix.yaml index b3afa0480c..dc5bf1a5c6 100644 --- a/examples/vlm_finetune/mistral/ministral3_3b_medpix.yaml +++ b/examples/vlm_finetune/mistral/ministral3_3b_medpix.yaml @@ -90,7 +90,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/mistral/ministral3_8b_medpix.yaml b/examples/vlm_finetune/mistral/ministral3_8b_medpix.yaml index 0a189d6aa8..921bb4b53f 100644 --- a/examples/vlm_finetune/mistral/ministral3_8b_medpix.yaml +++ b/examples/vlm_finetune/mistral/ministral3_8b_medpix.yaml @@ -90,7 +90,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/mistral3p5/mistral3p5_128b_medpix.yaml b/examples/vlm_finetune/mistral3p5/mistral3p5_128b_medpix.yaml index 279c2e8d8e..9a8c27cee4 100644 --- a/examples/vlm_finetune/mistral3p5/mistral3p5_128b_medpix.yaml +++ b/examples/vlm_finetune/mistral3p5/mistral3p5_128b_medpix.yaml @@ -79,7 +79,6 @@ distributed: # Freeze vision + projector; train only the language_model (standard MedPix setup). freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/mistral3p5/mistral3p5_128b_medpix_lora.yaml b/examples/vlm_finetune/mistral3p5/mistral3p5_128b_medpix_lora.yaml index 351a6eac58..c3d3d80a75 100644 --- a/examples/vlm_finetune/mistral3p5/mistral3p5_128b_medpix_lora.yaml +++ b/examples/vlm_finetune/mistral3p5/mistral3p5_128b_medpix_lora.yaml @@ -77,7 +77,6 @@ distributed: scale_grads_in_schedule: false freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/mistral4/mistral4_medpix.yaml b/examples/vlm_finetune/mistral4/mistral4_medpix.yaml index 2176dd709c..4c8e4fe0f9 100644 --- a/examples/vlm_finetune/mistral4/mistral4_medpix.yaml +++ b/examples/vlm_finetune/mistral4/mistral4_medpix.yaml @@ -72,7 +72,6 @@ distributed: reshard_after_forward: false freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/nemotron/nemotron_parse_v1_1.yaml b/examples/vlm_finetune/nemotron/nemotron_parse_v1_1.yaml index d7f0ed5f20..0272abb36a 100644 --- a/examples/vlm_finetune/nemotron/nemotron_parse_v1_1.yaml +++ b/examples/vlm_finetune/nemotron/nemotron_parse_v1_1.yaml @@ -97,7 +97,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/nemotron_omni/nemotron_omni_cord_v2.yaml b/examples/vlm_finetune/nemotron_omni/nemotron_omni_cord_v2.yaml index 0f57694bfa..6a9eea4881 100644 --- a/examples/vlm_finetune/nemotron_omni/nemotron_omni_cord_v2.yaml +++ b/examples/vlm_finetune/nemotron_omni/nemotron_omni_cord_v2.yaml @@ -59,7 +59,6 @@ distributed: sequence_parallel: false freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/nemotron_omni/nemotron_omni_cord_v2_peft.yaml b/examples/vlm_finetune/nemotron_omni/nemotron_omni_cord_v2_peft.yaml index 425b46313c..de81cea595 100644 --- a/examples/vlm_finetune/nemotron_omni/nemotron_omni_cord_v2_peft.yaml +++ b/examples/vlm_finetune/nemotron_omni/nemotron_omni_cord_v2_peft.yaml @@ -72,7 +72,6 @@ distributed: sequence_parallel: false freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/nemotron_omni/nemotron_omni_v3_cord_v2_ep8cp2.yaml b/examples/vlm_finetune/nemotron_omni/nemotron_omni_v3_cord_v2_ep8cp2.yaml index 935b9500f7..b0826d73f0 100644 --- a/examples/vlm_finetune/nemotron_omni/nemotron_omni_v3_cord_v2_ep8cp2.yaml +++ b/examples/vlm_finetune/nemotron_omni/nemotron_omni_v3_cord_v2_ep8cp2.yaml @@ -54,7 +54,6 @@ distributed: sequence_parallel: false freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/phi4/phi4_mm_cv17.yaml b/examples/vlm_finetune/phi4/phi4_mm_cv17.yaml index 7254ac6b8f..29c8c75778 100644 --- a/examples/vlm_finetune/phi4/phi4_mm_cv17.yaml +++ b/examples/vlm_finetune/phi4/phi4_mm_cv17.yaml @@ -92,7 +92,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: false freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/qwen2_5/qwen2_5_vl_3b_rdr.yaml b/examples/vlm_finetune/qwen2_5/qwen2_5_vl_3b_rdr.yaml index d18ce5cbc6..3cc9bd2beb 100644 --- a/examples/vlm_finetune/qwen2_5/qwen2_5_vl_3b_rdr.yaml +++ b/examples/vlm_finetune/qwen2_5/qwen2_5_vl_3b_rdr.yaml @@ -84,7 +84,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/qwen2_5/qwen2_5_vl_3b_shopify.yaml b/examples/vlm_finetune/qwen2_5/qwen2_5_vl_3b_shopify.yaml index 502cc1708f..79a7af1878 100644 --- a/examples/vlm_finetune/qwen2_5/qwen2_5_vl_3b_shopify.yaml +++ b/examples/vlm_finetune/qwen2_5/qwen2_5_vl_3b_shopify.yaml @@ -102,7 +102,6 @@ lr_scheduler: lr_decay_style: cosine freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/qwen3/qwen3_omni_moe_30b_te_deepep.yaml b/examples/vlm_finetune/qwen3/qwen3_omni_moe_30b_te_deepep.yaml index 11486c4f9b..f61a5a53b1 100644 --- a/examples/vlm_finetune/qwen3/qwen3_omni_moe_30b_te_deepep.yaml +++ b/examples/vlm_finetune/qwen3/qwen3_omni_moe_30b_te_deepep.yaml @@ -67,7 +67,6 @@ distributed: sequence_parallel: false freeze_config: - freeze_embeddings: false freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/qwen3/qwen3_vl_4b_instruct_rdr.yaml b/examples/vlm_finetune/qwen3/qwen3_vl_4b_instruct_rdr.yaml index 14503fa30d..9ffa320550 100644 --- a/examples/vlm_finetune/qwen3/qwen3_vl_4b_instruct_rdr.yaml +++ b/examples/vlm_finetune/qwen3/qwen3_vl_4b_instruct_rdr.yaml @@ -84,7 +84,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/qwen3/qwen3_vl_4b_neat_packing.yaml b/examples/vlm_finetune/qwen3/qwen3_vl_4b_neat_packing.yaml index abe12c9fdf..5ac80c0639 100644 --- a/examples/vlm_finetune/qwen3/qwen3_vl_4b_neat_packing.yaml +++ b/examples/vlm_finetune/qwen3/qwen3_vl_4b_neat_packing.yaml @@ -97,7 +97,6 @@ validation_dataloader: _target_: nemo_automodel.components.datasets.vlm.collate_fns.default_collate_fn freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/qwen3/qwen3_vl_8b_cp2_vision_frame_shard.yaml b/examples/vlm_finetune/qwen3/qwen3_vl_8b_cp2_vision_frame_shard.yaml index f405b1f65f..4229586424 100644 --- a/examples/vlm_finetune/qwen3/qwen3_vl_8b_cp2_vision_frame_shard.yaml +++ b/examples/vlm_finetune/qwen3/qwen3_vl_8b_cp2_vision_frame_shard.yaml @@ -92,7 +92,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/qwen3/qwen3_vl_8b_instruct_rdr.yaml b/examples/vlm_finetune/qwen3/qwen3_vl_8b_instruct_rdr.yaml index 41a7d98c53..e4157835b7 100644 --- a/examples/vlm_finetune/qwen3/qwen3_vl_8b_instruct_rdr.yaml +++ b/examples/vlm_finetune/qwen3/qwen3_vl_8b_instruct_rdr.yaml @@ -84,7 +84,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/qwen3/qwen3_vl_moe_30b_neat_packing.yaml b/examples/vlm_finetune/qwen3/qwen3_vl_moe_30b_neat_packing.yaml index 3e91749092..64dc0b48ab 100644 --- a/examples/vlm_finetune/qwen3/qwen3_vl_moe_30b_neat_packing.yaml +++ b/examples/vlm_finetune/qwen3/qwen3_vl_moe_30b_neat_packing.yaml @@ -102,7 +102,6 @@ validation_dataloader: _target_: nemo_automodel.components.datasets.vlm.collate_fns.default_collate_fn freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/qwen3/qwen3_vl_moe_30b_te_deepep.yaml b/examples/vlm_finetune/qwen3/qwen3_vl_moe_30b_te_deepep.yaml index ca74a917eb..984bc4b667 100644 --- a/examples/vlm_finetune/qwen3/qwen3_vl_moe_30b_te_deepep.yaml +++ b/examples/vlm_finetune/qwen3/qwen3_vl_moe_30b_te_deepep.yaml @@ -100,7 +100,6 @@ optimizer: weight_decay: 0.0 freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/qwen3_5_moe/qwen3_5_122b_128k_ep8cp32.yaml b/examples/vlm_finetune/qwen3_5_moe/qwen3_5_122b_128k_ep8cp32.yaml index d47dff48b9..9c22899372 100644 --- a/examples/vlm_finetune/qwen3_5_moe/qwen3_5_122b_128k_ep8cp32.yaml +++ b/examples/vlm_finetune/qwen3_5_moe/qwen3_5_122b_128k_ep8cp32.yaml @@ -125,7 +125,6 @@ distributed: cost_alpha: auto freeze_config: - freeze_embeddings: false freeze_vision_tower: false freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/qwen3_5_moe/qwen3_6_35b_lora.yaml b/examples/vlm_finetune/qwen3_5_moe/qwen3_6_35b_lora.yaml index 10f4c2075a..16c81db809 100644 --- a/examples/vlm_finetune/qwen3_5_moe/qwen3_6_35b_lora.yaml +++ b/examples/vlm_finetune/qwen3_5_moe/qwen3_6_35b_lora.yaml @@ -82,7 +82,6 @@ distributed: sequence_parallel: false freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/stepfun/step3p7_medpix_200b_ep32pp4.yaml b/examples/vlm_finetune/stepfun/step3p7_medpix_200b_ep32pp4.yaml index f979814fe2..7d872ad40e 100644 --- a/examples/vlm_finetune/stepfun/step3p7_medpix_200b_ep32pp4.yaml +++ b/examples/vlm_finetune/stepfun/step3p7_medpix_200b_ep32pp4.yaml @@ -82,7 +82,6 @@ distributed: wrap_outer_model: true freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/examples/vlm_finetune/stepfun/step3p7_medpix_200b_lora_pp8ep8_8node.yaml b/examples/vlm_finetune/stepfun/step3p7_medpix_200b_lora_pp8ep8_8node.yaml index af8a2adaaa..40e25df880 100644 --- a/examples/vlm_finetune/stepfun/step3p7_medpix_200b_lora_pp8ep8_8node.yaml +++ b/examples/vlm_finetune/stepfun/step3p7_medpix_200b_lora_pp8ep8_8node.yaml @@ -91,7 +91,6 @@ distributed: wrap_outer_model: true freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/nemo_automodel/_transformers/infrastructure.py b/nemo_automodel/_transformers/infrastructure.py index e5f7af196f..920555963b 100644 --- a/nemo_automodel/_transformers/infrastructure.py +++ b/nemo_automodel/_transformers/infrastructure.py @@ -24,6 +24,7 @@ """ import logging +from collections.abc import Callable from contextlib import nullcontext from dataclasses import is_dataclass, replace from functools import partial @@ -67,6 +68,7 @@ from nemo_automodel.components.quantization.qat import QATConfig from nemo_automodel.components.utils.compile_utils import compile_model from nemo_automodel.components.utils.model_utils import ( + FreezeConfig, _supports_logits_to_keep, apply_parameter_freezing, count_model_parameters, @@ -75,6 +77,7 @@ freeze_minimax_m3_indexer_params, freeze_unused_kv_sharing_params, init_empty_weights, + parse_freeze_config, print_trainable_parameters, ) from nemo_automodel.shared.tied_weights import ensure_tied_lm_head @@ -189,7 +192,7 @@ def _apply_runtime_compatibility_fixes(model): # Sharding helpers -def _shard_pp(autopipeline, model, loss_fn, parallelize_fn): +def _shard_pp(autopipeline, model, loss_fn, parallelize_fn, reapply_trainability): trainable_params, total_params = count_model_parameters(model) # Store param info on autopipeline before splitting so it can be accessed later # This captures the full model's param counts before PP shards it across ranks @@ -198,22 +201,25 @@ def _shard_pp(autopipeline, model, loss_fn, parallelize_fn): if get_world_size_safe() == 1: logger.info("World size is 1, skipping autopipeline.") else: + if parallelize_fn is not None: + parallelize_fn = partial(parallelize_fn, reapply_trainability=reapply_trainability) autopipeline.build(model, loss_fn=loss_fn, parallelize_fn=parallelize_fn) model = autopipeline return model -def _shard_ep_fsdp(model, model_wrapper, parallelize_fn, mesh: MeshContext): +def _shard_ep_fsdp(model, model_wrapper, parallelize_fn, mesh: MeshContext, reapply_trainability): """Apply EP + FSDP sharding (non-PP path).""" if parallelize_fn is not None and get_world_size_safe() > 1: parallelize_fn( model, world_mesh=mesh.device_mesh, moe_mesh=mesh.moe_mesh, + reapply_trainability=reapply_trainability, **mesh.parallelize_axis_kwargs(), ) elif callable(getattr(model_wrapper, "parallelize", None)): - model = model_wrapper.parallelize(model) + model = model_wrapper.parallelize(model, reapply_trainability=reapply_trainability) model = ( model[0] if isinstance(model, tuple) else model ) # MegatronFSDP will return (model, None) since we don't pass optimizer here @@ -318,6 +324,7 @@ def parallelize_for_pp( model: torch.nn.Module, *, model_wrapper: Union[FSDP2Manager, MegatronFSDPManager, DDPManager] | None = None, + reapply_trainability: Callable[[torch.nn.Module], None] | None = None, **kwargs, ) -> torch.nn.Module: """Parallelize model for pipeline parallelism (non-MoE case). @@ -328,6 +335,8 @@ def parallelize_for_pp( Args: model: The model to parallelize. model_wrapper: Distributed manager instance. + reapply_trainability: Callback that re-resolves the trainability policy + after pipeline-stage surgery and immediately before wrapping. **kwargs: Additional arguments (world_mesh, moe_mesh, axis names) passed by AutoPipeline but unused for non-MoE parallelization. @@ -336,7 +345,7 @@ def parallelize_for_pp( """ if model_wrapper is not None: if callable(getattr(model_wrapper, "parallelize", None)): - model = model_wrapper.parallelize(model) + model = model_wrapper.parallelize(model, reapply_trainability=reapply_trainability) return model @@ -462,6 +471,40 @@ def _uses_thd_only_te_attention(model) -> bool: ) +def _apply_trainability_policy( + model: torch.nn.Module, + *, + peft_enabled: bool, + freeze_config: FreezeConfig | None, + strict: bool, +) -> None: + """Resolve the complete trainability policy on the current module hierarchy. + + Parallelization and checkpoint loading can replace modules and parameters. + Re-running this policy after each such surgery selects the current objects by + module path instead of transferring stale parameter-name state. + + Args: + model: Model or pipeline stage whose trainability is being resolved. + peft_enabled: Whether the PEFT baseline should freeze non-LoRA parameters. + freeze_config: Optional user freeze/unfreeze policy. + strict: Whether every generic selector must match this model. Full-model + validation is strict; pipeline stages use non-strict rebinding because + each rank owns only part of the hierarchy. + """ + if peft_enabled: + for name, param in model.named_parameters(remove_duplicate=False): + param.requires_grad_("lora_" in name) + if freeze_config is not None: + apply_parameter_freezing(model, freeze_config, strict=strict) + + # These are framework invariants, so they are always applied last and cannot + # be overridden by a user unfreeze selector. + freeze_unused_kv_sharing_params(model) + freeze_deepseek_v4_indexer_params(model) + freeze_minimax_m3_indexer_params(model) + + # apply_model_infrastructure -- the main post-init orchestration function def apply_model_infrastructure( model, @@ -622,16 +665,21 @@ def apply_model_infrastructure( _maybe_adapt_state_dict_to_hf(model, model.state_dict(), quantization=False).keys() ) - # Apply freezing before sharding - freeze_config = _kwargs.get("freeze_config") - if freeze_config is not None: - apply_parameter_freezing(model, freeze_config) - - # Freeze dead K/V parameters in KV-shared layers (e.g. Gemma4 E2B/E4B) - # so the optimizer never tracks them and checkpoint save/resume stay consistent. - freeze_unused_kv_sharing_params(model) - freeze_deepseek_v4_indexer_params(model) - freeze_minimax_m3_indexer_params(model) + # Validate selectors on the complete pre-parallelization hierarchy. The + # same policy is rebound after model surgery and before DDP/FSDP capture. + freeze_config = parse_freeze_config(_kwargs.get("freeze_config")) + _apply_trainability_policy( + model, + peft_enabled=peft_config is not None, + freeze_config=freeze_config, + strict=True, + ) + reapply_trainability = partial( + _apply_trainability_policy, + peft_enabled=peft_config is not None, + freeze_config=freeze_config, + strict=False, + ) # NemotronOmni RADIO: opt into the fused SDPA path on ViT attention blocks. enable_radio_vit_fused_attn(model) @@ -644,11 +692,11 @@ def apply_model_infrastructure( # Note: AutoPipeline takes care of applying PP + EP + FSDP. _shard_ep_fsdp will take care of applying EP + FSDP if no PP. mfsdp_param_attrs = None if autopipeline is not None: - model = _shard_pp(autopipeline, model, loss_fn, parallelize_fn) + model = _shard_pp(autopipeline, model, loss_fn, parallelize_fn, reapply_trainability) for part in model.parts: setattr(part, "_pre_shard_hf_state_dict_keys", pre_shard_hf_state_dict_keys) else: - model = _shard_ep_fsdp(model, model_wrapper, parallelize_fn, mesh) + model = _shard_ep_fsdp(model, model_wrapper, parallelize_fn, mesh, reapply_trainability) # Megatron-FSDP stamps load-bearing per-parameter state (owning-model back-ref, # tied-weight ``_is_shared`` marker, ``orig_param`` and friends) during wrapping. # The lm-head re-tie and post-wrap checkpoint reload below rebuild Parameter @@ -724,14 +772,23 @@ def apply_model_infrastructure( checkpoint_loaded=bool(checkpoint_already_loaded or weights_already_loaded or should_load_checkpoint), ) - # Freeze parameters after checkpoint loading and parallelization - # This catches params created during parallelization (e.g., GroupedExpertsTE in init_token_dispatcher) - if peft_config is not None: - models_to_freeze = model.parts if hasattr(model, "parts") else [model] - for mp in models_to_freeze: - for name, param in mp.named_parameters(): - if "lora_" not in name and param.requires_grad: - param.requires_grad_(False) + # Checkpoint loading can perform another round of parameter replacement, so + # re-resolve once more on the final model parts before optimizer construction. + trainability_models: list[torch.nn.Module] + if hasattr(model, "parts"): + trainability_models = list(model.parts) + elif isinstance(model_wrapper, (DDPManager, MegatronFSDPManager)): + trainability_models = [getattr(model, "module", model)] + else: + trainability_models = [model] + for mp in trainability_models: + reapply_trainability(mp) + if peft_config is not None or freeze_config is not None: + if not any(param.requires_grad for mp in trainability_models for param in mp.parameters()): + logger.warning( + "The configured trainability policy left no trainable parameters; " + "check freeze_config and the PEFT configuration." + ) if autopipeline is None: print_trainable_parameters(model) # Once model's been sharded @@ -809,8 +866,10 @@ def apply_model_infrastructure( # module also keeps its storage-dtype params. Under fp32 master weights + bf16 compute # that leaves frozen fp32 tensors feeding bf16 trainable modules -> dtype-mismatch # matmul at the seam. Cast frozen params/buffers to the compute dtype so the whole - # forward runs uniformly. No-op for pure-fp32 / pure-bf16 runs and when no mp_policy - # is available (DDP/PP). + # forward runs uniformly. Freeze configuration only controls requires_grad; trainable + # parameters keep their storage dtype (compute dtype is owned by autocast or the + # distributed mixed-precision policy). No-op for pure-fp32 / pure-bf16 runs and when + # no mp_policy is available (DDP/PP). compute_dtype = getattr(getattr(model_wrapper, "mp_policy", None), "param_dtype", None) if compute_dtype is not None: for mp in model.parts if hasattr(model, "parts") else [model]: diff --git a/nemo_automodel/components/distributed/ddp.py b/nemo_automodel/components/distributed/ddp.py index 8410880d6f..1a77252a84 100644 --- a/nemo_automodel/components/distributed/ddp.py +++ b/nemo_automodel/components/distributed/ddp.py @@ -13,6 +13,7 @@ # limitations under the License. import logging +from collections.abc import Callable import torch import torch.distributed as dist @@ -90,7 +91,11 @@ def _setup_distributed(self): else: self.device = torch.device("cpu") - def parallelize(self, model): + def parallelize( + self, + model: torch.nn.Module, + reapply_trainability: Callable[[torch.nn.Module], None] | None = None, + ) -> torch.nn.Module: """ Wraps the given model with DistributedDataParallel (DDP). @@ -99,6 +104,8 @@ def parallelize(self, model): Args: model (torch.nn.Module): The PyTorch model to be wrapped. + reapply_trainability: Optional callback that re-resolves parameter + trainability after model surgery and before DDP construction. Returns: torch.nn.parallel.DistributedDataParallel: The DDP-wrapped model. @@ -124,6 +131,8 @@ def parallelize(self, model): model.gradient_checkpointing_enable() else: apply_submodule_checkpointing(layers, detect_kv_sharing_and_maybe_disable_cache(model)) + if reapply_trainability is not None: + reapply_trainability(model) return model if self.activation_checkpointing: @@ -152,4 +161,7 @@ def parallelize(self, model): if self.bucket_cap_mb is not None: ddp_kwargs["bucket_cap_mb"] = self.bucket_cap_mb - return DDP(model.to(self.device), **ddp_kwargs) + model = model.to(self.device) + if reapply_trainability is not None: + reapply_trainability(model) + return DDP(model, **ddp_kwargs) diff --git a/nemo_automodel/components/distributed/fsdp2.py b/nemo_automodel/components/distributed/fsdp2.py index 00f40a4fbb..ad2f241653 100644 --- a/nemo_automodel/components/distributed/fsdp2.py +++ b/nemo_automodel/components/distributed/fsdp2.py @@ -13,7 +13,9 @@ # limitations under the License. import logging +from collections.abc import Callable +from torch import nn from torch.distributed.device_mesh import DeviceMesh from nemo_automodel.components.distributed.activation_checkpointing import ( @@ -131,12 +133,18 @@ def __init__( self.fsdp2_forward_prefetch_depth = config.fsdp2_forward_prefetch_depth self.frozen_multimodal_sharding = config.multimodal.frozen_sharding - def parallelize(self, model): + def parallelize( + self, + model: nn.Module, + reapply_trainability: Callable[[nn.Module], None] | None = None, + ) -> nn.Module: """ Parallelizes the given model using FSDP2 and TP sharding strategies. Args: model (nn.Module): The model to be parallelized. + reapply_trainability: Optional callback that re-resolves parameter + trainability after model surgery and before FSDP construction. Returns: The parallelized model. @@ -168,6 +176,8 @@ def parallelize(self, model): model.gradient_checkpointing_enable() else: apply_submodule_checkpointing(layers, detect_kv_sharing_and_maybe_disable_cache(model)) + if reapply_trainability is not None: + reapply_trainability(model) return model if self.config.patch_is_packed_sequence: @@ -189,6 +199,7 @@ def parallelize(self, model): reshard_after_forward=self.reshard_after_forward, activation_checkpointing_scope=self.activation_checkpointing_scope, frozen_multimodal_sharding=self.frozen_multimodal_sharding, + reapply_trainability=reapply_trainability, ) return model diff --git a/nemo_automodel/components/distributed/megatron_fsdp.py b/nemo_automodel/components/distributed/megatron_fsdp.py index 05acfeab6a..d5dc310572 100644 --- a/nemo_automodel/components/distributed/megatron_fsdp.py +++ b/nemo_automodel/components/distributed/megatron_fsdp.py @@ -13,6 +13,7 @@ # limitations under the License. import logging +from collections.abc import Callable from typing import TYPE_CHECKING import torch @@ -91,7 +92,12 @@ def __init__( self.fsdp_double_buffer = config.fsdp_double_buffer self.activation_checkpointing = config.activation_checkpointing - def parallelize(self, model, optimizer=None): + def parallelize( + self, + model: nn.Module, + optimizer: torch.optim.Optimizer | None = None, + reapply_trainability: Callable[[nn.Module], None] | None = None, + ) -> tuple[nn.Module, torch.optim.Optimizer | None]: """ Parallelizes the given model using MegatronFSDP and TP sharding strategies. @@ -101,6 +107,8 @@ def parallelize(self, model, optimizer=None): model.finish_grad_sync() before optimizer.step(), model.install_optimized_model_weights() and model.zero_grad_buffer() after optimizer.zero_grad(). + reapply_trainability: Optional callback that re-resolves parameter + trainability after model surgery and before FSDP construction. Returns: tuple: (parallelized_model, optimizer) @@ -113,6 +121,8 @@ def parallelize(self, model, optimizer=None): model.gradient_checkpointing_enable() else: logger.error("Model does not support gradient checkpointing. Skipping.") + if reapply_trainability is not None: + reapply_trainability(model) return model, optimizer if self.activation_checkpointing: @@ -170,6 +180,7 @@ def parallelize(self, model, optimizer=None): fsdp_double_buffer=self.fsdp_double_buffer, dp_shard_dim=dp_shard_dim, tp_dim=tp_dim, + reapply_trainability=reapply_trainability, ) return model, optimizer diff --git a/nemo_automodel/components/distributed/parallelizer.py b/nemo_automodel/components/distributed/parallelizer.py index 4db6da3b0d..b8b6111a72 100644 --- a/nemo_automodel/components/distributed/parallelizer.py +++ b/nemo_automodel/components/distributed/parallelizer.py @@ -17,6 +17,7 @@ import logging import warnings from abc import ABC, abstractmethod +from collections.abc import Callable from contextlib import contextmanager from functools import lru_cache from types import FunctionType @@ -263,6 +264,7 @@ def parallelize( reshard_after_forward: bool | None = None, activation_checkpointing_scope: ActivationCheckpointingScope | None = "all", frozen_multimodal_sharding: FrozenMultimodalSharding = "root", + reapply_trainability: Callable[[nn.Module], None] | None = None, **kwargs, ) -> nn.Module: """Apply parallelization strategy to the model.""" @@ -292,20 +294,11 @@ def parallelize( reshard_after_forward: bool | None = None, activation_checkpointing_scope: ActivationCheckpointingScope | None = "all", frozen_multimodal_sharding: FrozenMultimodalSharding = "root", + reapply_trainability: Callable[[nn.Module], None] | None = None, fully_shard_fn=None, ) -> nn.Module: """Apply the default parallelization flow.""" frozen_multimodal_sharding = normalize_frozen_multimodal_sharding(frozen_multimodal_sharding) - frozen_multimodal_modules = [ - name for name, module in iter_multimodal_modules(model) if module_is_fully_frozen(module) - ] - if frozen_multimodal_sharding == "per_layer" and frozen_multimodal_modules: - logger.warning( - "distributed.multimodal.frozen_sharding='per_layer' selected for %s. Every rank in the FSDP " - "group must execute or skip these modules the same number of times and in the same order on every " - "microbatch; rank-asymmetric modality execution can hang or desynchronize FSDP collectives.", - ", ".join(frozen_multimodal_modules), - ) tp_mesh = device_mesh[tp_mesh_name] if fully_shard_fn is None: fully_shard_fn = fully_shard @@ -421,6 +414,22 @@ def parallelize( else: apply_submodule_checkpointing(ac_layers, _has_kv_sharing) + if reapply_trainability is not None: + reapply_trainability(model) + + # Evaluate frozen-module ownership only after TP/AC transformations and + # trainability rebinding so FSDP sees the final module hierarchy. + frozen_multimodal_modules = [ + name for name, module in iter_multimodal_modules(model) if module_is_fully_frozen(module) + ] + if frozen_multimodal_sharding == "per_layer" and frozen_multimodal_modules: + logger.warning( + "distributed.multimodal.frozen_sharding='per_layer' selected for %s. Every rank in the FSDP " + "group must execute or skip these modules the same number of times and in the same order on every " + "microbatch; rank-asymmetric modality execution can hang or desynchronize FSDP collectives.", + ", ".join(frozen_multimodal_modules), + ) + # Set up mixed precision policy if not mp_policy: mp_policy = MixedPrecisionPolicy( @@ -508,6 +517,7 @@ def parallelize( dp_shard_cp_mesh_name: str = "dp_shard_cp", tp_mesh_name: str = "tp", reshard_after_forward: bool | None = None, + reapply_trainability: Callable[[nn.Module], None] | None = None, **kwargs, ) -> nn.Module: """Apply NemotronH-specific parallelization.""" @@ -589,6 +599,9 @@ def parallelize( # Refresh the local handle so the FSDP wrap below sees the wrapped blocks. _, layers = _nemotronh_decoder_blocks(model) + if reapply_trainability is not None: + reapply_trainability(model) + dp_mesh = get_fsdp_dp_mesh(device_mesh, dp_replicate_mesh_name, dp_shard_cp_mesh_name) fp32_compute_module_names = tuple(getattr(model, "_keep_in_fp32_modules_strict", None) or ()) @@ -778,6 +791,7 @@ def parallelize( dp_replicate_mesh_name: str = "dp_replicate", dp_shard_cp_mesh_name: str = "dp_shard_cp", tp_mesh_name: str = "tp", + reapply_trainability: Callable[[nn.Module], None] | None = None, **kwargs, ) -> nn.Module: # Not using custom tp_shard_plan; apply Wan-specific plan @@ -855,6 +869,9 @@ def parallelize( output_dtype=torch.float32, ) + if reapply_trainability is not None: + reapply_trainability(model) + # Apply FSDP sharding recursively and to root apply_fsdp2_sharding_recursively( model, @@ -890,6 +907,7 @@ def parallelize( dp_replicate_mesh_name: str = "dp_replicate", dp_shard_cp_mesh_name: str = "dp_shard_cp", tp_mesh_name: str = "tp", + reapply_trainability: Callable[[nn.Module], None] | None = None, **kwargs, ) -> nn.Module: dp_mesh = get_fsdp_dp_mesh(device_mesh, dp_replicate_mesh_name, dp_shard_cp_mesh_name) @@ -909,6 +927,9 @@ def parallelize( checkpoint_impl=CheckpointImpl.NO_REENTRANT, ) + if reapply_trainability is not None: + reapply_trainability(model) + # Apply FSDP sharding recursively and to root apply_fsdp2_sharding_recursively( model, @@ -2356,7 +2377,8 @@ def fsdp2_strategy_parallelize( reshard_after_forward: bool | None = None, activation_checkpointing_scope: ActivationCheckpointingScope | None = "all", frozen_multimodal_sharding: FrozenMultimodalSharding = "root", -): + reapply_trainability: Callable[[nn.Module], None] | None = None, +) -> nn.Module: """ Apply parallelisms and activation checkpointing to the model. @@ -2391,6 +2413,9 @@ def fsdp2_strategy_parallelize( owned by the root FSDP unit (``"root"``), sharded normally (``"per_layer"``), or excluded from FSDP and copied on every rank (``"replicate"``). + reapply_trainability: Optional callback that re-resolves parameter + trainability after strategy-specific model surgery and immediately + before FSDP construction. Returns: The parallelized model. @@ -2421,6 +2446,7 @@ def fsdp2_strategy_parallelize( reshard_after_forward=reshard_after_forward, activation_checkpointing_scope=activation_checkpointing_scope, frozen_multimodal_sharding=frozen_multimodal_sharding, + reapply_trainability=reapply_trainability, ) @@ -2551,6 +2577,7 @@ def megatron_fsdp_strategy_parallelize( fsdp_double_buffer: bool = False, dp_shard_dim: str = "dp", tp_dim: str = "tp", + reapply_trainability: Callable[[nn.Module], None] | None = None, ): """ Apply tensor/data parallelism (MegatronFSDP) and optional activation-checkpointing to the model. @@ -2606,6 +2633,9 @@ def megatron_fsdp_strategy_parallelize( Defaults to "dp". tp_dim (str): Key name for the tensor parallel mesh in device_mesh. Defaults to "tp". + reapply_trainability: Optional callback that re-resolves parameter + trainability after tensor-parallel surgery and immediately before + Megatron-FSDP construction. NOTE: The passed-in model should preferably reside on the meta device. Otherwise, ensure the model fits into available GPU or CPU memory. @@ -2631,6 +2661,9 @@ def megatron_fsdp_strategy_parallelize( if tp_mesh.size() > 1: parallelize_module(model, tp_mesh, tp_shard_plan) + if reapply_trainability is not None: + reapply_trainability(model) + # MegatronFSDP requires a sharded DP dimension to create its param/grad buffers. # In practice, configurations like world_size=2,tp=2 -> dp=1 frequently hit # DTensor metadata assertions inside megatron_fsdp. In that case, we still diff --git a/nemo_automodel/components/moe/parallelizer.py b/nemo_automodel/components/moe/parallelizer.py index e95ed49221..c6b4feaa50 100644 --- a/nemo_automodel/components/moe/parallelizer.py +++ b/nemo_automodel/components/moe/parallelizer.py @@ -997,8 +997,15 @@ def parallelize_model( sequence_parallel: bool = False, enable_async_tensor_parallel: bool = False, frozen_multimodal_sharding: FrozenMultimodalSharding = "root", + reapply_trainability: Callable[[nn.Module], None] | None = None, ) -> None: - """Apply tensor, context, expert, activation-checkpointing, and FSDP parallelism.""" + """Apply tensor, context, expert, activation-checkpointing, and FSDP parallelism. + + Args: + reapply_trainability: Optional callback that re-resolves parameter + trainability after TP/EP/AC surgery and immediately before FSDP + construction. + """ tp_enabled = tp_axis_name is not None and world_mesh[tp_axis_name].size() > 1 if tp_enabled: @@ -1065,6 +1072,9 @@ def parallelize_model( activation_checkpointing_scope=activation_checkpointing_scope, ) + if reapply_trainability is not None: + reapply_trainability(model) + if ep_shard_axis_names is not None: ep_shard_mesh = moe_mesh[ep_shard_axis_names] else: diff --git a/nemo_automodel/components/utils/model_utils.py b/nemo_automodel/components/utils/model_utils.py index 0102930d76..a06d741a07 100644 --- a/nemo_automodel/components/utils/model_utils.py +++ b/nemo_automodel/components/utils/model_utils.py @@ -12,14 +12,18 @@ # See the License for the specific language governing permissions and # limitations under the License. +import fnmatch import inspect import logging import os +from collections.abc import Mapping from contextlib import contextmanager +from dataclasses import dataclass from functools import lru_cache from typing import Any, Callable from nemo_automodel.shared.import_utils import safe_import +from nemo_automodel.shared.parameter_names import canonical_parameter_fqn HAVE_TORCHAO, torch_ao = safe_import("torchao") HAVE_BNB, bnb = safe_import("bitsandbytes") @@ -308,23 +312,208 @@ def print_trainable_parameters(model: nn.Module, name: str = "Model") -> tuple[i def _freeze_module_by_attribute_and_patterns(model, attribute_name, name_patterns): - """Helper function to freeze parameters by attribute name and name patterns. + """Freeze a legacy model attribute and modules matching name substrings.""" + if attribute_name is not None and hasattr(model, attribute_name): + getattr(model, attribute_name).requires_grad_(False) + + for name, module in model.named_modules(): + if any(pattern in name.lower() for pattern in name_patterns): + module.requires_grad_(False) + + +@dataclass(frozen=True) +class ModuleSelector: + """Select modules by exact canonical path or case-sensitive shell-style glob. + + Exactly one of ``path`` or ``glob`` must be set. Both match against the full + canonical module path (activation-checkpoint and ``torch.compile`` wrapper + components stripped): ``path`` is an exact match, while ``glob`` uses + :func:`fnmatch.fnmatchcase` semantics where ``*`` also crosses ``.`` + separators. Matching is recursive: a selected module's entire subtree is + frozen or unfrozen. + """ + + path: str | None = None + glob: str | None = None + + def __post_init__(self) -> None: + """Validate that exactly one matching mode is configured.""" + if (self.path is None) == (self.glob is None): + raise ValueError( + f"ModuleSelector requires exactly one of `path` or `glob`; got path={self.path!r}, glob={self.glob!r}." + ) + value = self.path if self.path is not None else self.glob + if not isinstance(value, str) or not value: + raise ValueError(f"ModuleSelector values must be non-empty strings; got {value!r}.") + + def matches(self, module_path: str) -> bool: + """Return whether this selector matches a canonical module path. + + Args: + module_path: Canonical fully qualified module path. + + Returns: + Whether the selector matches ``module_path``. + """ + if self.path is not None: + return module_path == self.path + assert self.glob is not None + return fnmatch.fnmatchcase(module_path, self.glob) + + def describe(self) -> str: + """Return the selector in its ``key: value`` configuration form.""" + key, value = ("path", self.path) if self.path is not None else ("glob", self.glob) + return f"{key}: {value}" + + +@dataclass(frozen=True) +class FreezeConfig: + """Typed schema for the ``freeze_config`` recipe section. + + Trainability semantics, in application order: + + 1. Full fine-tuning preserves the model's existing trainability state; + PEFT establishes a LoRA-trainable, base-frozen baseline. + 2. ``freeze_modules`` recursively freezes the selected modules. + 3. ``unfreeze_modules`` recursively unfreezes the selected modules and wins + on overlap. + 4. Framework-required freezes (dead K/V projections, indexer parameters) + are applied by the infrastructure afterwards and remain protected. + 5. The final trainable parameter set is validated before the optimizer is + constructed. + + The legacy modality booleans remain supported. ``freeze_vision_tower`` + defaults to ``True`` only for legacy-only configurations; a configuration + that declares ``freeze_modules`` or ``unfreeze_modules`` uses + explicit-selector semantics and does not implicitly freeze vision modules. + """ + + freeze_modules: list[ModuleSelector] | None = None + unfreeze_modules: list[ModuleSelector] | None = None + freeze_vision_tower: bool | None = None + freeze_audio_tower: bool = False + freeze_language_model: bool = False + freeze_video_embedder: bool = False + + def __post_init__(self) -> None: + """Validate selectors and legacy compatibility options.""" + for field_name in ("freeze_modules", "unfreeze_modules"): + selectors = getattr(self, field_name) + if selectors is None: + continue + if not isinstance(selectors, list): + raise TypeError(f"FreezeConfig.{field_name} must be a list; got {type(selectors).__name__}.") + invalid = [selector for selector in selectors if not isinstance(selector, ModuleSelector)] + if invalid: + raise TypeError(f"FreezeConfig.{field_name} entries must be ModuleSelector instances; got {invalid!r}.") + + if self.freeze_vision_tower is not None and not isinstance(self.freeze_vision_tower, bool): + raise TypeError( + f"FreezeConfig.freeze_vision_tower must be a boolean or None; got {self.freeze_vision_tower!r}." + ) + for field_name in ("freeze_audio_tower", "freeze_language_model", "freeze_video_embedder"): + value = getattr(self, field_name) + if not isinstance(value, bool): + raise TypeError(f"FreezeConfig.{field_name} must be a boolean; got {value!r}.") + + def has_generic_selectors(self) -> bool: + """Return whether either generic selector field was explicitly declared.""" + return self.freeze_modules is not None or self.unfreeze_modules is not None + + +_FREEZE_CONFIG_OPTIONS = { + "freeze_modules", + "unfreeze_modules", + "freeze_vision_tower", + "freeze_audio_tower", + "freeze_language_model", + "freeze_video_embedder", +} + + +def _parse_module_selector(entry: Any, *, field_name: str) -> ModuleSelector: + """Parse one ``freeze_modules``/``unfreeze_modules`` entry into a ModuleSelector. Args: - model: The model to apply freezing to. - attribute_name: Name of the model attribute to freeze (e.g., 'vision_tower'). - name_patterns: List of patterns to match in module names. + entry: Mapping with exactly one of ``path`` or ``glob``, or an existing + ModuleSelector. + field_name: Owning list field name, used in error messages. + + Returns: + The validated ModuleSelector. + + Raises: + ValueError: If the entry is not a ``{path: ...}`` / ``{glob: ...}`` + mapping or contains unknown keys. """ - # Freeze by attribute name - if hasattr(model, attribute_name): - for param in getattr(model, attribute_name).parameters(): - param.requires_grad = False + if isinstance(entry, ModuleSelector): + return entry + if not isinstance(entry, Mapping): + raise ValueError( + f"freeze_config.{field_name} entries must be mappings with exactly one of `path` or `glob`, " + f"e.g. `- path: vision_tower` or `- glob: '*.audio_encoder'`; got {entry!r}." + ) + unknown = set(entry) - {"path", "glob"} + if unknown: + raise ValueError( + f"freeze_config.{field_name} entry {dict(entry)!r} has unsupported key(s) {sorted(unknown)}; " + "use exactly one of `path` or `glob`." + ) + return ModuleSelector(path=entry.get("path"), glob=entry.get("glob")) - # Freeze by name patterns - for name, module in model.named_modules(): - if any(pattern in name.lower() for pattern in name_patterns): - for param in module.parameters(): - param.requires_grad = False + +def parse_freeze_config(config: FreezeConfig | Mapping[str, Any] | None) -> FreezeConfig | None: + """Validate raw ``freeze_config`` YAML data and return a typed FreezeConfig. + + Args: + config: A FreezeConfig (returned unchanged), a raw mapping from the + recipe config, or None. + + Returns: + The validated FreezeConfig, or None when no freeze configuration was + provided. + + Raises: + TypeError: If ``config`` is neither a mapping nor a FreezeConfig. + ValueError: If ``config`` contains unknown options, malformed + selectors, or non-boolean legacy options. + """ + if config is None or isinstance(config, FreezeConfig): + return config + if not isinstance(config, Mapping): + raise TypeError(f"freeze_config must be a mapping of freeze options; got {type(config).__name__}: {config!r}.") + + unknown = set(config) - _FREEZE_CONFIG_OPTIONS + if unknown: + raise ValueError( + f"freeze_config has unsupported option(s) {sorted(unknown)}; " + f"valid options are {sorted(_FREEZE_CONFIG_OPTIONS)}." + ) + + def _selectors(field_name: str) -> list[ModuleSelector] | None: + if field_name not in config: + return None + entries = config[field_name] + if not isinstance(entries, list): + raise ValueError( + f"freeze_config.{field_name} must be a list of `path`/`glob` selector mappings; got {entries!r}." + ) + return [_parse_module_selector(entry, field_name=field_name) for entry in entries] + + legacy_booleans = {} + for option in ("freeze_vision_tower", "freeze_audio_tower", "freeze_language_model", "freeze_video_embedder"): + value = config.get(option, None) + if value is None: + continue + if not isinstance(value, bool): + raise ValueError(f"freeze_config.{option} must be a boolean; got {value!r}.") + legacy_booleans[option] = value + + return FreezeConfig( + freeze_modules=_selectors("freeze_modules"), + unfreeze_modules=_selectors("unfreeze_modules"), + **legacy_booleans, + ) def enable_radio_vit_fused_attn(model): @@ -361,50 +550,113 @@ def enable_radio_vit_fused_attn(model): logger.info("Enabled fused_attn on %d RADIO ViT blocks", flipped) -def apply_parameter_freezing(model, freeze_config): - """Apply parameter freezing based on configuration. +def _apply_module_selectors( + model: nn.Module, + selectors: list[ModuleSelector], + *, + requires_grad: bool, + strict: bool, + field_name: str, +) -> None: + """Set ``requires_grad`` on every module matched by ``selectors``, recursively. + + Module paths are canonicalized (activation-checkpoint and ``_orig_mod`` + wrapper components stripped) so selectors keep matching after model surgery + that wraps or renames parameter-holding modules. Args: - model: The model to apply freezing to. - freeze_config: Configuration dict specifying what to freeze. + model: The model (or pipeline-parallel model part) to modify. + selectors: Typed selectors to resolve against the module hierarchy. + requires_grad: Trainability to apply to matched modules. + strict: When True, raise if a selector matches no parameters. Rebinding + after parallelization uses False because sharding may relocate or + regroup the selected modules (e.g. pipeline stages hold only a part + of the model). + field_name: Owning FreezeConfig field name, used in error messages. + + Raises: + ValueError: If ``strict`` and a selector matches no parameters. + """ + named_modules = [ + (canonical_parameter_fqn(name).replace("_orig_mod.", ""), module) + for name, module in model.named_modules(remove_duplicate=False) + if name + ] + for selector in selectors: + matched_params = 0 + for name, module in named_modules: + if selector.matches(name): + matched_params += sum(1 for _ in module.parameters()) + module.requires_grad_(requires_grad) + if strict and matched_params == 0: + raise ValueError( + f"freeze_config.{field_name} selector `{selector.describe()}` matched no parameters in the model. " + "Selectors must match the full canonical module path." + ) + + +def apply_parameter_freezing( + model: nn.Module, + freeze_config: FreezeConfig | Mapping[str, Any], + *, + strict: bool = True, +) -> None: + """Apply parameter freezing based on a typed FreezeConfig. + + Application order: legacy modality booleans and ``freeze_modules`` freeze, + then ``unfreeze_modules`` unfreezes and wins on overlap. Framework-required + freezes (dead K/V projections, indexer parameters) are applied separately + by the infrastructure after this function, before and after sharding. - freeze_config can contain: - - freeze_vision_tower: bool (default True) + Args: + model: The model to apply freezing to. + freeze_config: Typed freeze configuration or a raw mapping retained for + compatibility with direct callers. Raw mappings are validated and + converted to FreezeConfig before use. + strict: When True, raise if a ``freeze_modules``/``unfreeze_modules`` + selector matches no parameters. Set False when rebinding the policy + onto a post-parallelization model part. + + Legacy modality booleans: + - freeze_vision_tower: bool (default True in legacy-only configurations) - freeze_audio_tower: bool (default False) - freeze_language_model: bool (default False) - freeze_video_embedder: bool (default False) """ - freeze_vision_tower = freeze_config.get("freeze_vision_tower", True) - freeze_audio_tower = freeze_config.get("freeze_audio_tower", False) - freeze_language_model = freeze_config.get("freeze_language_model", False) - freeze_video_embedder = freeze_config.get("freeze_video_embedder", False) - - # Freeze vision tower - if freeze_vision_tower: - _freeze_module_by_attribute_and_patterns(model, "vision_tower", ["vision", "visual", "image_encoder"]) - - # Freeze audio tower - if freeze_audio_tower: - _freeze_module_by_attribute_and_patterns(model, "audio_tower", ["audio", "audio_encoder", "speech", "sound"]) - - # Freeze language model backbone - if freeze_language_model: - _freeze_module_by_attribute_and_patterns(model, "language_model", ["language", "text", "llm"]) - - # NemotronOmni RADIO: patch_generator.video_embedder is only exercised on - # video inputs; on image-only training it sits in the optimizer without - # state (no grad → no lazy init), so dcp.load on resume raises + parsed_freeze_config = parse_freeze_config(freeze_config) + if parsed_freeze_config is None: + return + freeze_config = parsed_freeze_config + + # Preserve the legacy attribute and substring matching independently of the + # fnmatch semantics used by the generic selectors. + freeze_vision_tower = freeze_config.freeze_vision_tower + if freeze_vision_tower is None: + # Explicit-selector configurations establish trainability solely through + # freeze_modules/unfreeze_modules and skip the legacy implicit vision freeze. + freeze_vision_tower = not freeze_config.has_generic_selectors() + # freeze_video_embedder: NemotronOmni RADIO's patch_generator.video_embedder is only + # exercised on video inputs; on image-only training it sits in the optimizer without + # state (no grad -> no lazy init), so dcp.load on resume raises # "Missing key in checkpoint state_dict: optim.state.<...>.video_embedder.weight.step". # Independent of freeze_vision_tower so the image encoder can stay trainable # while the video branch is frozen out. - if freeze_video_embedder: - frozen = 0 - for name, param in model.named_parameters(): - if "patch_generator.video_embedder" in name: - param.requires_grad_(False) - frozen += 1 - if frozen: - logger.info("Froze %d patch_generator.video_embedder params", frozen) + legacy_freeze_aliases = ( + (freeze_vision_tower, "vision_tower", ("vision", "visual", "image_encoder")), + (freeze_config.freeze_audio_tower, "audio_tower", ("audio", "audio_encoder", "speech", "sound")), + (freeze_config.freeze_language_model, "language_model", ("language", "text", "llm")), + (freeze_config.freeze_video_embedder, None, ("patch_generator.video_embedder",)), + ) + for enabled, attribute_name, name_patterns in legacy_freeze_aliases: + if enabled: + _freeze_module_by_attribute_and_patterns(model, attribute_name, name_patterns) + + _apply_module_selectors( + model, freeze_config.freeze_modules or [], requires_grad=False, strict=strict, field_name="freeze_modules" + ) + _apply_module_selectors( + model, freeze_config.unfreeze_modules or [], requires_grad=True, strict=strict, field_name="unfreeze_modules" + ) # Phi4MM: cast internal fp32 LoRA adapters to bf16 for FSDP2 compatibility, # and disable KV cache (remote code uses legacy DynamicCache.key_cache diff --git a/nemo_automodel/recipes/base_recipe.py b/nemo_automodel/recipes/base_recipe.py index 69520f9e2e..77b475ad14 100644 --- a/nemo_automodel/recipes/base_recipe.py +++ b/nemo_automodel/recipes/base_recipe.py @@ -16,6 +16,7 @@ import logging import os import socket +from contextlib import nullcontext from datetime import datetime from pathlib import Path @@ -222,6 +223,13 @@ def untrack_state(self, *keys: str) -> None: for key in keys: tracked.discard(key) + def _autocast_context(self): + """Return the recipe-level autocast context configured by the strategy.""" + autocast_dtype = getattr(getattr(self, "distributed_config", None), "autocast_dtype", None) + if autocast_dtype is None: + return nullcontext() + return torch.autocast(device_type=self.dist_env.device.type, dtype=autocast_dtype) + def save_checkpoint( self, epoch: int, diff --git a/nemo_automodel/recipes/llm/train_ft.py b/nemo_automodel/recipes/llm/train_ft.py index 268cd5e80e..ad5118794e 100644 --- a/nemo_automodel/recipes/llm/train_ft.py +++ b/nemo_automodel/recipes/llm/train_ft.py @@ -54,6 +54,7 @@ ) from nemo_automodel._transformers.utils import apply_cache_compatibility_patches from nemo_automodel.components.config._arg_parser import parse_args_and_load_config +from nemo_automodel.components.config.loader import ConfigNode from nemo_automodel.components.cuda_graphs import PartialCudaGraphManager from nemo_automodel.components.datasets.loader import DataloaderConfig from nemo_automodel.components.distributed.config import DistributedSetup, FSDP2Config, MegatronFSDPConfig @@ -90,6 +91,7 @@ ) from nemo_automodel.components.utils.flops_utils import calculate_mfu from nemo_automodel.components.utils.model_utils import ( + FreezeConfig, _supports_logits_to_keep, _supports_seq_lens, filter_forward_kwargs, @@ -182,7 +184,7 @@ def build_model( cfg_quantization=None, distributed_setup: DistributedSetup | None = None, cfg_qat=None, - unfreeze_modules: list[str] | None = None, + cfg_freeze: ConfigNode | dict[str, Any] | FreezeConfig | None = None, sdpa_method: list[str] | None = None, device_mesh=None, ) -> tuple[nn.Module | AutoPipeline, list["Optimizer"]]: # noqa: F821 @@ -198,7 +200,8 @@ def build_model( cfg_quantization: Configuration for BitsAndBytes quantization. distributed_setup: Resolved distributed topology and policy object. cfg_qat: Configuration for QAT (will be instantiated to QATConfig). - unfreeze_modules: List of module names/substrings to unfreeze. + cfg_freeze: Freeze configuration (``freeze_config`` YAML section as a + mapping, or a typed FreezeConfig) controlling parameter trainability. sdpa_method: Explicit list of SDPA backend name strings (e.g. ``["flash_attention", "efficient_attention"]``), or ``None`` to auto-select based on CP / activation checkpointing. @@ -208,6 +211,9 @@ def build_model( kwargs = { "has_packed_sequence": has_packed_sequence, "peft_config": cfg_peft, + # ConfigNode lives at the recipe boundary; downstream components take + # plain mappings or typed FreezeConfig objects. + "freeze_config": cfg_freeze.to_dict() if isinstance(cfg_freeze, ConfigNode) else cfg_freeze, "sdpa_method": sdpa_method, } if distributed_setup is not None: @@ -289,18 +295,12 @@ def build_model( fp8_config=kwargs.get("fp8_config"), compile_config=kwargs.get("compile_config"), quantization_config=kwargs.get("quantization_config"), + freeze_config=kwargs.get("freeze_config"), pretrained_model_name_or_path=None, load_base_model=False, cache_dir=hf_constants.HF_HUB_CACHE, ) - # Explicitly unfreeze specified modules (e.g. task heads) that need full fine-tuning - if unfreeze_modules: - for name, param in model.named_parameters(): - if any(module_name in name for module_name in unfreeze_modules): - param.requires_grad_(True) - logging.info(f"Unfroze parameters matching: {unfreeze_modules}") - return model @@ -620,6 +620,7 @@ def setup(self): cfg_quantization=self.cfg.get("quantization", None), distributed_setup=self.distributed_setup, cfg_qat=self.cfg.get("qat", None), + cfg_freeze=self.cfg.get("freeze_config", None), sdpa_method=self.cfg.get("sdpa_method", None), ) self.embedding_row_repair_report = None diff --git a/nemo_automodel/recipes/llm/train_seq_cls.py b/nemo_automodel/recipes/llm/train_seq_cls.py index cf9c328a7c..92e08d45f0 100644 --- a/nemo_automodel/recipes/llm/train_seq_cls.py +++ b/nemo_automodel/recipes/llm/train_seq_cls.py @@ -32,7 +32,7 @@ from nemo_automodel.components.training.rng import ScopedRNG, StatefulRNG from nemo_automodel.components.training.utils import clip_grad_norm from nemo_automodel.components.utils.flops_utils import calculate_mfu -from nemo_automodel.components.utils.model_utils import filter_forward_kwargs +from nemo_automodel.components.utils.model_utils import FreezeConfig, ModuleSelector, filter_forward_kwargs from nemo_automodel.recipes._dist_utils import create_distributed_setup_from_config, shard_optimizers_for_megatron_fsdp from nemo_automodel.recipes._typed_config import RecipeConfig from nemo_automodel.recipes.base_recipe import BaseRecipe @@ -105,6 +105,11 @@ def setup(self): ) self.peft_config = self.cfg.instantiate_path("peft") + freeze_config = self.cfg.get("freeze_config", None) + if freeze_config is None and self.peft_config is not None: + # Preserve the pre-freeze_config behavior for existing PEFT sequence + # classification recipes; new configs declare this selector directly. + freeze_config = FreezeConfig(unfreeze_modules=[ModuleSelector(glob="*classifier")]) # fp32 master-weight default planned to be enabled in follow-up PR (resolve_storage_dtype). model = build_model( cfg_model=self.cfg.model, @@ -114,7 +119,7 @@ def setup(self): cfg_compile=self.cfg.get("compile", None), cfg_quantization=self.cfg.get("quantization", None), distributed_setup=self.distributed_setup, - unfreeze_modules=["classifier"] if self.peft_config is not None else None, + cfg_freeze=freeze_config, ) optimizer = self.cfg.optimizer.build(model, device_mesh=self.device_mesh, is_peft=self.peft_config is not None) allow_megatron_fsdp_sharding = getattr(self.cfg.optimizer, "supports_megatron_fsdp_sharding", True) @@ -230,9 +235,10 @@ def _run_train_optim_step(self, batches): } labels = batch.pop("labels") batch = filter_forward_kwargs(model, batch) - out = model(**batch) - logits = getattr(out, "logits", out) - loss = self.loss_fn(logits, labels.view(-1)) + with self._autocast_context(): + out = model(**batch) + logits = getattr(out, "logits", out) + loss = self.loss_fn(logits, labels.view(-1)) losses.append(loss.detach().clone()) # Collect predictions for accuracy calculation @@ -339,9 +345,10 @@ def _validate_one_epoch(self, dataloader): } labels = batch.pop("labels") batch = filter_forward_kwargs(model, batch) - out = model(**batch) - logits = getattr(out, "logits", out) - loss = self.loss_fn(logits, labels.view(-1)) + with self._autocast_context(): + out = model(**batch) + logits = getattr(out, "logits", out) + loss = self.loss_fn(logits, labels.view(-1)) total_loss += loss.detach() # Collect predictions for accuracy diff --git a/tests/functional_tests/hf_transformer/test_freeze_config_fsdp2.py b/tests/functional_tests/hf_transformer/test_freeze_config_fsdp2.py new file mode 100644 index 0000000000..54c6b134d1 --- /dev/null +++ b/tests/functional_tests/hf_transformer/test_freeze_config_fsdp2.py @@ -0,0 +1,195 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Real two-rank FSDP2 lifecycle coverage for generic freeze configuration.""" + +import copy +import socket +from types import SimpleNamespace + +import pytest +import torch +import torch.distributed as dist +import torch.multiprocessing as mp +from torch import nn +from torch.distributed.device_mesh import init_device_mesh +from torch.distributed.fsdp import MixedPrecisionPolicy +from torch.distributed.tensor import DTensor + +from nemo_automodel._transformers.infrastructure import apply_model_infrastructure +from nemo_automodel.components.distributed.config import FSDP2Config +from nemo_automodel.components.distributed.fsdp2 import FSDP2Manager +from nemo_automodel.components.distributed.mesh import MeshContext + +_WORLD_SIZE = 2 +_FEATURES = 4 + + +class _TinyBackbone(nn.Module): + """One-layer backbone exposing AutoModel's generic layer-container contract.""" + + def __init__(self) -> None: + super().__init__() + self.layers = nn.ModuleList([nn.Linear(_FEATURES, _FEATURES, bias=False)]) + + def forward(self, inputs: torch.Tensor) -> torch.Tensor: + """Run the backbone layer. + + Args: + inputs: Tensor of shape [batch, features]. + + Returns: + Tensor of shape [batch, features]. + """ + return self.layers[0](inputs) + + +class _TinyFreezeModel(nn.Module): + """Small model with frozen model-owned state and an explicitly selected head.""" + + def __init__(self) -> None: + super().__init__() + self.backbone = _TinyBackbone() + self.classifier = nn.Linear(_FEATURES, 1, bias=False) + self.classifier.requires_grad_(False) + self.model_constant = nn.Parameter(torch.ones(1), requires_grad=False) + self.config = SimpleNamespace(use_cache=False, num_kv_shared_layers=0) + + def forward(self, inputs: torch.Tensor) -> torch.Tensor: + """Run the frozen backbone and selected classifier. + + Args: + inputs: Tensor of shape [batch, features]. + + Returns: + Tensor of shape [batch, 1]. + """ + return self.classifier(self.backbone(inputs)) + + +def _free_port() -> int: + """Return an available localhost TCP port for the spawned process group.""" + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock: + sock.bind(("127.0.0.1", 0)) + return int(sock.getsockname()[1]) + + +def _full_tensor(tensor: torch.Tensor) -> torch.Tensor: + """Replicate a possibly sharded tensor on every rank. + + Args: + tensor: Tensor of arbitrary shape, either local or sharded on the FSDP mesh. + + Returns: + Tensor with the input's global shape, replicated on every rank. + """ + return tensor.full_tensor() if isinstance(tensor, DTensor) else tensor + + +def _worker(rank: int, port: int) -> None: + """Run one rank of the real FSDP2 trainability lifecycle regression.""" + torch.cuda.set_device(rank) + dist.init_process_group( + "nccl", + init_method=f"tcp://127.0.0.1:{port}", + rank=rank, + world_size=_WORLD_SIZE, + ) + try: + torch.manual_seed(1234) + model = _TinyFreezeModel().cuda(rank) + reference = copy.deepcopy(model) + reference.backbone.requires_grad_(False) + reference.classifier.requires_grad_(True) + + mesh = init_device_mesh( + "cuda", + (1, _WORLD_SIZE, 1), + mesh_dim_names=("dp_replicate", "dp_shard_cp", "tp"), + ) + config = FSDP2Config( + mp_policy=MixedPrecisionPolicy( + param_dtype=torch.float32, + reduce_dtype=torch.float32, + output_dtype=torch.float32, + ), + enable_fsdp2_prefetch=False, + ) + model = apply_model_infrastructure( + model=model, + is_meta_device=False, + device=torch.device("cuda", rank), + load_base_model=False, + model_wrapper=FSDP2Manager(config, device_mesh=mesh), + mesh=MeshContext.from_meshes(mesh), + freeze_config={ + "freeze_modules": [{"path": "backbone"}], + "unfreeze_modules": [{"path": "classifier"}], + }, + ) + + assert isinstance(model.backbone.layers[0].weight, DTensor) + assert isinstance(model.classifier.weight, DTensor) + assert isinstance(model.model_constant, DTensor) + assert not model.backbone.layers[0].weight.requires_grad + assert model.classifier.weight.requires_grad + assert not model.model_constant.requires_grad + assert [name for name, parameter in model.named_parameters() if parameter.requires_grad] == [ + "classifier.weight" + ] + + optimizer = torch.optim.SGD((parameter for parameter in model.parameters() if parameter.requires_grad), lr=0.1) + reference_optimizer = torch.optim.SGD( + (parameter for parameter in reference.parameters() if parameter.requires_grad), lr=0.1 + ) + + rank_inputs = torch.full((2, _FEATURES), float(rank + 1), device=rank) + model(rank_inputs).sum().backward() + + reference_loss = ( + sum( + reference(torch.full((2, _FEATURES), float(source_rank + 1), device=rank)).sum() + for source_rank in range(_WORLD_SIZE) + ) + / _WORLD_SIZE + ) + reference_loss.backward() + + assert model.backbone.layers[0].weight.grad is None + assert model.classifier.weight.grad is not None + torch.testing.assert_close( + _full_tensor(model.classifier.weight.grad), + reference.classifier.weight.grad, + ) + + optimizer.step() + reference_optimizer.step() + full_classifier_weight = _full_tensor(model.classifier.weight) + torch.testing.assert_close(full_classifier_weight, reference.classifier.weight) + + gathered_weights = [torch.empty_like(full_classifier_weight) for _ in range(_WORLD_SIZE)] + dist.all_gather(gathered_weights, full_classifier_weight) + for gathered_weight in gathered_weights: + torch.testing.assert_close(gathered_weight, full_classifier_weight) + finally: + dist.destroy_process_group() + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < _WORLD_SIZE, + reason="requires two CUDA GPUs", +) +def test_freeze_config_survives_real_fsdp2_forward_backward_and_optimizer_step() -> None: + """Selected FSDP2 parameters remain trainable and synchronized through one update.""" + mp.spawn(_worker, args=(_free_port(),), nprocs=_WORLD_SIZE, join=True) diff --git a/tests/functional_tests/llm_pretrain_and_kd/llm_seq_cls/seq_cls_bert_lora_fp32.yaml b/tests/functional_tests/llm_pretrain_and_kd/llm_seq_cls/seq_cls_bert_lora_fp32.yaml index 29c1ac3868..e81c37dec7 100644 --- a/tests/functional_tests/llm_pretrain_and_kd/llm_seq_cls/seq_cls_bert_lora_fp32.yaml +++ b/tests/functional_tests/llm_pretrain_and_kd/llm_seq_cls/seq_cls_bert_lora_fp32.yaml @@ -32,6 +32,7 @@ distributed: tp_size: 1 cp_size: 1 sequence_parallel: false + autocast_dtype: bfloat16 peft: _target_: nemo_automodel.components._peft.lora.PeftConfig @@ -41,6 +42,10 @@ peft: dim: 4 alpha: 8 +freeze_config: + unfreeze_modules: + - glob: "*classifier" + dataset: _target_: nemo_automodel/components/datasets/llm/mock_seq_cls.py:MockSequenceClassificationDataset num_samples: 4 diff --git a/tests/functional_tests/parallelism/gemma4_31b_proxy.yaml b/tests/functional_tests/parallelism/gemma4_31b_proxy.yaml index 9742f6331b..dd8ca5e56c 100644 --- a/tests/functional_tests/parallelism/gemma4_31b_proxy.yaml +++ b/tests/functional_tests/parallelism/gemma4_31b_proxy.yaml @@ -121,7 +121,6 @@ optimizer: betas: [0.9, 0.95] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_audio_tower: true freeze_language_model: false diff --git a/tests/unit_tests/_transformers/test_infrastructure.py b/tests/unit_tests/_transformers/test_infrastructure.py index 8aec00068e..481983314c 100644 --- a/tests/unit_tests/_transformers/test_infrastructure.py +++ b/tests/unit_tests/_transformers/test_infrastructure.py @@ -12,6 +12,9 @@ # See the License for the specific language governing permissions and # limitations under the License. +import os +from datetime import timedelta +from pathlib import Path from types import SimpleNamespace from unittest.mock import MagicMock, call, patch @@ -138,6 +141,19 @@ def test_moe_infrastructure_forwards_fsdp2_tp_sequence_and_offload_settings(): assert parallelize_fn.keywords["frozen_multimodal_sharding"] == "replicate" +def test_pipeline_parallelizer_forwards_trainability_rebind(): + """Each PP stage receives the same pre-wrapper trainability policy.""" + from nemo_automodel._transformers.infrastructure import parallelize_for_pp + + model = torch.nn.Linear(2, 2) + callback = MagicMock() + manager = MagicMock() + manager.parallelize.return_value = model + + assert parallelize_for_pp(model, model_wrapper=manager, reapply_trainability=callback) is model + manager.parallelize.assert_called_once_with(model, reapply_trainability=callback) + + # ============================================================================= # Tests for apply_model_infrastructure: post-shard initialize_model_weights # ============================================================================= @@ -150,6 +166,245 @@ def __init__(self): self.config = SimpleNamespace() +class _TinyTrainabilityModel(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.base = torch.nn.Linear(1, 1, bias=False) + self.extension = torch.nn.Linear(1, 1, bias=False) + self.extension.requires_grad_(False) + self.model_constant = torch.nn.Parameter(torch.ones(1), requires_grad=False) + self.config = SimpleNamespace() + + def forward(self, inputs: torch.Tensor) -> torch.Tensor: + """Apply both projections. + + Args: + inputs: Tensor of shape [batch, features]. + + Returns: + Tensor of shape [batch, features]. + """ + return self.base(inputs) + self.extension(inputs) + + +class _TinyClassifierModel(torch.nn.Module): + def __init__(self) -> None: + super().__init__() + self.backbone = torch.nn.Linear(2, 2) + self.classifier = torch.nn.Linear(2, 1) + self.config = SimpleNamespace() + + def forward(self, inputs: torch.Tensor) -> torch.Tensor: + """Apply the backbone and classifier. + + Args: + inputs: Tensor of shape [batch, features]. + + Returns: + Tensor of shape [batch, 1]. + """ + return self.classifier(self.backbone(inputs)) + + +_GENERIC_FREEZE_CONFIG = { + "freeze_modules": [{"glob": "b*"}, {"glob": "ext*"}], + "unfreeze_modules": [{"glob": "ext*"}], +} + + +def _run_freeze_config_ddp(rank: int, world_size: int, init_file: str, result_dir: str) -> None: + os.environ["GLOO_SOCKET_IFNAME"] = "lo" + torch.distributed.init_process_group( + "gloo", + init_method=f"file://{init_file}", + rank=rank, + world_size=world_size, + timeout=timedelta(seconds=30), + ) + try: + from nemo_automodel._transformers.infrastructure import apply_model_infrastructure + from nemo_automodel.components.distributed.config import DDPConfig + from nemo_automodel.components.distributed.ddp import DDPManager + + model = apply_model_infrastructure( + model=_TinyTrainabilityModel(), + is_meta_device=False, + device=torch.device("cpu"), + load_base_model=False, + model_wrapper=DDPManager(DDPConfig()), + freeze_config=_GENERIC_FREEZE_CONFIG, + ) + + for _ in range(2): + model.zero_grad(set_to_none=True) + model(torch.tensor([[float(rank + 1)]])).sum().backward() + + assert model.module.base.weight.grad is None + assert model.module.extension.weight.grad is not None + Path(result_dir, f"rank_{rank}.txt").write_text( + str(model.module.extension.weight.grad.item()), + encoding="utf-8", + ) + finally: + torch.distributed.destroy_process_group() + + +def test_freeze_config_applies_before_ddp_reducer_construction(tmp_path: Path): + """Configured trainable parameters participate in DDP gradient synchronization.""" + torch.multiprocessing.spawn( + _run_freeze_config_ddp, + args=(2, str(tmp_path / "gloo_init"), str(tmp_path)), + nprocs=2, + join=True, + ) + + gradients = [float((tmp_path / f"rank_{rank}.txt").read_text(encoding="utf-8")) for rank in range(2)] + assert gradients == [1.5, 1.5] + + +def _replace_and_wrap_parameters(model, *_args, **_kwargs): + """Simulate parallelization surgery: replace Parameters and change their FQNs. + + Wrapping ``extension`` in a Sequential changes its parameter FQN from + ``extension.weight`` to ``extension.0.weight`` and installs a new Parameter + object (which defaults to requires_grad=True), mirroring grouped-expert + regrouping and wrapper-injecting transforms. + """ + model.base.weight = torch.nn.Parameter(model.base.weight.detach().clone()) + model.extension = torch.nn.Sequential(model.extension) + model.extension[0].weight = torch.nn.Parameter(model.extension[0].weight.detach().clone()) + if hasattr(model, "lora_weight"): + model.lora_weight = torch.nn.Parameter(model.lora_weight.detach().clone()) + model.vision_adapter.lora_weight = torch.nn.Parameter(model.vision_adapter.lora_weight.detach().clone()) + model.register_parameter("parallel_parameter", torch.nn.Parameter(torch.ones(1))) + return model + + +def _run_trainability_infrastructure(model, freeze_config, peft_config=None): + from nemo_automodel._transformers.infrastructure import apply_model_infrastructure + + class _ReplacingWrapper: + mp_policy = None + + def parallelize(self, model, reapply_trainability=None): + model = _replace_and_wrap_parameters(model) + assert reapply_trainability is not None + reapply_trainability(model) + model.trainability_at_wrapper_construction = { + name: param.requires_grad for name, param in model.named_parameters(remove_duplicate=False) + } + return model + + freeze_kwargs = {"freeze_config": freeze_config} if freeze_config is not None else {} + with ( + patch(f"{_INFRA_MODULE}.get_world_size_safe", return_value=1), + patch(f"{_INFRA_MODULE}._supports_logits_to_keep", return_value=True), + patch(f"{_INFRA_MODULE}.print_trainable_parameters"), + patch(f"{_INFRA_MODULE}._should_load_before_shard", return_value=False), + patch(f"{_INFRA_MODULE}._apply_peft_and_lower_precision", return_value=model), + patch(f"{_INFRA_MODULE}.Checkpointer") as MockCheckpointer, + ): + MockCheckpointer.return_value.config.dequantize_base_checkpoint = False + return apply_model_infrastructure( + model=model, + is_meta_device=False, + device=torch.device("cpu"), + load_base_model=False, + model_wrapper=_ReplacingWrapper(), + peft_config=peft_config, + **freeze_kwargs, + ) + + +@pytest.mark.parametrize( + ("freeze_config", "extension_trainable", "vision_lora_trainable"), + [ + pytest.param(None, False, True, id="no-freeze-config"), + pytest.param({"freeze_vision_tower": False}, False, True, id="disabled-legacy-alias"), + pytest.param({"freeze_vision_tower": True}, False, False, id="enabled-legacy-alias"), + pytest.param(_GENERIC_FREEZE_CONFIG, True, True, id="generic-freeze-config"), + ], +) +def test_peft_freeze_config_survives_name_changing_parameter_replacement( + freeze_config: dict | None, extension_trainable: bool, vision_lora_trainable: bool +): + """Under PEFT, selector overrides rebind after surgery while new non-LoRA parameters freeze.""" + model = _TinyTrainabilityModel() + model.register_parameter("lora_weight", torch.nn.Parameter(torch.ones(1))) + model.vision_adapter = torch.nn.Module() + model.vision_adapter.register_parameter("lora_weight", torch.nn.Parameter(torch.ones(1))) + peft_config = SimpleNamespace(lora_A_init=None, use_triton=False) + + _run_trainability_infrastructure(model, freeze_config, peft_config=peft_config) + + assert not model.base.weight.requires_grad + assert model.extension[0].weight.requires_grad is extension_trainable + assert model.lora_weight.requires_grad + assert model.vision_adapter.lora_weight.requires_grad is vision_lora_trainable + assert not model.parallel_parameter.requires_grad + assert model.trainability_at_wrapper_construction["extension.0.weight"] is extension_trainable + assert model.trainability_at_wrapper_construction["vision_adapter.lora_weight"] is vision_lora_trainable + assert not model.trainability_at_wrapper_construction["parallel_parameter"] + + +def test_freeze_config_rebinds_after_name_changing_replacement_under_full_finetuning(): + """Under full fine-tuning, selectors re-resolve on the post-surgery module hierarchy.""" + freeze_config = { + "freeze_modules": [{"path": "base"}, {"glob": "ext*"}], + "unfreeze_modules": [{"path": "extension"}], + } + + model = _run_trainability_infrastructure(_TinyTrainabilityModel(), freeze_config) + + assert not model.base.weight.requires_grad + # The renamed extension parameter is re-selected through its parent module path. + assert model.extension[0].weight.requires_grad + # Parameters the policy does not select retain their model-owned trainability. + assert not model.model_constant.requires_grad + assert model.parallel_parameter.requires_grad + assert not model.trainability_at_wrapper_construction["base.weight"] + assert model.trainability_at_wrapper_construction["extension.0.weight"] + assert not model.trainability_at_wrapper_construction["model_constant"] + assert model.trainability_at_wrapper_construction["parallel_parameter"] + + +def test_freeze_config_unfreeze_keeps_parameter_storage_dtype(): + """Freeze configuration controls requires_grad only; it never casts trainable params.""" + from nemo_automodel._transformers.infrastructure import apply_model_infrastructure + + model = _TinyClassifierModel() + peft_config = SimpleNamespace(lora_A_init=None, use_triton=False) + model_wrapper = SimpleNamespace(mp_policy=SimpleNamespace(param_dtype=torch.bfloat16)) + + with ( + patch(f"{_INFRA_MODULE}.get_world_size_safe", return_value=1), + patch(f"{_INFRA_MODULE}._supports_logits_to_keep", return_value=True), + patch(f"{_INFRA_MODULE}.print_trainable_parameters"), + patch(f"{_INFRA_MODULE}._should_load_before_shard", return_value=False), + patch(f"{_INFRA_MODULE}._apply_peft_and_lower_precision", return_value=model), + patch(f"{_INFRA_MODULE}._shard_ep_fsdp", return_value=model), + patch(f"{_INFRA_MODULE}.Checkpointer") as MockCheckpointer, + ): + MockCheckpointer.return_value.config.dequantize_base_checkpoint = False + result = apply_model_infrastructure( + model=model, + is_meta_device=False, + device=torch.device("cpu"), + load_base_model=False, + model_wrapper=model_wrapper, + peft_config=peft_config, + freeze_config={"unfreeze_modules": [{"path": "classifier"}]}, + ) + + assert not result.backbone.weight.requires_grad + # Frozen plain params are cast to the compute dtype to avoid a mixed-dtype seam. + assert result.backbone.weight.dtype == torch.bfloat16 + assert result.classifier.weight.requires_grad + # Unfreezing must not turn fp32 resident/master parameters into bf16 storage; + # the compute dtype is owned by autocast or the distributed mp policy. + assert result.classifier.weight.dtype == torch.float32 + + def test_safe_moe_tp_requires_real_checkpoint_source_and_rejects_peft(): from nemo_automodel._transformers.infrastructure import ( _validate_safe_moe_tp_weight_source, diff --git a/tests/unit_tests/distributed/test_ddp_manager.py b/tests/unit_tests/distributed/test_ddp_manager.py index 22c9a34636..478ee245c3 100644 --- a/tests/unit_tests/distributed/test_ddp_manager.py +++ b/tests/unit_tests/distributed/test_ddp_manager.py @@ -53,6 +53,33 @@ def test_ddp_manager_forwards_ddp_constructor_flags(monkeypatch): assert ddp_ctor.call_args.kwargs["gradient_as_bucket_view"] is True +def test_ddp_manager_reapplies_trainability_before_constructor(monkeypatch): + monkeypatch.setattr(ddp_mod.dist, "is_available", lambda: True, raising=True) + monkeypatch.setattr(ddp_mod.dist, "is_initialized", lambda: True, raising=True) + monkeypatch.setattr(ddp_mod.dist, "get_rank", lambda: 0, raising=True) + monkeypatch.setattr(ddp_mod.dist, "get_world_size", lambda: 2, raising=True) + monkeypatch.setattr(ddp_mod.dist, "get_backend", lambda: "gloo", raising=True) + + model = nn.Linear(2, 2) + model.weight.requires_grad_(False) + events = [] + + def reapply_trainability(module): + events.append("reapply") + module.weight.requires_grad_(True) + + def construct_ddp(module, **_kwargs): + events.append("ddp") + assert module.weight.requires_grad + return "wrapped" + + monkeypatch.setattr(ddp_mod, "DDP", construct_ddp, raising=True) + manager = ddp_mod.DDPManager(DDPConfig()) + + assert manager.parallelize(model, reapply_trainability=reapply_trainability) == "wrapped" + assert events == ["reapply", "ddp"] + + def test_ddp_manager_applies_selective_activation_checkpointing(monkeypatch): monkeypatch.setattr(ddp_mod.dist, "is_available", lambda: True, raising=True) monkeypatch.setattr(ddp_mod.dist, "is_initialized", lambda: True, raising=True) diff --git a/tests/unit_tests/distributed/test_parallelization_strategies.py b/tests/unit_tests/distributed/test_parallelization_strategies.py index 5bb565d6f6..3a3409def6 100644 --- a/tests/unit_tests/distributed/test_parallelization_strategies.py +++ b/tests/unit_tests/distributed/test_parallelization_strategies.py @@ -323,6 +323,7 @@ def test_parallelize_method_signature(self, strategy): "dp_shard_cp_mesh_name", "tp_mesh_name", "frozen_multimodal_sharding", + "reapply_trainability", ] for param in required_params: @@ -368,6 +369,24 @@ def test_parallelize_with_tensor_parallel(self, strategy, mock_device_mesh, mock mock_distributed_env["get_plan"].assert_called_once() mock_distributed_env["parallelize_module"].assert_called_once() + def test_trainability_rebind_runs_after_tp_and_before_fsdp(self, strategy, mock_device_mesh, mock_distributed_env): + """FSDP captures the selector result on the post-TP hierarchy.""" + mesh, _, _, tp_mesh = mock_device_mesh + tp_mesh.size.return_value = 2 + model = MockModel() + events = [] + + mock_distributed_env["parallelize_module"].side_effect = lambda *_args, **_kwargs: events.append("tp") + mock_distributed_env["apply_fsdp"].side_effect = lambda *_args, **_kwargs: events.append("fsdp") + + strategy.parallelize( + model=model, + device_mesh=mesh, + reapply_trainability=lambda _model: events.append("trainability"), + ) + + assert events == ["tp", "trainability", "fsdp"] + def test_parallelize_with_activation_checkpointing(self, strategy, mock_device_mesh, mock_distributed_env): """Test parallelization with activation checkpointing enabled.""" mesh, dp_replicate_mesh, dp_shard_mesh, tp_mesh = mock_device_mesh @@ -1237,6 +1256,7 @@ def test_preserves_function_signature(self): "dp_shard_cp_mesh_name", "tp_mesh_name", "frozen_multimodal_sharding", + "reapply_trainability", ] for param in expected_params: diff --git a/tests/unit_tests/moe/test_parallelizer.py b/tests/unit_tests/moe/test_parallelizer.py index d4770b5793..5a68ee4179 100644 --- a/tests/unit_tests/moe/test_parallelizer.py +++ b/tests/unit_tests/moe/test_parallelizer.py @@ -1621,9 +1621,10 @@ def test_parallelize_model_applies_tp_before_cp_ep_ac_and_fsdp(monkeypatch): tp_axis_name="tp", ep_axis_name="ep", activation_checkpointing=True, + reapply_trainability=lambda _model: calls.append("trainability"), ) - assert calls == ["tp", "tie", "cp", "ep", "ac", "fsdp"] + assert calls == ["tp", "tie", "cp", "ep", "ac", "trainability", "fsdp"] assert model._nemo_moe_tp_requires_replica_sync is True assert model._nemo_moe_tp_requires_pretrained_weights is True P._resolve_moe_tp_plan.assert_called_once_with( diff --git a/tests/unit_tests/recipes/test_base_recipe.py b/tests/unit_tests/recipes/test_base_recipe.py index 38351971cb..f8bd2da0aa 100644 --- a/tests/unit_tests/recipes/test_base_recipe.py +++ b/tests/unit_tests/recipes/test_base_recipe.py @@ -303,6 +303,20 @@ def __init__(self, checkpoint_dir, cfg_dict=None, max_recent_checkpoints="defaul self.cfg = ConfigNode(cfg_dict) +def test_recipe_autocast_preserves_fp32_parameter_storage(): + """Configured compute autocast must not mutate resident/master parameters.""" + recipe = BaseRecipe.__new__(BaseRecipe) + recipe.__dict__["distributed_config"] = SimpleNamespace(autocast_dtype=torch.bfloat16) + recipe.__dict__["dist_env"] = SimpleNamespace(device=torch.device("cpu")) + linear = nn.Linear(2, 2, dtype=torch.float32) + + with recipe._autocast_context(): + output = linear(torch.ones(1, 2, dtype=torch.float32)) + + assert output.dtype == torch.bfloat16 + assert linear.weight.dtype == torch.float32 + + def test_dp_allreduce_uses_world_group_without_device_mesh(tmp_path, monkeypatch): """ DDP does not create a device mesh, so DP reductions should use the default diff --git a/tests/unit_tests/recipes/test_train_ft.py b/tests/unit_tests/recipes/test_train_ft.py index c197a2b93d..dc7fba69dd 100644 --- a/tests/unit_tests/recipes/test_train_ft.py +++ b/tests/unit_tests/recipes/test_train_ft.py @@ -854,6 +854,41 @@ def __exit__(self, exc_type, exc, tb): return False +def test_build_model_passes_freeze_config(monkeypatch): + """LLM model construction forwards freeze_config to NeMoAutoModel.""" + from nemo_automodel._transformers import NeMoAutoModelForCausalLM + + captured_kwargs = {} + + class CapturingModelConfig: + def __init__(self): + self._target_ = NeMoAutoModelForCausalLM.from_pretrained + + def instantiate(self, **kwargs): + captured_kwargs.update(kwargs) + return DummyModel() + + def get(self, key, default=None): + return getattr(self, key, default) + + freeze_config = ConfigNode( + { + "unfreeze_modules": [{"path": "layer2"}], + } + ) + monkeypatch.setattr("nemo_automodel.recipes.llm.train_ft.ScopedRNG", lambda **kwargs: nullcontext()) + + build_model( + cfg_model=CapturingModelConfig(), + cfg_peft=None, + cfg_freeze=freeze_config, + seed=123, + ) + + # ConfigNode is unwrapped at the recipe boundary; downstream receives a mapping. + assert captured_kwargs["freeze_config"] == {"unfreeze_modules": [{"path": "layer2"}]} + + @requires_cuda def test_force_hf_true_disables_meta_init(monkeypatch): """When cfg_model.force_hf=True, meta-device init (init_empty_weights) should not be used. @@ -1109,6 +1144,50 @@ def test_setup_does_not_change_storage_dtype_for_non_kd_recipe(monkeypatch): assert not hasattr(cfg.model, "torch_dtype") +def test_freeze_config_applies_before_optimizer_build(monkeypatch): + """The optimizer sees the trainability selected through freeze_config.""" + from nemo_automodel.components.utils.model_utils import apply_parameter_freezing, parse_freeze_config + + cfg = _minimal_cfg_with_nvtx(nvtx_value=False) + cfg.freeze_config = ConfigNode( + { + "freeze_modules": [{"glob": "layer*"}], + "unfreeze_modules": [{"glob": "*2"}], + } + ) + _patch_setup_minimals(monkeypatch, lambda *args, **kwargs: None) + + model = DummyModel() + freeze_configs = [] + + def _build_model(*args, cfg_freeze=None, **kwargs): + freeze_configs.append(cfg_freeze) + if cfg_freeze is not None: + apply_parameter_freezing(model, parse_freeze_config(cfg_freeze.to_dict())) + return model + + trainable_at_optimizer_build = [] + + def _build_optimizer(model, *args, **kwargs): + trainable_at_optimizer_build.extend( + name for name, parameter in model.named_parameters() if parameter.requires_grad + ) + return [SimpleNamespace(param_groups=[{"lr": 0.01}], step=lambda: None, zero_grad=lambda: None)] + + monkeypatch.setattr("nemo_automodel.recipes.llm.train_ft.build_model", _build_model) + monkeypatch.setattr( + "nemo_automodel.recipes._typed_config.RecipeConfig.optimizer", + property(lambda self: SimpleNamespace(build=_build_optimizer)), + ) + + trainer = TrainFinetuneRecipeForNextTokenPrediction(cfg) + trainer.setup() + + assert freeze_configs[0] is not None + assert freeze_configs[0].to_dict() == cfg.freeze_config.to_dict() + assert trainable_at_optimizer_build == ["layer2.weight"] + + def test_nvtx_true_pipeline_patches_all_parts(monkeypatch): cfg = _minimal_cfg_with_nvtx(nvtx_value=True) patch_calls = [] diff --git a/tests/unit_tests/utils/test_model_utils.py b/tests/unit_tests/utils/test_model_utils.py index c4ba8d28b2..794769dc2e 100644 --- a/tests/unit_tests/utils/test_model_utils.py +++ b/tests/unit_tests/utils/test_model_utils.py @@ -14,6 +14,7 @@ from __future__ import annotations +import warnings from types import MethodType from typing import Dict @@ -143,7 +144,7 @@ def test_apply_parameter_freezing(dummy_model, freeze_cfg: Dict, expect: Dict): for p in dummy_model.parameters(): p.requires_grad = True - model_utils.apply_parameter_freezing(dummy_model, freeze_cfg) + model_utils.apply_parameter_freezing(dummy_model, model_utils.parse_freeze_config(freeze_cfg)) # vision tower(s) assert _all_requires_grad(dummy_model.vision_tower) is expect["vision"] @@ -159,6 +160,241 @@ def test_apply_parameter_freezing(dummy_model, freeze_cfg: Dict, expect: Dict): assert dummy_model.token_embed.weight.requires_grad is True +def test_apply_parameter_freezing_supports_generic_module_selectors(): + """Generic selectors match fully qualified module names and unfreeze wins on overlap.""" + + class GenericModel(nn.Module): + def __init__(self) -> None: + super().__init__() + self.backbone = nn.ModuleDict( + { + "q_proj": nn.Linear(4, 4), + "k_proj": nn.Linear(4, 4), + "dense": nn.Linear(4, 4), + } + ) + self.backbone["k_proj"].requires_grad_(False) + + model = GenericModel() + + model_utils.apply_parameter_freezing( + model, + model_utils.parse_freeze_config( + { + "freeze_modules": [{"glob": "*_proj"}], + "unfreeze_modules": [{"glob": "*.k_proj"}], + } + ), + ) + + assert not _any_requires_grad(model.backbone["q_proj"]) + assert _all_requires_grad(model.backbone["k_proj"]) + assert _all_requires_grad(model.backbone["dense"]) + + +def test_generic_module_selector_matches_shared_module_alias(): + """A selector can address any fully qualified alias of a shared module.""" + + class SharedModel(nn.Module): + def __init__(self) -> None: + super().__init__() + shared = nn.Linear(4, 4) + self.primary = shared + self.alias = shared + + model = SharedModel() + model_utils.apply_parameter_freezing( + model, + model_utils.FreezeConfig(freeze_modules=[model_utils.ModuleSelector(path="alias")]), + ) + + assert not _any_requires_grad(model.primary) + + +def test_generic_unfreeze_selector_overrides_legacy_freeze_alias(dummy_model): + """Generic unfreeze selectors take precedence over legacy freeze aliases.""" + model_utils.apply_parameter_freezing( + dummy_model, + model_utils.parse_freeze_config( + { + "freeze_vision_tower": True, + "unfreeze_modules": [{"path": "vision_tower"}], + } + ), + ) + + assert _all_requires_grad(dummy_model.vision_tower) + assert not _any_requires_grad(dummy_model.visual_extra) + + +@pytest.mark.parametrize( + "legacy_option", + ["freeze_vision_tower", "freeze_audio_tower", "freeze_language_model", "freeze_video_embedder"], +) +def test_legacy_freeze_options_remain_supported_without_deprecation(dummy_model, legacy_option: str): + """Legacy modality booleans are supported and do not emit deprecation warnings.""" + with warnings.catch_warnings(): + warnings.simplefilter("error", FutureWarning) + model_utils.apply_parameter_freezing( + dummy_model, model_utils.parse_freeze_config({legacy_option: False, "freeze_vision_tower": False}) + ) + + assert _all_requires_grad(dummy_model.other) + + +def test_legacy_freeze_alias_uses_substring_matching(dummy_model): + """Legacy aliases retain their original substring matcher ("visual" matches "visual_extra").""" + model_utils.apply_parameter_freezing(dummy_model, model_utils.FreezeConfig(freeze_vision_tower=True)) + + assert not _any_requires_grad(dummy_model.vision_tower) + assert not _any_requires_grad(dummy_model.visual_extra) + + +def test_generic_selectors_disable_implicit_legacy_vision_freeze(dummy_model): + """Explicit-selector configurations do not implicitly freeze vision modules.""" + model_utils.apply_parameter_freezing( + dummy_model, + model_utils.FreezeConfig(unfreeze_modules=[model_utils.ModuleSelector(path="other")]), + ) + + assert _all_requires_grad(dummy_model.vision_tower) + assert _all_requires_grad(dummy_model.visual_extra) + assert _all_requires_grad(dummy_model.other) + + +def test_empty_generic_selector_list_disables_implicit_legacy_vision_freeze(dummy_model): + """An explicitly empty selector field opts out of legacy implicit freezing.""" + config = model_utils.parse_freeze_config({"freeze_modules": []}) + + assert config.freeze_modules == [] + assert config.unfreeze_modules is None + model_utils.apply_parameter_freezing(dummy_model, config) + + assert _all_requires_grad(dummy_model.vision_tower) + assert _all_requires_grad(dummy_model.visual_extra) + + +def test_glob_selector_treats_dot_as_literal(): + """Glob matching is shell-style: `.` is literal, not a regex wildcard.""" + + class DotModel(nn.Module): + def __init__(self) -> None: + super().__init__() + self.axb = nn.Linear(4, 4) # would match the regular expression r"^a.b$" + self.a = nn.ModuleDict({"b": nn.Linear(4, 4)}) # canonical module path "a.b" + + model = DotModel() + model_utils.apply_parameter_freezing( + model, + model_utils.FreezeConfig(freeze_modules=[model_utils.ModuleSelector(glob="a.b")]), + ) + + assert not _any_requires_grad(model.a["b"]) + assert _all_requires_grad(model.axb) + + +def test_glob_selector_is_case_sensitive(): + """Glob matching uses fnmatchcase and never folds case.""" + assert model_utils.ModuleSelector(glob="*Encoder").matches("audio_encoder") is False + assert model_utils.ModuleSelector(glob="*Encoder").matches("audio_Encoder") is True + assert model_utils.ModuleSelector(path="audio_encoder").matches("model.audio_encoder") is False + + +class TestParseFreezeConfig: + """Strict validation of the freeze_config schema.""" + + def test_none_and_typed_passthrough(self): + typed = model_utils.FreezeConfig() + assert model_utils.parse_freeze_config(None) is None + assert model_utils.parse_freeze_config(typed) is typed + + def test_empty_mapping_retains_legacy_selector_absence(self): + config = model_utils.parse_freeze_config({}) + + assert config.freeze_modules is None + assert config.unfreeze_modules is None + assert not config.has_generic_selectors() + + def test_parses_typed_selectors(self): + config = model_utils.parse_freeze_config( + { + "freeze_modules": [{"path": "vision_tower"}, {"glob": "*.audio_encoder"}], + "unfreeze_modules": [{"path": "multi_modal_projector"}], + } + ) + assert config.freeze_modules == [ + model_utils.ModuleSelector(path="vision_tower"), + model_utils.ModuleSelector(glob="*.audio_encoder"), + ] + assert config.unfreeze_modules == [model_utils.ModuleSelector(path="multi_modal_projector")] + + def test_unknown_option_rejected(self): + with pytest.raises(ValueError, match="unsupported option.*freeze_vision"): + model_utils.parse_freeze_config({"freeze_vision": True}) + + def test_hydra_meta_option_rejected(self): + with pytest.raises(ValueError, match="unsupported option.*_target_"): + model_utils.parse_freeze_config({"_target_": "example.FreezeConfig"}) + + def test_bare_string_selector_rejected(self): + with pytest.raises(ValueError, match="must be mappings with exactly one of `path` or `glob`"): + model_utils.parse_freeze_config({"freeze_modules": ["vision_tower"]}) + + def test_selector_list_required(self): + with pytest.raises(ValueError, match="freeze_config.freeze_modules must be a list"): + model_utils.parse_freeze_config({"freeze_modules": {"path": "vision_tower"}}) + with pytest.raises(ValueError, match="freeze_config.freeze_modules must be a list"): + model_utils.parse_freeze_config({"freeze_modules": ({"path": "vision_tower"},)}) + + @pytest.mark.parametrize("field_name", ["freeze_modules", "unfreeze_modules"]) + @pytest.mark.parametrize("invalid", [None, False, 0]) + def test_empty_like_selector_values_rejected(self, field_name, invalid): + with pytest.raises(ValueError, match=rf"freeze_config\.{field_name} must be a list"): + model_utils.parse_freeze_config({field_name: invalid}) + + def test_selector_with_both_path_and_glob_rejected(self): + with pytest.raises(ValueError, match="exactly one of `path` or `glob`"): + model_utils.parse_freeze_config({"freeze_modules": [{"path": "a", "glob": "b"}]}) + + def test_selector_with_neither_path_nor_glob_rejected(self): + with pytest.raises(ValueError, match="exactly one of `path` or `glob`"): + model_utils.parse_freeze_config({"unfreeze_modules": [{}]}) + + def test_selector_unknown_key_rejected(self): + with pytest.raises(ValueError, match="unsupported key.*regex"): + model_utils.parse_freeze_config({"freeze_modules": [{"regex": "vision"}]}) + + def test_non_boolean_legacy_option_rejected(self): + with pytest.raises(ValueError, match="freeze_config.freeze_audio_tower must be a boolean"): + model_utils.parse_freeze_config({"freeze_audio_tower": "yes"}) + + def test_non_mapping_rejected(self): + with pytest.raises(TypeError, match="freeze_config must be a mapping"): + model_utils.parse_freeze_config("vision_tower") + + def test_direct_typed_config_rejects_untyped_values(self): + with pytest.raises(TypeError, match="entries must be ModuleSelector"): + model_utils.FreezeConfig(freeze_modules=[{"path": "vision_tower"}]) + with pytest.raises(TypeError, match="freeze_audio_tower must be a boolean"): + model_utils.FreezeConfig(freeze_audio_tower="yes") + + def test_apply_parameter_freezing_accepts_raw_mapping_for_compatibility(self, dummy_model): + model_utils.apply_parameter_freezing(dummy_model, {"freeze_modules": [{"path": "other"}]}) + assert not _any_requires_grad(dummy_model.other) + + def test_unmatched_selector_raises_in_strict_mode(self, dummy_model): + config = model_utils.FreezeConfig(freeze_modules=[model_utils.ModuleSelector(path="nonexistent")]) + with pytest.raises( + ValueError, match=r"freeze_config\.freeze_modules selector `path: nonexistent` matched no parameters" + ): + model_utils.apply_parameter_freezing(dummy_model, config, strict=True) + + def test_unmatched_selector_allowed_when_rebinding(self, dummy_model): + config = model_utils.FreezeConfig(freeze_modules=[model_utils.ModuleSelector(path="nonexistent")]) + model_utils.apply_parameter_freezing(dummy_model, config, strict=False) + assert _all_requires_grad(dummy_model.other) + + def test_init_empty_weights_moves_params_to_meta_and_preserves_requires_grad(): """ Creating parameters inside the context should place them on meta device and @@ -333,7 +569,9 @@ def forward(self, x): def test_freeze_audio_tower_freezes_speech_pattern(audio_model): """freeze_audio_tower=True should freeze modules matching 'audio' and 'speech'.""" - model_utils.apply_parameter_freezing(audio_model, {"freeze_audio_tower": True, "freeze_vision_tower": False}) + model_utils.apply_parameter_freezing( + audio_model, model_utils.FreezeConfig(freeze_audio_tower=True, freeze_vision_tower=False) + ) assert not _any_requires_grad(audio_model.audio_encoder) assert not _any_requires_grad(audio_model.speech_adapter) @@ -342,12 +580,25 @@ def test_freeze_audio_tower_freezes_speech_pattern(audio_model): def test_freeze_audio_tower_false_keeps_speech_trainable(audio_model): """freeze_audio_tower=False should keep speech modules trainable.""" - model_utils.apply_parameter_freezing(audio_model, {"freeze_audio_tower": False, "freeze_vision_tower": False}) + model_utils.apply_parameter_freezing( + audio_model, model_utils.FreezeConfig(freeze_audio_tower=False, freeze_vision_tower=False) + ) assert _all_requires_grad(audio_model.audio_encoder) assert _all_requires_grad(audio_model.speech_adapter) +def test_freeze_audio_encoder_with_generic_selector(audio_model): + """A glob selector freezes only the matched module, without the legacy substrings.""" + model_utils.apply_parameter_freezing( + audio_model, + model_utils.FreezeConfig(freeze_modules=[model_utils.ModuleSelector(glob="*audio_encoder")]), + ) + + assert not _any_requires_grad(audio_model.audio_encoder) + assert _all_requires_grad(audio_model.speech_adapter) + + @pytest.fixture() def sound_model() -> nn.Module: """Model with a NemotronOmni-style ``sound_*`` submodule.""" @@ -367,7 +618,9 @@ def forward(self, x): def test_freeze_audio_tower_freezes_sound_pattern(sound_model): """``freeze_audio_tower=True`` must also freeze NemotronOmni's ``sound_*`` modules.""" - model_utils.apply_parameter_freezing(sound_model, {"freeze_audio_tower": True, "freeze_vision_tower": False}) + model_utils.apply_parameter_freezing( + sound_model, model_utils.FreezeConfig(freeze_audio_tower=True, freeze_vision_tower=False) + ) assert not _any_requires_grad(sound_model.sound_encoder) assert not _any_requires_grad(sound_model.sound_projection) @@ -376,7 +629,9 @@ def test_freeze_audio_tower_freezes_sound_pattern(sound_model): def test_freeze_audio_tower_false_keeps_sound_trainable(sound_model): """Sound modules stay trainable when audio freeze is off.""" - model_utils.apply_parameter_freezing(sound_model, {"freeze_audio_tower": False, "freeze_vision_tower": False}) + model_utils.apply_parameter_freezing( + sound_model, model_utils.FreezeConfig(freeze_audio_tower=False, freeze_vision_tower=False) + ) assert _all_requires_grad(sound_model.sound_encoder) assert _all_requires_grad(sound_model.sound_projection) @@ -420,7 +675,7 @@ def test_phi4mm_use_cache_disabled(): m = nn.Linear(4, 4) m.config = types.SimpleNamespace(model_type="phi4mm", use_cache=True) - model_utils.apply_parameter_freezing(m, {"freeze_vision_tower": False}) + model_utils.apply_parameter_freezing(m, model_utils.FreezeConfig()) assert m.config.use_cache is False @@ -431,7 +686,7 @@ def test_non_phi4mm_use_cache_unchanged(): m = nn.Linear(4, 4) m.config = types.SimpleNamespace(model_type="gemma3", use_cache=True) - model_utils.apply_parameter_freezing(m, {"freeze_vision_tower": False}) + model_utils.apply_parameter_freezing(m, model_utils.FreezeConfig()) assert m.config.use_cache is True @@ -720,7 +975,7 @@ def test_enable_radio_vit_fused_attn_skips_blocks_without_attn(): # ============================================================================= -# Tests for freeze_video_embedder branch in apply_parameter_freezing +# Tests for the freeze_video_embedder compatibility alias # ============================================================================= @@ -747,7 +1002,7 @@ def test_freeze_video_embedder_true_freezes_only_video_embedder(video_embedder_m """``freeze_video_embedder=True`` freezes ``patch_generator.video_embedder.*`` only.""" model_utils.apply_parameter_freezing( video_embedder_model, - {"freeze_vision_tower": False, "freeze_video_embedder": True}, + model_utils.FreezeConfig(freeze_vision_tower=False, freeze_video_embedder=True), ) assert not _any_requires_grad(video_embedder_model.patch_generator.video_embedder) @@ -759,7 +1014,7 @@ def test_freeze_video_embedder_false_keeps_video_embedder_trainable(video_embedd """Default ``freeze_video_embedder=False`` leaves the video embedder trainable.""" model_utils.apply_parameter_freezing( video_embedder_model, - {"freeze_vision_tower": False, "freeze_video_embedder": False}, + model_utils.FreezeConfig(freeze_vision_tower=False, freeze_video_embedder=False), ) assert _all_requires_grad(video_embedder_model.patch_generator.video_embedder) diff --git a/tutorials/nemotron-parse/nemotron_parse_config.yaml b/tutorials/nemotron-parse/nemotron_parse_config.yaml index 6b8b20b0ef..b550edc975 100644 --- a/tutorials/nemotron-parse/nemotron_parse_config.yaml +++ b/tutorials/nemotron-parse/nemotron_parse_config.yaml @@ -88,6 +88,5 @@ optimizer: betas: [0.9, 0.999] freeze_config: - freeze_embeddings: true freeze_vision_tower: true freeze_language_model: false