Spaces:
Running on L40S
Running on L40S
| # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. | |
| # SPDX-License-Identifier: OpenMDW-1.1 | |
| """Dataloader config options.""" | |
| from hydra.core.config_store import ConfigStore | |
| from cosmos_framework.callbacks.manual_gc import ManualGarbageCollection | |
| from cosmos_framework.utils.lazy_config import PLACEHOLDER | |
| from cosmos_framework.utils.lazy_config import LazyCall as L | |
| from cosmos_framework.utils.callback import LowPrecisionCallback, WandBCallback | |
| from cosmos_framework.callbacks.compile_tokenizer import CompileTokenizer | |
| from cosmos_framework.callbacks.device_monitor import DeviceMonitor | |
| from cosmos_framework.callbacks.every_n_draw_sample import EveryNDrawSample | |
| from cosmos_framework.callbacks.expert_heatmap import ExpertHeatmap | |
| from cosmos_framework.callbacks.grad_clip import GradClip | |
| from cosmos_framework.callbacks.heart_beat import HeartBeat | |
| from cosmos_framework.callbacks.iter_speed import IterSpeed | |
| from cosmos_framework.callbacks.load_pretrained import LoadPretrained | |
| from cosmos_framework.callbacks.mfu import MFUCallback | |
| from cosmos_framework.callbacks.moe_specialization_callback import MoESpecializationCallback | |
| from cosmos_framework.callbacks.moe_stability_callback import MoEStabilityCallback | |
| from cosmos_framework.callbacks.norm_monitor import NormMonitor | |
| from cosmos_framework.callbacks.ofu import OFUCallback | |
| from cosmos_framework.callbacks.param_count import ParamCount | |
| from cosmos_framework.callbacks.sequence_packing_padding import SequencePackingPadding | |
| from cosmos_framework.callbacks.sigma_loss_analysis import SigmaLossAnalysis | |
| from cosmos_framework.callbacks.skip_nan_step import SkipNaNStep | |
| from cosmos_framework.callbacks.termination_signal_checkpoint import TerminationSignalCheckpoint | |
| from cosmos_framework.callbacks.training_stats import TrainingStatsCallback | |
| from cosmos_framework.callbacks.wandb_log import WandbCallback as WandBCallbackMultiplier | |
| from cosmos_framework.callbacks.wandb_log_eval import WandbCallback as WandBCallbackEval | |
| BASIC_CALLBACKS = dict( | |
| iter_speed=L(IterSpeed)( # does not use model or optimizer | |
| every_n="${trainer.logging_iter}", | |
| save_s3="${upload_reproducible_setup}", | |
| save_s3_every_log_n=500, | |
| hit_thres=50, | |
| ), | |
| manual_gc=L(ManualGarbageCollection)(every_n=5), # does not use model or optimizer | |
| wandb=L(WandBCallback)(), | |
| wandb_2x=L(WandBCallbackMultiplier)( | |
| logging_iter_multipler=2, | |
| save_logging_iter_multipler=1, | |
| save_s3="${upload_reproducible_setup}", | |
| ), | |
| param_count=L(ParamCount)( # use model | |
| save_s3="${upload_reproducible_setup}", | |
| ), | |
| wandb_val=L(WandBCallbackEval)( | |
| save_s3="${upload_reproducible_setup}", | |
| ), | |
| moe_stability=L(MoEStabilityCallback)(every_n=250), | |
| moe_specialization=L(MoESpecializationCallback)(every_n=250), | |
| expert_heatmap=L(ExpertHeatmap)(), | |
| load_pretrained=L(LoadPretrained)(), | |
| compile_tokenizer=L(CompileTokenizer)(enabled=False, compile_after_iterations=3), | |
| norm_monitor=L(NormMonitor)( | |
| every_n=5000, | |
| log_stat_wandb=True, | |
| save_s3="${upload_reproducible_setup}", | |
| track_activations=True, | |
| ), | |
| sigma_loss_analysis=L(SigmaLossAnalysis)( | |
| every_n=5000, | |
| every_n_viz=5000, | |
| save_s3="${upload_reproducible_setup}", | |
| ), | |
| sequence_packing_padding=L(SequencePackingPadding)(every_n="${trainer.logging_iter}"), | |
| mfu=L(MFUCallback)(every_n="${trainer.logging_iter}", grad_accum_iter="${trainer.grad_accum_iter}"), | |
| ofu=L(OFUCallback)(every_n="${trainer.logging_iter}"), | |
| ) | |
| JOB_MONITOR_CALLBACKS = dict( | |
| heart_beat=L(HeartBeat)( | |
| every_n=200, | |
| update_interval_in_minute=20, | |
| save_s3="${upload_reproducible_setup}", | |
| ), | |
| device_monitor=L(DeviceMonitor)( | |
| every_n=200, | |
| save_s3="${upload_reproducible_setup}", | |
| upload_every_n_mul=5, | |
| ), | |
| termination_signal_checkpoint=L(TerminationSignalCheckpoint)( | |
| min_save_fraction=1 / 3, | |
| ), | |
| ) | |
| OPTIMIZATION_CALLBACKS = dict( | |
| skip_nan_step=L(SkipNaNStep)(max_consecutive_nan=100), | |
| grad_clip=L(GradClip)(clip_norm=1.0, track_per_modality=True), # image/video grad-norm split | |
| low_precision=L(LowPrecisionCallback)(update_iter=1, config=PLACEHOLDER, trainer=PLACEHOLDER), # use model | |
| ) | |
| VIZ_ONLINE_SAMPLING_CALLBACKS = dict( | |
| every_n_sample_reg=L(EveryNDrawSample)( | |
| every_n=5000, | |
| save_s3=True, | |
| do_x0_prediction=False, | |
| ), | |
| every_n_sample_ema=L(EveryNDrawSample)( | |
| every_n=5000, | |
| is_ema=True, | |
| save_s3=True, | |
| do_x0_prediction=False, | |
| ), | |
| ) | |
| def register_callbacks(): | |
| cs = ConfigStore.instance() | |
| cs.store(group="callbacks", package="trainer.callbacks", name="basic", node=BASIC_CALLBACKS) | |
| cs.store(group="callbacks", package="trainer.callbacks", name="job_monitor", node=JOB_MONITOR_CALLBACKS) | |
| cs.store(group="callbacks", package="trainer.callbacks", name="optimization", node=OPTIMIZATION_CALLBACKS) | |
| # Online sampling generation callback | |
| cs.store( | |
| group="callbacks", package="trainer.callbacks", name="viz_online_sampling", node=VIZ_ONLINE_SAMPLING_CALLBACKS | |
| ) | |
| # Register "generation" as alias for "viz_online_sampling" (expected by base config.py defaults) | |
| cs.store(group="callbacks", package="trainer.callbacks", name="generation", node=VIZ_ONLINE_SAMPLING_CALLBACKS) | |
| TRAINING_STATS_CALLBACKS = dict( | |
| training_stats=L(TrainingStatsCallback)( | |
| log_freq=100, | |
| ) | |
| ) | |
| cs.store(group="callbacks", package="trainer.callbacks", name="training_stats", node=TRAINING_STATS_CALLBACKS) | |