| """Wan14B DMD training for Predictor-v4, with optional Full-Generator tuning.""" |
|
|
| from __future__ import annotations |
|
|
| import gc |
| import json |
| import math |
| import os |
| import random |
| import shutil |
| import time |
| from contextlib import contextmanager |
| from pathlib import Path |
| from typing import Any, Iterator |
|
|
| import torch |
| import torch.distributed as dist |
| from omegaconf import OmegaConf |
| from torch.distributed.fsdp import FullyShardedDataParallel as FSDP |
| from torch.nn.parallel import DistributedDataParallel as DDP |
| from torch.optim import AdamW |
| from torch.optim.lr_scheduler import LambdaLR |
| from torch.utils.data import DataLoader |
| from torch.utils.data.distributed import DistributedSampler |
|
|
| from model.dmd import DMD |
| from model.predictor_v4 import SelfForcingPredictorV4 |
| from pipeline.predictor_v4_dmd_training import PredictorV4DMDTrainingPipeline |
| from predictor_training.checkpoint import ( |
| atomic_torch_save, |
| save_predictor_weights, |
| unwrap_model, |
| ) |
| from trainer.predictor_v4_rollout import ( |
| EXPECTED_TIMESTEPS, |
| _configure_predictor_precision, |
| _load_stage1_model_state, |
| _optimizer_groups, |
| ) |
| from utils.dataset import TextDataset |
| from utils.distributed import fsdp_state_dict, fsdp_wrap, launch_distributed_job |
| from utils.misc import set_seed |
|
|
|
|
| def _cosine_with_linear_warmup(step: int, warmup: int, total: int) -> float: |
| if step < warmup: |
| return max(1e-8, float(step + 1) / max(1, warmup)) |
| progress = min(1.0, float(step - warmup) / max(1, total - warmup)) |
| return 0.5 * (1.0 + math.cos(math.pi * progress)) |
|
|
|
|
| def _make_scheduler( |
| optimizer: torch.optim.Optimizer, |
| *, |
| warmup_steps: int, |
| max_steps: int, |
| ) -> LambdaLR: |
| return LambdaLR( |
| optimizer, |
| lambda step: _cosine_with_linear_warmup( |
| step, warmup_steps, max_steps |
| ), |
| ) |
|
|
|
|
| class _LocalEMA: |
| """EMA over the rank-local parameter views, including FSDP shards.""" |
|
|
| def __init__(self, module: torch.nn.Module, decay: float) -> None: |
| self.decay = float(decay) |
| self.parameters = [ |
| parameter for parameter in module.parameters() |
| if parameter.requires_grad |
| ] |
| self.shadow = [ |
| parameter.detach().float().clone() |
| for parameter in self.parameters |
| ] |
|
|
| @torch.no_grad() |
| def update(self) -> None: |
| for shadow, parameter in zip(self.shadow, self.parameters): |
| shadow.mul_(self.decay).add_( |
| parameter.detach().float(), alpha=1.0 - self.decay |
| ) |
|
|
| @contextmanager |
| def apply(self): |
| with torch.no_grad(): |
| backups = [ |
| parameter.detach().clone() for parameter in self.parameters |
| ] |
| try: |
| for parameter, shadow in zip(self.parameters, self.shadow): |
| parameter.copy_(shadow.to(dtype=parameter.dtype)) |
| yield |
| finally: |
| for parameter, backup in zip(self.parameters, backups): |
| parameter.copy_(backup) |
|
|
|
|
| def _distributed_mean(values: torch.Tensor) -> torch.Tensor: |
| dist.all_reduce(values, op=dist.ReduceOp.SUM) |
| return values / dist.get_world_size() |
|
|
|
|
| def _atomic_json(path: Path, payload: dict[str, Any]) -> None: |
| path.parent.mkdir(parents=True, exist_ok=True) |
| temporary = path.with_suffix(path.suffix + f".tmp.{os.getpid()}") |
| temporary.write_text( |
| json.dumps(payload, ensure_ascii=False, indent=2, sort_keys=True, default=str) |
| + "\n", |
| encoding="utf-8", |
| ) |
| os.replace(temporary, path) |
|
|
|
|
| class Trainer: |
| """Distributed Predictor-only or joint Full+Predictor DMD trainer.""" |
|
|
| def __init__(self, config) -> None: |
| self.root_config = config |
| self.cfg = config.predictor_v4_dmd |
| self.mode = str(self.cfg.training_mode).lower() |
| if self.mode not in {"predictor_only", "joint"}: |
| raise ValueError(f"Unknown training_mode {self.mode!r}") |
|
|
| torch.backends.cuda.matmul.allow_tf32 = True |
| torch.backends.cudnn.allow_tf32 = True |
| torch.set_float32_matmul_precision("high") |
| launch_distributed_job() |
| self.rank = dist.get_rank() |
| self.world_size = dist.get_world_size() |
| self.local_rank = int(os.environ["LOCAL_RANK"]) |
| self.device = torch.device("cuda", self.local_rank) |
| self.is_main = self.rank == 0 |
| if self.world_size != int(self.cfg.expected_world_size): |
| raise ValueError( |
| f"Expected {self.cfg.expected_world_size} ranks, got {self.world_size}" |
| ) |
|
|
| seed = int(config.seed) |
| set_seed(seed + self.rank) |
| random.seed(seed + self.rank) |
| self.output_dir = Path(str(self.cfg.output_dir)).resolve() |
| if self.is_main: |
| self.output_dir.mkdir(parents=True, exist_ok=True) |
| dist.barrier() |
| self.log_path = self.output_dir / "train_log.jsonl" |
|
|
| self.model = DMD(config, device=self.device) |
| checkpoint = torch.load( |
| str(self.cfg.full_checkpoint), |
| map_location="cpu", |
| mmap=True, |
| weights_only=False, |
| ) |
| checkpoint_key = str(self.cfg.full_checkpoint_key) |
| if checkpoint_key not in checkpoint: |
| raise KeyError( |
| f"{self.cfg.full_checkpoint} lacks {checkpoint_key!r}" |
| ) |
| self.model.generator.load_state_dict( |
| checkpoint[checkpoint_key], |
| strict=bool(self.cfg.strict_full_load), |
| ) |
| del checkpoint |
|
|
| predictor = SelfForcingPredictorV4.from_teacher( |
| self.model.generator.model, |
| source_block_ids=tuple( |
| int(value) for value in self.cfg.source_block_ids |
| ), |
| ) |
| if self.mode == "predictor_only": |
| self.model.generator.requires_grad_(False) |
| else: |
| self.model.generator.requires_grad_(True) |
|
|
| self.model.generator = fsdp_wrap( |
| self.model.generator, |
| sharding_strategy=str(config.sharding_strategy), |
| mixed_precision=bool(config.mixed_precision), |
| wrap_strategy=str(config.generator_fsdp_wrap_strategy), |
| ) |
| self.model.real_score = fsdp_wrap( |
| self.model.real_score, |
| sharding_strategy=str(config.sharding_strategy), |
| mixed_precision=bool(config.mixed_precision), |
| wrap_strategy=str(config.real_score_fsdp_wrap_strategy), |
| ) |
| self.model.fake_score = fsdp_wrap( |
| self.model.fake_score, |
| sharding_strategy=str(config.sharding_strategy), |
| mixed_precision=bool(config.mixed_precision), |
| wrap_strategy=str(config.fake_score_fsdp_wrap_strategy), |
| ) |
| self.model.text_encoder = fsdp_wrap( |
| self.model.text_encoder, |
| sharding_strategy=str(config.sharding_strategy), |
| mixed_precision=bool(config.mixed_precision), |
| wrap_strategy=str(config.text_encoder_fsdp_wrap_strategy), |
| cpu_offload=bool(getattr(config, "text_encoder_cpu_offload", False)), |
| ) |
|
|
| _configure_predictor_precision(predictor, device=self.device) |
| _load_stage1_model_state( |
| predictor, |
| Path(str(self.cfg.stage1_training_state)).resolve(), |
| expected_step=int(self.cfg.stage1_expected_step), |
| ) |
| predictor.train() |
| self.predictor = DDP( |
| predictor, |
| device_ids=[self.local_rank], |
| output_device=self.local_rank, |
| broadcast_buffers=False, |
| gradient_as_bucket_view=True, |
| find_unused_parameters=False, |
| ) |
|
|
| actual_timesteps = self.model.denoising_step_list.detach().float().cpu() |
| if not torch.equal(actual_timesteps, EXPECTED_TIMESTEPS): |
| raise ValueError( |
| "Predictor DMD requires exact warped timesteps " |
| f"{EXPECTED_TIMESTEPS.tolist()}, got {actual_timesteps.tolist()}" |
| ) |
| forced_exit = getattr(self.cfg, "forced_exit_step", None) |
| self.rollout = PredictorV4DMDTrainingPipeline( |
| denoising_step_list=self.model.denoising_step_list, |
| scheduler=self.model.scheduler, |
| generator=self.model.generator, |
| predictor=self.predictor, |
| training_mode=self.mode, |
| context_noise=int(config.context_noise), |
| forced_exit_step=( |
| None if forced_exit is None else int(forced_exit) |
| ), |
| ) |
| self.model.inference_pipeline = self.rollout |
|
|
| fusion, blocks, group_names = _optimizer_groups( |
| unwrap_model(self.predictor) |
| ) |
| self.predictor_optimizer = AdamW( |
| [ |
| { |
| "params": fusion, |
| "lr": float(self.cfg.predictor_fusion_lr), |
| "name": "fusion_residual", |
| }, |
| { |
| "params": blocks, |
| "lr": float(self.cfg.predictor_blocks_lr), |
| "name": "blocks", |
| }, |
| ], |
| betas=( |
| float(self.cfg.student_beta1), |
| float(self.cfg.student_beta2), |
| ), |
| weight_decay=float(self.cfg.weight_decay), |
| ) |
| self.predictor_scheduler = _make_scheduler( |
| self.predictor_optimizer, |
| warmup_steps=int(self.cfg.warmup_steps), |
| max_steps=int(self.cfg.target_predictor_updates), |
| ) |
| self.full_optimizer = None |
| self.full_scheduler = None |
| if self.mode == "joint": |
| self.full_optimizer = AdamW( |
| [ |
| parameter |
| for parameter in self.model.generator.parameters() |
| if parameter.requires_grad |
| ], |
| lr=float(self.cfg.full_lr), |
| betas=( |
| float(self.cfg.student_beta1), |
| float(self.cfg.student_beta2), |
| ), |
| weight_decay=float(self.cfg.weight_decay), |
| ) |
| self.full_scheduler = _make_scheduler( |
| self.full_optimizer, |
| warmup_steps=int(self.cfg.warmup_steps), |
| max_steps=int(self.cfg.max_student_steps), |
| ) |
| self.critic_optimizer = AdamW( |
| [ |
| parameter |
| for parameter in self.model.fake_score.parameters() |
| if parameter.requires_grad |
| ], |
| lr=float(self.cfg.fake_score_lr), |
| betas=( |
| float(self.cfg.critic_beta1), |
| float(self.cfg.critic_beta2), |
| ), |
| weight_decay=float(self.cfg.weight_decay), |
| ) |
|
|
| dataset = TextDataset(prompt_path=str(self.cfg.data_path)) |
| self.sampler = DistributedSampler( |
| dataset, |
| num_replicas=self.world_size, |
| rank=self.rank, |
| shuffle=True, |
| seed=seed, |
| drop_last=True, |
| ) |
| self.loader = DataLoader( |
| dataset, |
| batch_size=int(config.batch_size), |
| sampler=self.sampler, |
| num_workers=int(self.cfg.num_workers), |
| pin_memory=bool(self.cfg.pin_memory), |
| drop_last=True, |
| ) |
| self.data_iterator: Iterator[dict[str, Any]] | None = None |
| self.data_epoch = 0 |
|
|
| self.student_step = 0 |
| self.predictor_step = 0 |
| self.critic_step = 0 |
| self.predictor_ema: _LocalEMA | None = None |
| self.full_ema: _LocalEMA | None = None |
| self.unconditional_dict: dict[str, torch.Tensor] | None = None |
| self.run_config = { |
| **OmegaConf.to_container(self.cfg, resolve=True), |
| "world_size": self.world_size, |
| "effective_global_batch": int(config.batch_size) * self.world_size, |
| "parameter_groups": group_names, |
| "rollout": ( |
| "random_exit_P1_P2_P3" |
| if self.mode == "predictor_only" |
| else "random_exit_F0_P1_P2_P3" |
| ), |
| "dmd_frames": 21, |
| "clean_context_grad": False, |
| } |
| if self.is_main: |
| _atomic_json(self.output_dir / "train_config.json", self.run_config) |
| self.swanlab_run = self._initialize_swanlab() |
|
|
| def _initialize_swanlab(self): |
| if not self.is_main or not bool(self.cfg.use_swanlab): |
| return None |
| import swanlab |
|
|
| mode = str(self.cfg.swanlab_mode) |
| if mode == "cloud": |
| api_key = os.environ.get("SWANLAB_API_KEY") |
| if api_key: |
| swanlab.login(api_key=api_key, save=False) |
| else: |
| swanlab.login() |
| workspace = self.cfg.swanlab_workspace |
| return swanlab.init( |
| project=str(self.cfg.swanlab_project), |
| workspace=None if workspace is None else str(workspace), |
| experiment_name=str(self.cfg.swanlab_experiment), |
| description=str(self.cfg.swanlab_description), |
| tags=list(self.cfg.swanlab_tags), |
| config=self.run_config, |
| logdir=str(self.output_dir / "swanlab"), |
| mode=mode, |
| ) |
|
|
| def _next_batch(self) -> dict[str, Any]: |
| if self.data_iterator is None: |
| self.sampler.set_epoch(self.data_epoch) |
| self.data_iterator = iter(self.loader) |
| try: |
| return next(self.data_iterator) |
| except StopIteration: |
| self.data_epoch += 1 |
| self.sampler.set_epoch(self.data_epoch) |
| self.data_iterator = iter(self.loader) |
| return next(self.data_iterator) |
|
|
| @torch.no_grad() |
| def _conditional_dicts( |
| self, prompts: list[str] |
| ) -> tuple[dict[str, torch.Tensor], dict[str, torch.Tensor]]: |
| conditional = self.model.text_encoder(text_prompts=prompts) |
| if self.unconditional_dict is None: |
| unconditional = self.model.text_encoder( |
| text_prompts=[str(self.root_config.negative_prompt)] |
| * len(prompts) |
| ) |
| self.unconditional_dict = { |
| key: value.detach() for key, value in unconditional.items() |
| } |
| return conditional, self.unconditional_dict |
|
|
| def _image_shape(self, batch_size: int) -> list[int]: |
| shape = list(self.root_config.image_or_video_shape) |
| shape[0] = int(batch_size) |
| return shape |
|
|
| def _critic_update(self) -> dict[str, float]: |
| batch = self._next_batch() |
| prompts = list(batch["prompts"]) |
| conditional, unconditional = self._conditional_dicts(prompts) |
| self.critic_optimizer.zero_grad(set_to_none=True) |
| loss, logs = self.model.critic_loss( |
| image_or_video_shape=self._image_shape(len(prompts)), |
| conditional_dict=conditional, |
| unconditional_dict=unconditional, |
| clean_latent=None, |
| initial_latent=None, |
| ) |
| loss.backward() |
| grad_norm = self.model.fake_score.clip_grad_norm_( |
| float(self.cfg.critic_grad_clip) |
| ) |
| if not torch.isfinite(grad_norm): |
| raise FloatingPointError(f"Non-finite Fake Score grad norm {grad_norm}") |
| self.critic_optimizer.step() |
| self.critic_step += 1 |
| result = { |
| "critic_loss": float(loss.detach()), |
| "critic_grad_norm": float(grad_norm), |
| } |
| del batch, conditional, loss, logs |
| return result |
|
|
| def _maybe_initialize_ema(self) -> None: |
| if ( |
| self.predictor_ema is None |
| and self.predictor_step >= int(self.cfg.ema_start_step) |
| ): |
| self.predictor_ema = _LocalEMA( |
| unwrap_model(self.predictor), float(self.cfg.ema_decay) |
| ) |
| if ( |
| self.mode == "joint" |
| and self.full_ema is None |
| and self.student_step >= int(self.cfg.ema_start_step) |
| ): |
| self.full_ema = _LocalEMA( |
| self.model.generator, float(self.cfg.ema_decay) |
| ) |
|
|
| @torch.no_grad() |
| def _sync_predictor_frozen_from_full(self) -> None: |
| if self.mode != "joint": |
| return |
| with FSDP.summon_full_params( |
| self.model.generator, |
| recurse=True, |
| writeback=False, |
| rank0_only=False, |
| ): |
| unwrap_model(self.predictor).sync_frozen_from_teacher( |
| self.model.generator.module.model |
| ) |
|
|
| def _student_update(self) -> dict[str, float]: |
| batch = self._next_batch() |
| prompts = list(batch["prompts"]) |
| conditional, unconditional = self._conditional_dicts(prompts) |
| self.predictor_optimizer.zero_grad(set_to_none=True) |
| if self.full_optimizer is not None: |
| self.full_optimizer.zero_grad(set_to_none=True) |
|
|
| loss, logs = self.model.generator_loss( |
| image_or_video_shape=self._image_shape(len(prompts)), |
| conditional_dict=conditional, |
| unconditional_dict=unconditional, |
| clean_latent=None, |
| initial_latent=None, |
| ) |
| exit_step = int(self.rollout.last_exit_step) |
| loss.backward() |
|
|
| predictor_updated = exit_step > 0 |
| predictor_grad_norm = torch.zeros((), device=self.device) |
| if predictor_updated: |
| predictor_grad_norm = torch.nn.utils.clip_grad_norm_( |
| unwrap_model(self.predictor).parameters(), |
| float(self.cfg.predictor_grad_clip), |
| ) |
| if not torch.isfinite(predictor_grad_norm): |
| raise FloatingPointError( |
| f"Non-finite Predictor grad norm {predictor_grad_norm}" |
| ) |
| self.predictor_optimizer.step() |
| self.predictor_scheduler.step() |
| self.predictor_step += 1 |
|
|
| full_grad_norm = torch.zeros((), device=self.device) |
| if self.full_optimizer is not None: |
| full_grad_norm = self.model.generator.clip_grad_norm_( |
| float(self.cfg.full_grad_clip) |
| ) |
| if not torch.isfinite(full_grad_norm): |
| raise FloatingPointError( |
| f"Non-finite Full grad norm {full_grad_norm}" |
| ) |
| self.full_optimizer.step() |
| self.full_scheduler.step() |
| self._sync_predictor_frozen_from_full() |
|
|
| self.student_step += 1 |
| self._maybe_initialize_ema() |
| if predictor_updated and self.predictor_ema is not None: |
| self.predictor_ema.update() |
| if self.full_optimizer is not None and self.full_ema is not None: |
| self.full_ema.update() |
|
|
| result = { |
| "dmd_loss": float(loss.detach()), |
| "dmd_gradient_norm": float(logs["dmdtrain_gradient_norm"]), |
| "dmd_score_timestep": float(logs["timestep"].float().mean()), |
| "exit_step": float(exit_step), |
| "predictor_updated": float(predictor_updated), |
| "predictor_grad_norm": float(predictor_grad_norm), |
| "full_grad_norm": float(full_grad_norm), |
| } |
| del batch, conditional, loss, logs |
| return result |
|
|
| def _checkpoint_metadata(self) -> dict[str, Any]: |
| model = unwrap_model(self.predictor) |
| return { |
| "source_block_ids": list(model.source_block_ids), |
| "student_step": self.student_step, |
| "predictor_step": self.predictor_step, |
| "training_mode": self.mode, |
| "training_rollout": ( |
| "P1_P2_P3" if self.mode == "predictor_only" else "F0_P1_P2_P3" |
| ), |
| "teacher_checkpoint_key": "generator_ema", |
| "predictor_config": model.config_dict, |
| } |
|
|
| def _save(self, *, final: bool = False) -> None: |
| dist.barrier() |
| suffix = "final" if final else f"step_{self.student_step:05d}" |
| checkpoint_dir = self.output_dir / f"checkpoint_{suffix}" |
| if self.is_main: |
| checkpoint_dir.mkdir(parents=True, exist_ok=True) |
| save_predictor_weights( |
| self.predictor, |
| checkpoint_dir / "predictor.safetensors", |
| metadata=self._checkpoint_metadata(), |
| ) |
| if self.predictor_ema is not None: |
| with self.predictor_ema.apply(): |
| save_predictor_weights( |
| self.predictor, |
| checkpoint_dir / "predictor_ema.safetensors", |
| metadata={ |
| **self._checkpoint_metadata(), |
| "ema_decay": float(self.cfg.ema_decay), |
| }, |
| ) |
| dist.barrier() |
|
|
| critic_state = fsdp_state_dict(self.model.fake_score) |
| generator_state = None |
| generator_ema_state = None |
| if self.mode == "joint": |
| generator_state = fsdp_state_dict(self.model.generator) |
| if self.full_ema is not None: |
| with self.full_ema.apply(): |
| generator_ema_state = fsdp_state_dict(self.model.generator) |
| else: |
| generator_ema_state = generator_state |
| if self.is_main: |
| payload: dict[str, Any] = { |
| "critic": critic_state, |
| "student_step": self.student_step, |
| "predictor_step": self.predictor_step, |
| "training_mode": self.mode, |
| } |
| if generator_state is not None: |
| payload["generator"] = generator_state |
| payload["generator_ema"] = generator_ema_state |
| atomic_torch_save(payload, checkpoint_dir / "model.pt") |
| _atomic_json( |
| checkpoint_dir / "state.json", |
| { |
| "student_step": self.student_step, |
| "predictor_step": self.predictor_step, |
| "critic_step": self.critic_step, |
| "training_mode": self.mode, |
| "final": final, |
| }, |
| ) |
| latest = self.output_dir / "latest" |
| temporary = self.output_dir / f".latest.{os.getpid()}" |
| if temporary.exists() or temporary.is_symlink(): |
| temporary.unlink() |
| temporary.symlink_to(checkpoint_dir.name) |
| os.replace(temporary, latest) |
| keep = int(self.cfg.keep_checkpoints) |
| snapshots = sorted( |
| path for path in self.output_dir.glob("checkpoint_step_*") |
| if path.is_dir() |
| ) |
| for old in snapshots[:-keep] if keep > 0 else snapshots: |
| shutil.rmtree(old) |
| dist.barrier() |
|
|
| def train(self) -> None: |
| max_student_steps = int(self.cfg.max_student_steps) |
| critic_updates = int(self.cfg.critic_updates_per_student) |
| exit_counts = torch.zeros(4, device=self.device, dtype=torch.float64) |
| try: |
| while self.student_step < max_student_steps: |
| started = time.perf_counter() |
| critic_loss_sum = 0.0 |
| critic_grad_sum = 0.0 |
| for _ in range(critic_updates): |
| metrics = self._critic_update() |
| critic_loss_sum += metrics["critic_loss"] |
| critic_grad_sum += metrics["critic_grad_norm"] |
| student = self._student_update() |
| exit_counts[int(student["exit_step"])] += 1 |
| values = torch.tensor( |
| [ |
| student["dmd_loss"], |
| student["dmd_gradient_norm"], |
| student["dmd_score_timestep"], |
| student["predictor_grad_norm"], |
| student["full_grad_norm"], |
| critic_loss_sum / critic_updates, |
| critic_grad_sum / critic_updates, |
| time.perf_counter() - started, |
| ], |
| device=self.device, |
| dtype=torch.float64, |
| ) |
| averaged = _distributed_mean(values) |
| global_exit_counts = exit_counts.clone() |
| dist.all_reduce(global_exit_counts, op=dist.ReduceOp.SUM) |
| global_exit_counts /= self.world_size |
|
|
| if ( |
| self.student_step == 1 |
| or self.student_step % int(self.cfg.log_every) == 0 |
| ): |
| record = { |
| "student_step": self.student_step, |
| "predictor_step": self.predictor_step, |
| "critic_step": self.critic_step, |
| "training_mode": self.mode, |
| "exit_step": int(student["exit_step"]), |
| "dmd_loss": float(averaged[0]), |
| "dmd_gradient_norm": float(averaged[1]), |
| "dmd_score_timestep": float(averaged[2]), |
| "predictor_grad_norm": float(averaged[3]), |
| "full_grad_norm": float(averaged[4]), |
| "critic_loss": float(averaged[5]), |
| "critic_grad_norm": float(averaged[6]), |
| "step_time_s": float(averaged[7]), |
| "lr_predictor_fusion": self.predictor_optimizer.param_groups[0]["lr"], |
| "lr_predictor_blocks": self.predictor_optimizer.param_groups[1]["lr"], |
| "lr_full": ( |
| 0.0 |
| if self.full_optimizer is None |
| else self.full_optimizer.param_groups[0]["lr"] |
| ), |
| "lr_fake_score": self.critic_optimizer.param_groups[0]["lr"], |
| "exit_count_f0": int(global_exit_counts[0]), |
| "exit_count_p1": int(global_exit_counts[1]), |
| "exit_count_p2": int(global_exit_counts[2]), |
| "exit_count_p3": int(global_exit_counts[3]), |
| "peak_memory_gib": ( |
| torch.cuda.max_memory_allocated(self.device) / 2**30 |
| ), |
| } |
| if self.is_main: |
| with self.log_path.open("a", encoding="utf-8") as handle: |
| handle.write(json.dumps(record, sort_keys=True) + "\n") |
| print(json.dumps(record, sort_keys=True), flush=True) |
| if self.swanlab_run is not None: |
| import swanlab |
|
|
| swanlab.log( |
| { |
| key: value |
| for key, value in record.items() |
| if not isinstance(value, str) |
| }, |
| step=self.student_step, |
| ) |
|
|
| should_save = ( |
| not bool(self.root_config.no_save) |
| and ( |
| self.student_step % int(self.cfg.save_every) == 0 |
| or self.student_step == max_student_steps |
| ) |
| ) |
| if should_save: |
| self._save(final=self.student_step == max_student_steps) |
| del values, averaged, global_exit_counts, student |
| if self.student_step % int(self.root_config.gc_interval) == 0: |
| gc.collect() |
| torch.cuda.empty_cache() |
| except BaseException as error: |
| if self.swanlab_run is not None: |
| try: |
| import swanlab |
|
|
| swanlab.finish(error=str(error)) |
| except Exception: |
| pass |
| raise |
| else: |
| if self.swanlab_run is not None: |
| import swanlab |
|
|
| swanlab.finish() |
| finally: |
| gc.collect() |
| torch.cuda.empty_cache() |
| dist.destroy_process_group() |
|
|
|
|
| __all__ = ["Trainer"] |
|
|