CoolFace
Modelpublic

RASMUS/Finnish-ASR-Canary-v2

sourceHugging Facemitupdated 7mo agoView on Hugging Face
0likes2.2kdownloads
dit_train.py557 linesDownload Raw Back to dit
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