yzt15806542928's picture
Upload folder using huggingface_hub
bd3493c verified
Raw
History Blame Contribute Delete
9.25 kB
"""Train the official Aardvark Day-1 TAS model on official-schema tasks."""
from __future__ import annotations
import argparse
import json
import math
import random
import sys
from pathlib import Path
import numpy as np
import torch
import yaml
from torch.utils.data import DataLoader
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.aardvark_adapter import build_one_day_model
from model.sample_dataset import AardvarkTaskDataset, collate_tasks, discover_samples, split_samples
def parse_args() -> argparse.Namespace:
pre_parser = argparse.ArgumentParser(add_help=False)
pre_parser.add_argument("--config", type=Path, default=ROOT / "conf" / "config.yaml")
known, _ = pre_parser.parse_known_args()
config = yaml.safe_load(known.config.read_text())["training"]
parser = argparse.ArgumentParser(description=__doc__, parents=[pre_parser])
parser.add_argument("--data", type=Path, default=ROOT / config["data"])
parser.add_argument("--output-dir", type=Path, default=ROOT / config["output_dir"])
parser.add_argument("--epochs", type=int, default=config["epochs"])
parser.add_argument("--train-steps", type=int, default=config["train_steps"])
parser.add_argument("--validation-steps", type=int, default=config["validation_steps"])
parser.add_argument("--validation-fraction", type=float, default=config["validation_fraction"])
parser.add_argument("--batch-size", type=int, default=config["batch_size"])
parser.add_argument("--learning-rate", type=float, default=config["learning_rate"])
parser.add_argument("--weight-decay", type=float, default=config["weight_decay"])
parser.add_argument("--gradient-clip", type=float, default=config["gradient_clip"])
parser.add_argument("--patience", type=int, default=config["patience"])
parser.add_argument("--seed", type=int, default=config["seed"])
parser.add_argument("--train-modules", choices=("decoder", "all"), default=config["train_modules"])
parser.add_argument("--resume", type=Path)
return parser.parse_args()
def masked_metrics(prediction: torch.Tensor, target: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, int]:
valid = torch.isfinite(target)
count = int(valid.sum())
if count == 0:
raise RuntimeError("Aardvark task has no finite station targets")
error = prediction[valid] - target[valid]
return torch.sqrt(error.square().mean()), error.abs().mean(), count
def configure_trainable_parameters(model: torch.nn.Module, train_modules: str) -> list[torch.nn.Parameter]:
for parameter in model.parameters():
parameter.requires_grad_(train_modules == "all")
if train_modules == "decoder":
for parameter in model.sf_model.parameters():
parameter.requires_grad_(True)
parameters = [parameter for parameter in model.parameters() if parameter.requires_grad]
if not parameters:
raise RuntimeError("No trainable parameters selected")
return parameters
def checkpoint_payload(model, optimizer, scheduler, args, epoch, best_validation_rmse, history):
model_state = model.state_dict() if args.train_modules == "all" else model.sf_model.state_dict()
return {
"model": model_state,
"optimizer": optimizer.state_dict(),
"scheduler": scheduler.state_dict(),
"epoch": epoch,
"best_validation_rmse": best_validation_rmse,
"history": history,
"train_modules": args.train_modules,
"config": {key: str(value) if isinstance(value, Path) else value for key, value in vars(args).items()},
"torch_rng_state": torch.get_rng_state(),
"numpy_rng_state": np.random.get_state(),
"python_rng_state": random.getstate(),
}
def main() -> None:
args = parse_args()
if not torch.cuda.is_available():
raise RuntimeError("The official Aardvark model requires CUDA")
if min(args.epochs, args.train_steps, args.validation_steps, args.batch_size) < 1:
raise ValueError("epochs, steps and batch_size must be at least 1")
random.seed(args.seed)
np.random.seed(args.seed)
torch.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
args.output_dir.mkdir(parents=True, exist_ok=True)
samples = discover_samples(args.data)
train_samples, validation_samples = split_samples(samples, args.validation_fraction, args.seed)
train_loader = DataLoader(
AardvarkTaskDataset(train_samples, args.train_steps * args.batch_size),
batch_size=args.batch_size,
shuffle=True,
collate_fn=collate_tasks,
)
validation_loader = DataLoader(
AardvarkTaskDataset(validation_samples, args.validation_steps * args.batch_size),
batch_size=args.batch_size,
collate_fn=collate_tasks,
)
model = build_one_day_model(ROOT / "weights", ROOT / "official-src", "cuda")
model.return_gridded = False
parameters = configure_trainable_parameters(model, args.train_modules)
optimizer = torch.optim.AdamW(parameters, lr=args.learning_rate, weight_decay=args.weight_decay)
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(
optimizer,
mode="min",
factor=0.5,
patience=max(1, args.patience // 2),
)
start_epoch = 0
best_validation_rmse = math.inf
history: list[dict[str, float | int]] = []
if args.resume:
payload = torch.load(args.resume, map_location="cuda", weights_only=False)
if payload["train_modules"] != args.train_modules:
raise ValueError("--train-modules must match the resumed checkpoint")
target_model = model if args.train_modules == "all" else model.sf_model
target_model.load_state_dict(payload["model"])
optimizer.load_state_dict(payload["optimizer"])
scheduler.load_state_dict(payload["scheduler"])
start_epoch = int(payload["epoch"]) + 1
best_validation_rmse = float(payload["best_validation_rmse"])
history = payload["history"]
torch.set_rng_state(payload["torch_rng_state"])
np.random.set_state(payload["numpy_rng_state"])
random.setstate(payload["python_rng_state"])
epochs_without_improvement = 0
last_path = args.output_dir / "last.pth"
best_path = args.output_dir / "best.pth"
for epoch in range(start_epoch, args.epochs):
model.train()
if args.train_modules == "decoder":
model.se_model.eval()
model.forecast_model.eval()
model.sf_model.train()
train_rmse = []
train_mae = []
for task in train_loader:
optimizer.zero_grad(set_to_none=True)
prediction = model(task)
target = task["y_target"].to(prediction.device)
rmse, mae, _ = masked_metrics(prediction, target)
rmse.backward()
torch.nn.utils.clip_grad_norm_(parameters, args.gradient_clip)
optimizer.step()
train_rmse.append(float(rmse.detach()))
train_mae.append(float(mae.detach()))
model.eval()
validation_rmse = []
validation_mae = []
valid_stations = 0
with torch.inference_mode():
for task in validation_loader:
prediction = model(task)
rmse, mae, valid_stations = masked_metrics(prediction, task["y_target"].to(prediction.device))
validation_rmse.append(float(rmse))
validation_mae.append(float(mae))
record = {
"epoch": epoch,
"train_rmse": float(np.mean(train_rmse)),
"train_mae": float(np.mean(train_mae)),
"validation_rmse": float(np.mean(validation_rmse)),
"validation_mae": float(np.mean(validation_mae)),
"learning_rate": optimizer.param_groups[0]["lr"],
"valid_stations": valid_stations,
}
history.append(record)
scheduler.step(record["validation_rmse"])
improved = record["validation_rmse"] < best_validation_rmse
if improved:
best_validation_rmse = record["validation_rmse"]
epochs_without_improvement = 0
else:
epochs_without_improvement += 1
payload = checkpoint_payload(model, optimizer, scheduler, args, epoch, best_validation_rmse, history)
torch.save(payload, last_path)
if improved:
torch.save(payload, best_path)
(args.output_dir / "history.json").write_text(json.dumps(history, indent=2) + "\n")
print(json.dumps(record, sort_keys=True))
if epochs_without_improvement >= args.patience:
break
report = {
"status": "completed",
"epochs_completed": len(history),
"train_modules": args.train_modules,
"train_samples": [str(path) for path in train_samples],
"validation_samples": [str(path) for path in validation_samples],
"best_validation_rmse": best_validation_rmse,
"best_checkpoint": str(best_path),
"last_checkpoint": str(last_path),
}
(args.output_dir / "train.json").write_text(json.dumps(report, indent=2) + "\n")
print(json.dumps(report, indent=2))
if __name__ == "__main__":
main()