RASMUS/Finnish-ASR-Canary-v2
02.2k
1# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7# http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14 15import os16 17import lightning.pytorch as pl18import nemo_run as run19import torch20from lightning.pytorch.loggers import WandbLogger21from megatron.core.distributed import DistributedDataParallelConfig22from megatron.core.optimizer import OptimizerConfig23from megatron.core.transformer.enums import AttnMaskType24 25from nemo import lightning as nl26from nemo.collections import llm27from nemo.collections.diffusion.data.diffusion_energon_datamodule import DiffusionDataModule28from nemo.collections.diffusion.data.diffusion_fake_datamodule import VideoLatentFakeDataModule29from nemo.collections.diffusion.data.diffusion_taskencoder import BasicDiffusionTaskEncoder30from nemo.collections.diffusion.models.model import (31 DiT7BConfig,32 DiTConfig,33 DiTLConfig,34 DiTLlama1BConfig,35 DiTLlama5BConfig,36 DiTLlama30BConfig,37 DiTModel,38 DiTXLConfig,39 ECDiTLlama1BConfig,40)41from nemo.collections.multimodal.data.energon.base import EnergonMultiModalDataModule42from nemo.lightning.pytorch.callbacks import ModelCheckpoint, PreemptionCallback43from nemo.lightning.pytorch.callbacks.megatron_comm_overlap import MegatronCommOverlapCallback44from nemo.lightning.pytorch.callbacks.model_transform import ModelTransform45from nemo.lightning.pytorch.callbacks.nsys import NsysCallback46from nemo.lightning.pytorch.strategies.utils import RestoreConfig47from nemo.utils.exp_manager import TimingCallback48 49 50@run.cli.factory51@run.autoconvert52def multimodal_datamodule() -> pl.LightningDataModule:53 """Multimodal Datamodule Initialization"""54 data_module = DiffusionDataModule(55 seq_length=2048,56 task_encoder=run.Config(BasicDiffusionTaskEncoder, seq_length=2048),57 micro_batch_size=1,58 global_batch_size=32,59 )60 return data_module61 62 63@run.cli.factory64@run.autoconvert65def simple_datamodule() -> pl.LightningDataModule:66 """Simple Datamodule Initialization"""67 data_module = EnergonMultiModalDataModule(68 seq_length=2048,69 micro_batch_size=1,70 global_batch_size=32,71 num_workers=16,72 tokenizer=None,73 image_processor=None,74 task_encoder=run.Config(BasicDiffusionTaskEncoder, seq_length=2048),75 )76 return data_module77 78 79@run.cli.factory80@run.autoconvert81def multimodal_fake_datamodule() -> pl.LightningDataModule:82 """Multimodal Mock Datamodule Initialization"""83 data_module = VideoLatentFakeDataModule(84 seq_length=None, # Set None to dectect the sequence length automatically.85 task_encoder=run.Config(BasicDiffusionTaskEncoder, seq_length=2048),86 micro_batch_size=1,87 global_batch_size=32,88 )89 return data_module90 91 92@run.cli.factory93@run.autoconvert94def peft(args) -> ModelTransform:95 """Parameter Efficient Fine Tuning"""96 return llm.peft.LoRA(97 target_modules=['linear_qkv', 'linear_proj'], # , 'linear_fc1', 'linear_fc2'],98 dim=args.lora_dim,99 )100 101 102@run.cli.factory(target=llm.train)103def pretrain() -> run.Partial:104 """Base Pretraining Config"""105 return run.Partial(106 llm.train,107 model=run.Config(108 DiTModel,109 config=run.Config(DiTConfig),110 ),111 data=multimodal_datamodule(),112 trainer=run.Config(113 nl.Trainer,114 devices='auto',115 num_nodes=int(os.environ.get('SLURM_NNODES', 1)),116 accelerator="gpu",117 strategy=run.Config(118 nl.MegatronStrategy,119 tensor_model_parallel_size=1,120 pipeline_model_parallel_size=1,121 context_parallel_size=1,122 sequence_parallel=False,123 pipeline_dtype=torch.bfloat16,124 ddp=run.Config(125 DistributedDataParallelConfig,126 check_for_nan_in_grad=True,127 grad_reduce_in_fp32=True,128 overlap_grad_reduce=True,129 overlap_param_gather=True,130 ),131 ),132 plugins=nl.MegatronMixedPrecision(precision="bf16-mixed"),133 num_sanity_val_steps=0,134 limit_val_batches=1,135 val_check_interval=1000,136 max_epochs=10000,137 log_every_n_steps=1,138 callbacks=[139 run.Config(140 ModelCheckpoint,141 monitor='global_step',142 filename='{global_step}',143 every_n_train_steps=1000,144 save_top_k=3,145 mode='max',146 ),147 run.Config(PreemptionCallback),148 run.Config(TimingCallback),149 run.Config(150 MegatronCommOverlapCallback,151 tp_comm_overlap=False,152 ),153 ],154 ),155 log=nl.NeMoLogger(wandb=(WandbLogger() if "WANDB_API_KEY" in os.environ else None)),156 optim=run.Config(157 nl.MegatronOptimizerModule,158 config=run.Config(159 OptimizerConfig,160 lr=1e-4,161 bf16=True,162 params_dtype=torch.bfloat16,163 use_distributed_optimizer=True,164 weight_decay=0,165 ),166 ),167 tokenizer=None,168 resume=run.Config(169 nl.AutoResume,170 resume_if_exists=True,171 resume_ignore_no_checkpoint=True,172 resume_past_end=True,173 ),174 model_transform=None,175 )176 177 178@run.cli.factory(target=llm.train)179def pretrain_xl() -> run.Partial:180 """DiT-XL Pretraining Recipe"""181 recipe = pretrain()182 recipe.model.config = run.Config(DiTXLConfig)183 return recipe184 185 186@run.cli.factory(target=llm.train)187def pretrain_l() -> run.Partial:188 """DiT-L Pretraining Recipe"""189 recipe = pretrain()190 recipe.model.config = run.Config(DiTLConfig)191 return recipe192 193 194def set_use_megatron_fsdp(recipe):195 try:196 recipe.trainer.strategy.ddp.use_megatron_fsdp = True197 except AttributeError:198 recipe.trainer.strategy.ddp.use_custom_fsdp = True199 200 201@run.cli.factory(target=llm.train)202def train_mock() -> run.Partial:203 """DiT Mock Pretraining Recipe"""204 recipe = pretrain()205 recipe.model.config = run.Config(DiTLlama5BConfig, max_frames=1)206 recipe.data = multimodal_fake_datamodule()207 recipe.model.config.num_layers = 16208 recipe.data.seq_length = 73728209 recipe.data.task_encoder.seq_length = 73728210 recipe.trainer.strategy.tensor_model_parallel_size = 4211 recipe.trainer.strategy.sequence_parallel = True212 recipe.trainer.strategy.context_parallel_size = 2213 recipe.data.micro_batch_size = 1214 recipe.data.global_batch_size = 1215 recipe.trainer.limit_val_batches = 0216 recipe.trainer.val_check_interval = 1.0217 recipe.data.model_config = recipe.model.config218 recipe.log.log_dir = 'nemo_experiments/train_mock'219 220 set_use_megatron_fsdp(recipe=recipe)221 recipe.trainer.strategy.ddp.data_parallel_sharding_strategy = 'optim_grads_params'222 recipe.trainer.strategy.ddp.overlap_param_gather = True223 recipe.trainer.strategy.ddp.overlap_grad_reduce = True224 recipe.model.config.use_cpu_initialization = True225 226 return recipe227 228 229@run.cli.factory(target=llm.train)230def mock_ditllama5b_8k() -> run.Partial:231 """DiT-5B mock Recipe"""232 recipe = pretrain()233 recipe.model.config = run.Config(DiTLlama5BConfig, max_frames=1)234 recipe.data = multimodal_fake_datamodule()235 recipe.data.seq_length = recipe.data.task_encoder.seq_length = 8192236 recipe.trainer.strategy.tensor_model_parallel_size = 2237 recipe.trainer.strategy.sequence_parallel = True238 recipe.trainer.strategy.context_parallel_size = 1239 recipe.data.micro_batch_size = 1240 recipe.data.global_batch_size = 32241 recipe.trainer.limit_val_batches = 0242 recipe.trainer.val_check_interval = 1.0243 recipe.data.model_config = recipe.model.config244 recipe.log.log_dir = 'nemo_experiments/mock_ditllama5b_8k'245 recipe.model.config.attn_mask_type = AttnMaskType.no_mask246 set_use_megatron_fsdp(recipe=recipe)247 recipe.trainer.strategy.ddp.data_parallel_sharding_strategy = 'optim_grads_params'248 recipe.trainer.strategy.ddp.overlap_param_gather = True249 recipe.trainer.strategy.ddp.overlap_grad_reduce = True250 recipe.model.config.use_cpu_initialization = True251 recipe.trainer.max_steps = 15252 recipe.trainer.callbacks.pop(0)253 recipe.trainer.enable_checkpointing = False254 recipe.trainer.callbacks.append(255 run.Config(256 NsysCallback,257 start_step=10,258 end_step=11,259 )260 )261 recipe.resume = None262 return recipe263 264 265@run.cli.factory(target=llm.train)266def mock_dit7b_8k() -> run.Partial:267 """DiT-7B mock Recipe"""268 recipe = mock_ditllama5b_8k()269 recipe.model.config = run.Config(DiT7BConfig, max_frames=1)270 recipe.data.model_config = recipe.model.config271 recipe.model.config.attn_mask_type = AttnMaskType.no_mask272 recipe.model.config.use_cpu_initialization = True273 recipe.log.log_dir = 'nemo_experiments/mock_dit7b_8k'274 return recipe275 276 277@run.cli.factory(target=llm.train)278def pretrain_7b() -> run.Partial:279 """DiT-7B Pretraining Recipe"""280 recipe = pretrain()281 recipe.model.config = run.Config(DiT7BConfig)282 recipe.data.global_batch_size = 4608283 recipe.data.micro_batch_size = 9284 recipe.data.num_workers = 15285 recipe.data.use_train_split_for_val = True286 recipe.data.seq_length = 260287 recipe.data.task_encoder.seq_length = 260288 recipe.trainer.val_check_interval = 1000289 recipe.log.log_dir = 'nemo_experiments/dit7b'290 recipe.optim.lr_scheduler = run.Config(nl.lr_scheduler.WarmupHoldPolicyScheduler, warmup_steps=100, hold_steps=1e9)291 recipe.optim.config.weight_decay = 0.1292 recipe.optim.config.adam_beta1 = 0.9293 recipe.optim.config.adam_beta2 = 0.95294 295 return recipe296 297 298@run.cli.factory(target=llm.train)299def pretrain_7b_pack() -> run.Partial:300 """DiT-7B Pretraining Recipe with Packing"""301 recipe = pretrain_7b()302 recipe.data.global_batch_size = 4608 // 9303 recipe.data.micro_batch_size = 1304 recipe.data.num_workers = 15305 recipe.data.use_train_split_for_val = True306 recipe.data.seq_length = 256 * 9307 recipe.data.packing_buffer_size = 1000308 recipe.data.task_encoder.seq_length = None309 recipe.data.task_encoder.max_seq_length = recipe.data.seq_length310 recipe.model.config.qkv_format = 'thd'311 return recipe312 313 314@run.cli.factory(target=llm.train)315def pretrain_7b_256p_joint() -> run.Partial:316 """DiT-7B Pretraining Recipe 256p Stage 1"""317 recipe = pretrain_7b()318 recipe.data.global_batch_size = 256 # 768319 recipe.data.micro_batch_size = 1320 recipe.data.seq_length = 8192321 recipe.data.task_encoder.seq_length = 8192322 recipe.model.config.seq_length = 8192323 324 recipe.optim.config.lr = 6e-5325 recipe.trainer.strategy.tensor_model_parallel_size = 2326 recipe.trainer.strategy.sequence_parallel = True327 recipe.trainer.strategy.ddp.overlap_grad_reduce = True328 # recipe.resume.restore_config = run.Config(RestoreConfig, path='', load_optim_state=True)329 recipe.log.log_dir = 'nemo_experiments/pretrain_7b_256p_joint'330 return recipe331 332 333@run.cli.factory(target=llm.train)334def pretrain_7b_256p_joint_pack() -> run.Partial:335 """DiT-7B Pretraining Recipe 256p Stage 1 with Packing"""336 recipe = pretrain_7b_256p_joint()337 recipe.data.global_batch_size = 128338 recipe.data.micro_batch_size = 1339 recipe.data.num_workers = 10340 recipe.data.seq_length = recipe.model.config.seq_length = recipe.data.task_encoder.max_seq_length = 10240341 recipe.data.task_encoder.seq_length = None342 recipe.data.packing_buffer_size = 1000343 recipe.data.virtual_epoch_length = 0344 recipe.model.config.qkv_format = 'thd'345 return recipe346 347 348@run.cli.factory(target=llm.train)349def pretrain_ditllama5b() -> run.Partial:350 """MovieGen 5B Training"""351 recipe = pretrain_7b()352 recipe.data.micro_batch_size = 12353 recipe.model.config = run.Config(DiTLlama5BConfig)354 recipe.log.log_dir = 'nemo_experiments/ditllama5b'355 return recipe356 357 358@run.cli.factory(target=llm.train)359def pretrain_ditllama30b() -> run.Partial:360 """MovieGen 30B Stage 1 Training"""361 recipe = pretrain_ditllama5b()362 recipe.model.config = run.Config(DiTLlama30BConfig)363 recipe.data.global_batch_size = 9216364 recipe.data.micro_batch_size = 6365 recipe.data.task_encoder.aethetic_score = 4.0366 recipe.data.seq_length = 256367 recipe.data.task_encoder.seq_length = 256368 recipe.data.virtual_epoch_length = 0369 recipe.log.log_dir = 'nemo_experiments/ditllama30b_stage1_mock'370 set_use_megatron_fsdp(recipe=recipe)371 recipe.trainer.strategy.ddp.data_parallel_sharding_strategy = 'optim_grads_params'372 recipe.trainer.strategy.ddp.overlap_param_gather = True373 recipe.trainer.strategy.ddp.overlap_grad_reduce = True374 recipe.model.config.use_cpu_initialization = True375 return recipe376 377 378@run.cli.factory(target=llm.train)379def pretrain_ditllama30b_stage2_mock() -> run.Partial:380 """MovieGen 30B Stage 2 Training"""381 recipe = pretrain_ditllama5b()382 recipe.model.config = run.Config(DiTLlama30BConfig)383 recipe.data = multimodal_fake_datamodule()384 recipe.data.model_config = recipe.model.config385 recipe.data.seq_length = 8192386 recipe.data.task_encoder.seq_length = 8192387 recipe.data.global_batch_size = 256388 recipe.data.micro_batch_size = 1389 recipe.trainer.strategy.tensor_model_parallel_size = 2390 recipe.trainer.strategy.context_parallel_size = 4391 recipe.trainer.strategy.sequence_parallel = True392 recipe.trainer.limit_val_batches = 0393 recipe.trainer.val_check_interval = 1.0394 recipe.data.model_config = recipe.model.config395 recipe.log.log_dir = 'nemo_experiments/ditllama30b_stage2_mock'396 set_use_megatron_fsdp(recipe=recipe)397 recipe.trainer.strategy.ddp.data_parallel_sharding_strategy = 'optim_grads_params'398 recipe.trainer.strategy.ddp.overlap_param_gather = True399 recipe.trainer.strategy.ddp.overlap_grad_reduce = True400 recipe.model.config.use_cpu_initialization = True401 return recipe402 403 404@run.cli.factory(target=llm.train)405def pretrain_ditllama30b_stage3_mock() -> run.Partial:406 """MovieGen 30B Stage 3 Training"""407 recipe = pretrain_ditllama5b()408 recipe.model.config = run.Config(DiTLlama30BConfig)409 recipe.data = multimodal_fake_datamodule()410 recipe.data.model_config = recipe.model.config411 recipe.data.seq_length = 73728412 recipe.data.task_encoder.seq_length = 73728413 recipe.data.global_batch_size = 256414 recipe.data.micro_batch_size = 1415 recipe.trainer.strategy.tensor_model_parallel_size = 2416 recipe.trainer.strategy.context_parallel_size = 8417 recipe.trainer.strategy.sequence_parallel = True418 recipe.trainer.limit_val_batches = 0419 recipe.trainer.val_check_interval = 1.0420 recipe.data.model_config = recipe.model.config421 recipe.log.log_dir = 'nemo_experiments/ditllama30b_stage3_mock'422 set_use_megatron_fsdp(recipe=recipe)423 recipe.trainer.strategy.ddp.data_parallel_sharding_strategy = 'optim_grads_params'424 recipe.trainer.strategy.ddp.overlap_param_gather = True425 recipe.trainer.strategy.ddp.overlap_grad_reduce = True426 recipe.model.config.use_cpu_initialization = True427 return recipe428 429 430@run.cli.factory(target=llm.train)431def pretrain_ditllama5b_stage3_mock_with_pp() -> run.Partial:432 """MovieGen 30B Stage 3 Training"""433 recipe = pretrain_ditllama5b()434 recipe.data = multimodal_fake_datamodule()435 recipe.data.model_config = recipe.model.config436 recipe.data.seq_length = 8192437 recipe.data.task_encoder.seq_length = 8192438 recipe.data.global_batch_size = 1439 recipe.data.micro_batch_size = 1440 recipe.trainer.strategy.tensor_model_parallel_size = 2441 recipe.trainer.strategy.pipeline_model_parallel_size = 2442 recipe.trainer.strategy.context_parallel_size = 2443 recipe.trainer.strategy.sequence_parallel = True444 recipe.trainer.limit_val_batches = 0445 recipe.trainer.val_check_interval = 1.0446 recipe.data.model_config = recipe.model.config447 recipe.log.log_dir = 'nemo_experiments/ditllama30b_stage5_mock_with_pp'448 return recipe449 450 451@run.cli.factory(target=llm.train)452def pretrain_ditllama30b_stage3_mock_with_pp() -> run.Partial:453 """MovieGen 30B Stage 3 Training with Pipeline Parallelism"""454 recipe = pretrain_ditllama5b()455 recipe.model.config = run.Config(DiTLlama30BConfig)456 recipe.data = multimodal_fake_datamodule()457 recipe.data.model_config = recipe.model.config458 recipe.data.seq_length = 73728459 recipe.data.task_encoder.seq_length = 73728460 recipe.data.global_batch_size = 256461 recipe.data.micro_batch_size = 1462 recipe.trainer.strategy.tensor_model_parallel_size = 4463 recipe.trainer.strategy.pipeline_model_parallel_size = 4464 recipe.trainer.strategy.context_parallel_size = 8465 recipe.trainer.strategy.sequence_parallel = True466 recipe.trainer.limit_val_batches = 0467 recipe.trainer.val_check_interval = 1.0468 recipe.data.model_config = recipe.model.config469 recipe.log.log_dir = 'nemo_experiments/ditllama30b_stage3_mock_with_pp'470 return recipe471 472 473@run.cli.factory(target=llm.train)474def pretrain_ditllama1b() -> run.Partial:475 """MovieGen 1B Stage 1 Training"""476 recipe = pretrain_ditllama5b()477 recipe.model.config = run.Config(DiTLlama1BConfig)478 recipe.data.task_encoder.aethetic_score = 4.0479 recipe.data.seq_length = 256480 recipe.data.task_encoder.seq_length = 256481 recipe.model.config.seq_length = 256482 recipe.data.global_batch_size = 1536483 recipe.data.micro_batch_size = 96484 recipe.trainer.strategy.ddp.overlap_grad_reduce = True485 recipe.log.log_dir = 'nemo_experiments/ditllama1b'486 recipe.trainer.val_check_interval = 3000487 recipe.trainer.callbacks[0].every_n_train_steps = 3000488 recipe.trainer.callbacks[0].monitor = 'global_step'489 recipe.trainer.callbacks[0].save_top_k = 3490 recipe.trainer.callbacks[0].mode = 'max'491 return recipe492 493 494@run.cli.factory(target=llm.train)495def pretrain_ditllama3b() -> run.Partial:496 """MovieGen 3B Stage 1 Training"""497 recipe = pretrain_ditllama1b()498 recipe.data.micro_batch_size = 48499 recipe.model.config = run.Config(500 DiTLlama1BConfig,501 hidden_size=3072,502 num_layers=28,503 num_attention_heads=24,504 ffn_hidden_size=8192,505 )506 recipe.log.log_dir = 'nemo_experiments/ditllama3b'507 508 return recipe509 510 511@run.cli.factory(target=llm.train)512def pretrain_ecditllama1b() -> run.Partial:513 """EC-DiT 1B Training"""514 recipe = pretrain_ditllama1b()515 recipe.data.task_encoder.aethetic_score = 5.0516 recipe.data.micro_batch_size = 72517 recipe.data.global_batch_size = 2304518 recipe.model.config = run.Config(ECDiTLlama1BConfig)519 recipe.log.log_dir = 'nemo_experiments/ecditllama1b'520 recipe.trainer.val_check_interval = 3000521 522 set_use_megatron_fsdp(recipe=recipe)523 recipe.trainer.strategy.ddp.data_parallel_sharding_strategy = 'optim_grads_params'524 recipe.trainer.strategy.ddp.overlap_param_gather = True525 recipe.trainer.strategy.ddp.overlap_grad_reduce = True526 recipe.model.config.use_cpu_initialization = True527 528 return recipe529 530 531@run.cli.factory(target=llm.train)532def dreambooth() -> run.Partial:533 """Dreambooth Fine Tuning"""534 recipe = pretrain()535 recipe.optim.config.lr = 1e-6536 recipe.data = multimodal_datamodule()537 recipe.model.config = run.Config(DiTConfig)538 recipe.trainer.max_steps = 1000539 recipe.trainer.strategy.tensor_model_parallel_size = 8540 recipe.trainer.strategy.sequence_parallel = True541 recipe.resume.restore_config = run.Config(RestoreConfig)542 recipe.resume.resume_if_exists = False543 return recipe544 545 546if __name__ == "__main__":547 OOM_DEBUG = False548 if OOM_DEBUG:549 torch.cuda.memory._record_memory_history(550 True,551 # Keep 100,000 alloc/free events from before the snapshot552 trace_alloc_max_entries=100000,553 # Record stack information for the trace events554 trace_alloc_record_context=True,555 )556 run.cli.main(llm.train, default_factory=dreambooth)557 