/
sususer
/
ColonyGEN
Обзор
Документация
Войти
/
sususer
/
ColonyGEN
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
main
ml/vae_cli.py
193 строки
9 KB
Chekr
f11
03 июн 2026, 13:53
03 июн 2026, 13:53
d260831
Код
Авторство
О чём код?
from __future__ import annotations import argparse import json import sys import time from pathlib import Path from ml.vae.dataset import MapDataset from ml.vae.model import NumpyVAE from ml.vae.training import VAETrainer def log(message: str) -> None: print(message, flush=True) def progress_interval(total: int) -> int: return max(1, total // 5) def build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description="Custom NumPy VAE pipeline for ColonyGEN maps") subparsers = parser.add_subparsers(dest="command", required=True) train_parser = subparsers.add_parser("train", help="Train the VAE on the generated dataset") train_parser.add_argument("--dataset", default="data/datasets/maps", help="Dataset root with manifest and split folders") train_parser.add_argument("--output", default="data/models/vae", help="Output directory for model artifacts") train_parser.add_argument("--epochs", type=int, default=40, help="Number of training epochs") train_parser.add_argument("--batch-size", type=int, default=16, help="Batch size") train_parser.add_argument("--learning-rate", type=float, default=1e-3, help="Adam learning rate") train_parser.add_argument("--beta-max", type=float, default=0.02, help="Maximum KL coefficient") train_parser.add_argument("--beta-warmup", type=int, default=10, help="Epochs to reach beta_max") train_parser.add_argument("--seed", type=int, default=42, help="Random seed") train_parser.add_argument("--train-limit", type=int, default=None, help="Optional limit of train samples") train_parser.add_argument("--val-limit", type=int, default=None, help="Optional limit of validation samples") evaluate_parser = subparsers.add_parser("evaluate", help="Evaluate a trained model on a dataset split") evaluate_parser.add_argument("--dataset", default="data/datasets/maps", help="Dataset root with manifest and split folders") evaluate_parser.add_argument("--model", default="data/models/vae/vae_best.npz", help="Path to trained model") evaluate_parser.add_argument("--split", default="test", choices=["train", "val", "test"], help="Dataset split") evaluate_parser.add_argument("--batch-size", type=int, default=16, help="Batch size") evaluate_parser.add_argument("--beta", type=float, default=0.02, help="KL coefficient for reporting") evaluate_parser.add_argument("--limit", type=int, default=None, help="Optional limit of evaluated samples") sample_parser = subparsers.add_parser("sample", help="Generate new maps from the latent prior") sample_parser.add_argument("--dataset", default="data/datasets/maps", help="Dataset root used for codec configuration") sample_parser.add_argument("--model", default="data/models/vae/vae_best.npz", help="Path to trained model") sample_parser.add_argument("--output", default="data/generated/vae", help="Output directory for generated maps") sample_parser.add_argument("--count", type=int, default=8, help="Number of samples to generate") sample_parser.add_argument("--temperature", type=float, default=1.0, help="Latent sampling temperature") sample_parser.add_argument("--seed", type=int, default=42, help="Random seed") sample_parser.add_argument("--difficulty", type=float, default=None, help="Target difficulty score for conditional VAE sampling") sample_parser.add_argument("--entropy", type=float, default=None, help="Target entropy score for conditional VAE sampling") reconstruct_parser = subparsers.add_parser("reconstruct", help="Reconstruct existing maps through the VAE") reconstruct_parser.add_argument("--dataset", default="data/datasets/maps", help="Dataset root used for inputs") reconstruct_parser.add_argument("--model", default="data/models/vae/vae_best.npz", help="Path to trained model") reconstruct_parser.add_argument("--output", default="data/generated/vae_recon", help="Output directory for reconstructed maps") reconstruct_parser.add_argument("--split", default="test", choices=["train", "val", "test"], help="Dataset split") reconstruct_parser.add_argument("--count", type=int, default=8, help="Number of files to reconstruct") return parser def command_train(args: argparse.Namespace) -> None: trainer = VAETrainer( dataset_root=args.dataset, output_dir=args.output, batch_size=args.batch_size, epochs=args.epochs, learning_rate=args.learning_rate, beta_max=args.beta_max, beta_warmup_epochs=args.beta_warmup, seed=args.seed, train_limit=args.train_limit, val_limit=args.val_limit, ) summary = trainer.train() print(json.dumps(summary, indent=2), flush=True) def command_evaluate(args: argparse.Namespace) -> None: started = time.perf_counter() log( f"[VAE][EVAL] loading model={args.model} dataset={args.dataset} " f"split={args.split} batch_size={args.batch_size} beta={args.beta:.5f}" ) model, codec_config, _ = NumpyVAE.load(args.model) dataset = MapDataset( args.dataset, args.split, limit=args.limit, ) if codec_config: dataset.codec = dataset.codec.__class__.from_manifest(codec_config) data = dataset.load_array() conditions = dataset.load_conditions() batch_size = args.batch_size total_batches = max(1, (len(data) + batch_size - 1) // batch_size) interval = progress_interval(total_batches) sums = { "loss": 0.0, "reconstruction_loss": 0.0, "kl_loss": 0.0, "terrain_accuracy": 0.0, "biome_accuracy": 0.0, "base_accuracy": 0.0, "resource_mae": 0.0, } batches = 0 log(f"[VAE][EVAL] dataset loaded samples={len(data)} batches={total_batches}") for start in range(0, len(data), batch_size): batch = data[start:start + batch_size] batch_condition = conditions[start:start + batch_size] batch_metrics = model.evaluate_batch(batch, condition=batch_condition, beta=args.beta) reconstructed = model.reconstruct(batch, condition=batch_condition) recon_metrics = {key: 0.0 for key in ("terrain_accuracy", "biome_accuracy", "base_accuracy", "resource_mae")} for src, dst in zip(batch, reconstructed, strict=False): sample_metrics = dataset.codec.reconstruction_metrics(src, dst) for key, value in sample_metrics.items(): recon_metrics[key] += value for key in recon_metrics: recon_metrics[key] /= max(1, len(batch)) for key in ("loss", "reconstruction_loss", "kl_loss"): sums[key] += batch_metrics[key] for key in recon_metrics: sums[key] += recon_metrics[key] batches += 1 if batches == 1 or batches == total_batches or batches % interval == 0: log( f"[VAE][EVAL] batch={batches}/{total_batches} " f"loss={batch_metrics['loss']:.5f} terrain={recon_metrics['terrain_accuracy']:.3f} " f"biome={recon_metrics['biome_accuracy']:.3f} resource_mae={recon_metrics['resource_mae']:.4f} " f"elapsed={time.perf_counter() - started:.1f}s" ) result = {key: value / max(1, batches) for key, value in sums.items()} log( f"[VAE][EVAL] completed loss={result['loss']:.5f} terrain={result['terrain_accuracy']:.3f} " f"biome={result['biome_accuracy']:.3f} resource_mae={result['resource_mae']:.4f} " f"elapsed={time.perf_counter() - started:.1f}s" ) print(json.dumps(result, indent=2), flush=True) def command_sample(args: argparse.Namespace) -> None: trainer = VAETrainer(dataset_root=args.dataset, output_dir=Path(args.output).parent) if args.difficulty is not None and args.entropy is not None: written = trainer.sample_conditioned( model_path=args.model, output_dir=args.output, count=args.count, temperature=args.temperature, seed=args.seed, difficulty_score=args.difficulty, entropy_score=args.entropy, ) else: written = trainer.sample( model_path=args.model, output_dir=args.output, count=args.count, temperature=args.temperature, seed=args.seed, ) print(json.dumps({"written": written}, indent=2), flush=True) def command_reconstruct(args: argparse.Namespace) -> None: trainer = VAETrainer(dataset_root=args.dataset, output_dir=Path(args.output).parent) written = trainer.reconstruct_files( model_path=args.model, output_dir=args.output, split=args.split, count=args.count, ) print(json.dumps({"written": written}, indent=2), flush=True) def main() -> None: if hasattr(sys.stdout, "reconfigure"): sys.stdout.reconfigure(line_buffering=True) parser = build_parser() args = parser.parse_args() if args.command == "train": command_train(args) elif args.command == "evaluate": command_evaluate(args) elif args.command == "sample": command_sample(args) elif args.command == "reconstruct": command_reconstruct(args) else: raise ValueError(f"Unsupported command: {args.command}") if __name__ == "__main__": main()