Download code/training/src/training_validation/train.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 13.7 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/train.py
- Command line
-
hf download hf://lsh9034/ci-net/code/training/src/training_validation/train.py
-
curl -L -o train.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/train.py
13.7 kB
| 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() | |