from __future__ import annotations import argparse import sys import time from pathlib import Path import torch from tqdm import tqdm try: from .common import ( append_csv_row, build_dataloader, build_dataset, build_model, build_optimizer, build_scheduler, is_ram_chunk_dataset, load_config, pack_inputs, resume_full_checkpoint, save_epoch_checkpoints, set_seed, shutdown_dataloader, ) from .logger import ExperimentLogger from .loss import build_loss except ImportError: code_root = Path(__file__).resolve().parents[2] if str(code_root) not in sys.path: sys.path.insert(0, str(code_root)) from src.training_validation.common import ( # type: ignore append_csv_row, build_dataloader, build_dataset, build_model, build_optimizer, build_scheduler, is_ram_chunk_dataset, load_config, pack_inputs, resume_full_checkpoint, save_epoch_checkpoints, set_seed, shutdown_dataloader, ) from src.training_validation.logger import ExperimentLogger # type: ignore from src.training_validation.loss import build_loss # type: ignore def train_one_epoch( model: torch.nn.Module, loader, criterion: torch.nn.Module, optimizer: torch.optim.Optimizer, device: torch.device, input_sources: list[str], grad_clip_norm: float | None = None, gradient_accumulation_steps: int = 1, preload_after_iter=None, ) -> dict: if gradient_accumulation_steps < 1: raise ValueError( f"gradient_accumulation_steps must be >= 1, got {gradient_accumulation_steps}" ) model.train() total_loss = 0.0 total_samples = 0 skipped_batches = 0 pending_micro_batches = 0 optimizer_steps = 0 component_totals: dict[str, float] = {} iterator = iter(loader) if preload_after_iter is not None: preload_after_iter() progress = tqdm(iterator, total=len(loader), desc="train", dynamic_ncols=True) optimizer.zero_grad(set_to_none=True) for batch in progress: if batch is None: skipped_batches += 1 continue x = pack_inputs(batch, input_sources, device) y = { key: value.to(device=device, dtype=torch.float32, non_blocking=True) for key, value in batch["labels"].items() if value is not None } outputs = model(x) loss = criterion(outputs, y) (loss / gradient_accumulation_steps).backward() pending_micro_batches += 1 if pending_micro_batches == gradient_accumulation_steps: if grad_clip_norm is not None and grad_clip_norm > 0: torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_norm) optimizer.step() optimizer.zero_grad(set_to_none=True) optimizer_steps += 1 pending_micro_batches = 0 batch_size = int(x.shape[0]) total_samples += batch_size total_loss += float(loss.detach().cpu()) * batch_size for name, value in getattr(criterion, "last_components", {}).items(): component_totals[name] = component_totals.get(name, 0.0) + float(value) * batch_size progress.set_postfix(loss=total_loss / max(total_samples, 1)) # Preserve the mean-gradient scale for a final incomplete accumulation group. if pending_micro_batches > 0: correction = gradient_accumulation_steps / pending_micro_batches for parameter in model.parameters(): if parameter.grad is not None: parameter.grad.mul_(correction) if grad_clip_norm is not None and grad_clip_norm > 0: torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_norm) optimizer.step() optimizer.zero_grad(set_to_none=True) optimizer_steps += 1 summary = { "loss": total_loss / max(total_samples, 1), "samples": int(total_samples), "skipped_batches": int(skipped_batches), "gradient_accumulation_steps": int(gradient_accumulation_steps), "optimizer_steps": int(optimizer_steps), } for name, value in component_totals.items(): summary[f"loss_{name}"] = value / max(total_samples, 1) return summary def _resume_ram_chunk_id(resume_path: Path, start_epoch: int, num_chunks: int) -> int: fallback = int(start_epoch) % int(num_chunks) if start_epoch <= 0: return 0 try: payload = torch.load(resume_path, map_location="cpu", weights_only=True) except Exception as exc: print(f"RAM chunk resume fallback to epoch modulo: failed to read {resume_path}: {exc}", flush=True) return fallback if not isinstance(payload, dict): return fallback train_summary = payload.get("train_summary") if not isinstance(train_summary, dict): return fallback for key in ("current_chunk_id_after_swap", "chunk_id"): value = train_summary.get(key) if value is None: continue chunk_id = int(value) if chunk_id >= 0: return chunk_id % int(num_chunks) return fallback def main() -> None: parser = argparse.ArgumentParser(description="Train a CI model with BasicDataset or FastDataset.") parser.add_argument("--config", required=True, help="Experiment YAML path") parser.add_argument("--device", default=None, help="Override device, e.g. cuda:0 or cpu") parser.add_argument("--output-dir", default=None, help="Override training output directory") args = parser.parse_args() config = load_config(args.config) if args.output_dir is not None: config["output_dir"] = str(Path(args.output_dir).resolve()) seed_cfg = dict(config.get("seed", {})) set_seed(int(seed_cfg.get("value", 42)), deterministic=bool(seed_cfg.get("deterministic", True))) requested_device = str(args.device or config.get("device") or "auto") if requested_device == "auto": requested_device = "cuda" if torch.cuda.is_available() else "cpu" device = torch.device(requested_device) train_cfg = dict(config.get("train", {})) input_sources = list(train_cfg.get("input_sources", config.get("input_sources", config.get("required_inputs", ["concat"])))) label_key = str(train_cfg.get("label_key", train_cfg.get("target_label", "ci"))) loss_cfg = dict(config.get("loss", {"name": "binary_focal"})) loss_required_labels = _required_labels_from_loss(loss_cfg) config.setdefault("train", {}) config["train"].setdefault("input_sources", input_sources) if loss_required_labels and ("losses" in loss_cfg or "required_labels" not in config["train"]): config["train"]["required_labels"] = loss_required_labels else: config["train"].setdefault("required_labels", [label_key]) config["_defer_ram_chunk_initial_load"] = True dataset = build_dataset(config, split=str(train_cfg.get("split", "train")), mode="train") model = build_model(config).to(device) if "losses" not in loss_cfg: loss_cfg.setdefault("label_key", label_key) criterion = build_loss(loss_cfg).to(device) optimizer = build_optimizer(config, model) scheduler = build_scheduler(config, optimizer) out_dir = Path(config.get("output_dir", config.get("checkpoint_dir", "runs/default"))) log_path = out_dir / "train_log.csv" epochs = int(train_cfg.get("epochs", config.get("epochs", 1))) resume_cfg_path = train_cfg.get("resume_path") resume_path = Path(resume_cfg_path) if resume_cfg_path else out_dir / "checkpoints" / "latest_full.pt" start_epoch = 0 if bool(train_cfg.get("resume", True)): start_epoch = resume_full_checkpoint(resume_path, model, optimizer, scheduler, device, criterion=criterion) if start_epoch > 0: print(f"resume from {resume_path} | start_epoch={start_epoch}", flush=True) else: print(f"resume skip | no checkpoint: {resume_path}", flush=True) if is_ram_chunk_dataset(dataset) and dataset.num_chunks > 0: # type: ignore[attr-defined] initial_chunk_id = _resume_ram_chunk_id(resume_path, start_epoch, int(dataset.num_chunks)) # type: ignore[attr-defined] print(f"RAM chunk initial load: chunk {initial_chunk_id}/{int(dataset.num_chunks) - 1}", flush=True) # type: ignore[attr-defined] dataset.load_chunk_sync(initial_chunk_id, free_current_before_load=True) # type: ignore[attr-defined] loader = build_dataloader(config, dataset, mode="train") grad_clip_norm = train_cfg.get("grad_clip_norm", config.get("grad_clip_norm")) grad_clip_norm = None if grad_clip_norm is None else float(grad_clip_norm) gradient_accumulation_steps = int(train_cfg.get("gradient_accumulation_steps", 1)) if gradient_accumulation_steps < 1: raise ValueError( "train.gradient_accumulation_steps must be >= 1, " f"got {gradient_accumulation_steps}" ) print( "training batch configuration: " f"physical_batch_size={int(train_cfg.get('batch_size', 1))}, " f"gradient_accumulation_steps={gradient_accumulation_steps}, " f"effective_batch_size=" f"{int(train_cfg.get('batch_size', 1)) * gradient_accumulation_steps}", flush=True, ) logger = ExperimentLogger(config, mode="train") logger.start() try: for epoch in range(start_epoch + 1, epochs + 1): start = time.time() chunk_id_before = getattr(dataset, "current_chunk_id", None) preload_status_before = ( dataset.get_preload_status() if is_ram_chunk_dataset(dataset) else {} # type: ignore[attr-defined] ) def _start_next_chunk_preload() -> None: if not is_ram_chunk_dataset(dataset): return if dataset.num_chunks <= 1: # type: ignore[attr-defined] return next_chunk = (int(dataset.current_chunk_id) + 1) % int(dataset.num_chunks) # type: ignore[attr-defined] dataset.start_preload(next_chunk) # type: ignore[attr-defined] summary = train_one_epoch( model=model, loader=loader, criterion=criterion, optimizer=optimizer, device=device, input_sources=input_sources, grad_clip_norm=grad_clip_norm, gradient_accumulation_steps=gradient_accumulation_steps, preload_after_iter=_start_next_chunk_preload if is_ram_chunk_dataset(dataset) else None, ) if scheduler is not None: if isinstance(scheduler, torch.optim.lr_scheduler.ReduceLROnPlateau): scheduler.step(summary["loss"]) else: scheduler.step() lr = float(optimizer.param_groups[0]["lr"]) swapped = False if is_ram_chunk_dataset(dataset) and dataset.num_chunks > 1: # type: ignore[attr-defined] swapped = bool(dataset.swap_if_preload_ready()) # type: ignore[attr-defined] if swapped: shutdown_dataloader(loader) loader = build_dataloader(config, dataset, mode="train") preload_status_after = ( dataset.get_preload_status() if is_ram_chunk_dataset(dataset) else {} # type: ignore[attr-defined] ) row = { "epoch": int(epoch), **summary, "lr": lr, "seconds": time.time() - start, } if is_ram_chunk_dataset(dataset): row.update( { "chunk_id": -1 if chunk_id_before is None else int(chunk_id_before), "num_chunks": int(dataset.num_chunks), # type: ignore[attr-defined] "chunk_samples": int(summary["samples"]), "preload_running": bool(preload_status_after.get("preload_running", False)), "preload_ready": bool(preload_status_after.get("preload_ready", False)), "preload_chunk_id": int(preload_status_after.get("preload_chunk_id", -1)), "preload_ready_before": bool(preload_status_before.get("preload_ready", False)), "swapped": bool(swapped), "current_chunk_id_after_swap": int(dataset.current_chunk_id), # type: ignore[attr-defined] } ) append_csv_row(log_path, row) logger.log(row, step=epoch, prefix="train") save_epoch_checkpoints(config, model, optimizer, scheduler, epoch, row, criterion=criterion) print(f"epoch {epoch:04d}: loss={row['loss']:.6g}, samples={row['samples']}, lr={lr:.3g}") finally: shutdown_dataloader(loader) if is_ram_chunk_dataset(dataset): dataset.shutdown_preload() # type: ignore[attr-defined] logger.finish() def _required_labels_from_loss(loss_cfg: dict) -> list[str]: if "losses" in loss_cfg: labels = [] for item in dict(loss_cfg["losses"]).values(): label_key = str(dict(item)["label_key"]) if label_key not in labels: labels.append(label_key) return labels label_key = loss_cfg.get("label_key") or loss_cfg.get("target_label") return [str(label_key)] if label_key else [] if __name__ == "__main__": main()