Wire AMP into Deep CFR trainer behind run.use_amp flag (default off)

Adds torch.autocast(fp16) + GradScaler around _train_advantage and
_train_strategy when run.use_amp=true and device=cuda. CPU/non-CUDA
falls back to fp32 no-op. Mitigations:
- scaler.unscale_(optimizer) before grad_clip.
- nonfinite-loss guard skips overflowing batches and counts them.
- diff.float().square() in advantage loss to avoid fp16 overflow.
- strategy mask/log_softmax kept in fp32.

New metrics: amp/grad_scale, amp/nonfinite_loss_count.

Tests: AMP CUDA smoke + CPU fallback in test_deep_cfr_trainer.py.

Bench: scripts/bench_amp_trainer.py micro-benches train phases under
synthetic replay memory. smoke.yaml result is fp32 3.22ms / AMP 3.92ms
(0.82×, regression). 100-iter A/B on default.yaml deliberately
skipped: smoke regression mirrors the 2026-05-07 torch.compile
regression dynamic (dispatch overhead > kernel benefit at this model
size) and re-confirming on the same size adds no information.

Default stays run.use_amp: false. Re-enable trigger documented in
docs/performance.md: hidden_size >= 1024 or num_layers >= 6, then run
the bench script + 100-iter A/B before flipping default.
This commit is contained in:
2026-05-07 20:21:17 +09:00
parent ad0be89857
commit 6c976f468a
4 changed files with 351 additions and 33 deletions
+55
View File
@@ -358,6 +358,59 @@ batched-traversal-inference work in Optimization Priorities #5 lands —
that is the change that would put compile on the dominant phase, not that is the change that would put compile on the dominant phase, not
just on the trainer's optimization steps. Not enabled on `main`. just on the trainer's optimization steps. Not enabled on `main`.
### AMP on trainer networks (2026-05-07, regression)
Wrapped the trainer optimization phases with `torch.autocast(fp16)` and
`torch.amp.GradScaler`: `_train_advantage` and `_train_strategy` now run
their network forward/backward/optimizer step through the AMP path when
`run.use_amp=true` and the trainer device is CUDA.
Safety mitigations included in the implementation:
- `GradScaler.unscale_(optimizer)` is called before `clip_grad_norm_`.
- Non-finite loss guard increments `amp/nonfinite_loss_count` and skips the
bad step instead of applying it.
- Advantage squared loss computes `diff.float().square()` so the loss
reduction is fp32 even when the forward path is autocast to fp16.
- Strategy logits are cast back to fp32 before `masked_fill` and
`log_softmax`.
- Metrics now expose `amp/grad_scale` and `amp/nonfinite_loss_count`.
Measurement used the small `smoke.yaml` config with synthetic replay-memory
samples via:
```bash
uv run python scripts/bench_amp_trainer.py \
--config configs/deep_cfr/smoke.yaml \
--runs 3 \
--warmup 1 \
--device cuda
```
| | mean ms/call | speedup vs fp32 |
| --- | ---: | ---: |
| fp32 | 3.22 | 1.00× |
| AMP (fp16) | 3.92 | 0.82× |
Net result: regression. This matches the same dispatch-overhead-vs-kernel
benefit dynamic as the `torch.compile` regression above: the current trainer
model and smoke workload are too small for AMP's lower-precision kernels to
pay back autocast and scaler bookkeeping overhead.
The full `default.yaml` 100-iteration A/B was intentionally skipped. Given
the small-model regression and the matching `torch.compile` precedent on the
same model family, there is no current evidence that spending GPU time on the
longer A/B would produce a different decision. The infrastructure is kept
merged but default-off: `run.use_amp=false` remains the default, and
re-enabling is a one-field config flip.
Re-measure AMP only after the model grows to at least `hidden_size >= 1024`
or `num_layers >= 6`. At that point run both the fast
`scripts/bench_amp_trainer.py` micro-bench and the formal 100-iteration
fp32-vs-AMP A/B. If AMP still provides less than 5% speedup at that larger
model size, keep it default-off and raise the next re-measure trigger to an
even larger model.
### GPU forward profiling for batched traversal (2026-05-07, decision support) ### GPU forward profiling for batched traversal (2026-05-07, decision support)
To decide whether Optimization Priorities #5 (batched traversal inference) is To decide whether Optimization Priorities #5 (batched traversal inference) is
@@ -556,6 +609,8 @@ Do this in order. Skipping ahead is the failure mode that creates misleading
2. **Next**: experiment with a larger network config. Measure compute vs 2. **Next**: experiment with a larger network config. Measure compute vs
learning-curve trade-off with the existing toolchain (no compile/TRT yet). learning-curve trade-off with the existing toolchain (no compile/TRT yet).
This step decides the model size that future optimizations target. This step decides the model size that future optimizations target.
It is also the prerequisite for revisiting AMP, `torch.compile`, and
TensorRT: all three are dispatch-overhead-bound on the current small model.
3. **Then**: re-measure `torch.compile` on the trainer at the chosen model 3. **Then**: re-measure `torch.compile` on the trainer at the chosen model
size. The earlier regression was size-bound; expect a different result. size. The earlier regression was size-bound; expect a different result.
4. **Then**: integrate TensorRT into the inference server (covers traversal 4. **Then**: integrate TensorRT into the inference server (covers traversal
+155
View File
@@ -0,0 +1,155 @@
from __future__ import annotations
import argparse
import math
import statistics
import time
from pathlib import Path
import numpy as np
import torch
from coolrl_lost_cities.games.classic.deep_cfr.cli import _with_overrides
from coolrl_lost_cities.games.classic.deep_cfr.config import load_config
from coolrl_lost_cities.games.classic.deep_cfr.memory import TrainingSample
from coolrl_lost_cities.games.classic.deep_cfr.tracking import NullRunTracker
from coolrl_lost_cities.games.classic.deep_cfr.trainer import DeepCFRTrainer
REPO_ROOT = Path(__file__).resolve().parent.parent
DEFAULT_CONFIG = REPO_ROOT / "configs" / "deep_cfr" / "default.yaml"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Benchmark trainer-side AMP train phases.")
parser.add_argument("--config", default=str(DEFAULT_CONFIG))
parser.add_argument("--runs", type=int, default=20)
parser.add_argument("--warmup", type=int, default=2)
parser.add_argument("--device", default="cuda")
return parser.parse_args()
def _seed_memories(trainer: DeepCFRTrainer) -> None:
rng = np.random.default_rng(trainer.config.run.seed + 909)
action_size = trainer.action_size
input_dim = trainer.input_dim
advantage_count = (
trainer.config.optimization.advantage_batch_size
* trainer.config.optimization.advantage_updates_per_iteration
)
strategy_count = (
trainer.config.optimization.strategy_batch_size
* trainer.config.optimization.strategy_updates_per_iteration
)
for player in range(2):
for index in range(max(advantage_count, trainer.config.optimization.advantage_batch_size)):
legal = rng.random(action_size) > 0.25
legal[int(rng.integers(0, action_size))] = True
trainer.advantage_memories[player].add(
TrainingSample(
info_state=rng.normal(size=input_dim).astype(np.float32),
target=rng.normal(scale=20.0, size=action_size).astype(np.float32),
legal_mask=legal,
iteration=index + 1,
player=player,
),
rng,
)
for index in range(max(strategy_count, trainer.config.optimization.strategy_batch_size)):
legal = rng.random(action_size) > 0.25
legal[int(rng.integers(0, action_size))] = True
target = np.zeros(action_size, dtype=np.float32)
weights = rng.random(np.count_nonzero(legal)).astype(np.float32)
weights /= weights.sum()
target[legal] = weights
trainer.strategy_memory.add(
TrainingSample(
info_state=rng.normal(size=input_dim).astype(np.float32),
target=target,
legal_mask=legal,
iteration=index + 1,
player=-1,
),
rng,
)
def _new_trainer(*, use_amp: bool, config_path: str, device: str) -> DeepCFRTrainer:
config = _with_overrides(
load_config(config_path),
{
"run": {"use_amp": use_amp, "device": device, "max_iterations": 1},
"checkpoint": {"save_every": 0, "save_latest": False},
"evaluation": {"eval_every": 0},
},
)
trainer = DeepCFRTrainer(
config=config,
game_config=config.rules.to_lost_cities_config(seed=config.run.seed),
device=device,
tracker=NullRunTracker(),
)
_seed_memories(trainer)
return trainer
def _measure(trainer: DeepCFRTrainer, runs: int) -> list[float]:
durations: list[float] = []
for iteration in range(runs):
trainer.iteration = iteration + 1
trainer._runtime_metrics = {}
if trainer.device.type == "cuda":
torch.cuda.synchronize()
started = time.perf_counter()
trainer._train_advantage_networks()
trainer._train_strategy_network()
if trainer.device.type == "cuda":
torch.cuda.synchronize()
durations.append(time.perf_counter() - started)
return durations
def _summary(values: list[float], warmup: int) -> tuple[float, float, float]:
measured = values[warmup:]
if not measured:
raise ValueError("warmup must be less than runs")
return (
statistics.mean(measured),
statistics.median(measured),
sorted(measured)[math.ceil(0.95 * (len(measured) - 1))],
)
def main() -> None:
args = parse_args()
if args.runs <= 0:
raise SystemExit("--runs must be positive")
if args.warmup < 0 or args.warmup >= args.runs:
raise SystemExit("--warmup must be non-negative and less than --runs")
if args.device == "cuda" and not torch.cuda.is_available():
raise SystemExit("CUDA is not available; pass --device cpu for CPU fallback smoke.")
print(f"Config: {args.config}")
print(f"Device: {args.device}")
print(f"Runs: {args.runs} Warmup: {args.warmup}")
print()
rows: list[tuple[str, float, float, float]] = []
for use_amp, label in ((False, "fp32"), (True, "amp")):
trainer = _new_trainer(use_amp=use_amp, config_path=args.config, device=args.device)
values = _measure(trainer, args.runs)
mean, p50, p95 = _summary(values, args.warmup)
rows.append((label, mean, p50, p95))
del trainer
if torch.cuda.is_available():
torch.cuda.empty_cache()
print(f"{'mode':<8} {'mean_ms':>10} {'p50_ms':>10} {'p95_ms':>10}")
for label, mean, p50, p95 in rows:
print(f"{label:<8} {mean * 1000.0:>10.2f} {p50 * 1000.0:>10.2f} {p95 * 1000.0:>10.2f}")
fp32_mean = rows[0][1]
amp_mean = rows[1][1]
print()
print(f"speedup: {fp32_mean / amp_mean:.2f}x")
if __name__ == "__main__":
main()
@@ -230,6 +230,9 @@ class DeepCFRTrainer:
lr=self.config.optimization.learning_rate, lr=self.config.optimization.learning_rate,
weight_decay=self.config.optimization.weight_decay, weight_decay=self.config.optimization.weight_decay,
) )
self._amp_enabled = bool(self.config.run.use_amp) and self.device.type == "cuda"
self._amp_dtype = torch.float16
self._scaler = torch.amp.GradScaler("cuda", enabled=self._amp_enabled)
self.advantage_memories = [ self.advantage_memories = [
ReservoirMemory(self.config.memory.advantage_capacity) for _ in range(2) ReservoirMemory(self.config.memory.advantage_capacity) for _ in range(2)
] ]
@@ -258,6 +261,8 @@ class DeepCFRTrainer:
self._runtime_metrics: dict[str, float | int] = {} self._runtime_metrics: dict[str, float | int] = {}
self._inference_server: InferenceServerController | None = None self._inference_server: InferenceServerController | None = None
self._last_inference_weight_sync_iteration: int | None = None self._last_inference_weight_sync_iteration: int | None = None
if self.config.run.use_amp and not self._amp_enabled:
self.tracker.log_event("AMP requested but trainer device is not CUDA; running fp32.")
def checkpoint_payload(self, metrics: IterationMetrics | None = None) -> dict: def checkpoint_payload(self, metrics: IterationMetrics | None = None) -> dict:
return { return {
@@ -339,6 +344,9 @@ class DeepCFRTrainer:
self._runtime_metrics["time/strategy_train_seconds"] = ( self._runtime_metrics["time/strategy_train_seconds"] = (
time.perf_counter() - strategy_started time.perf_counter() - strategy_started
) )
if self.config.run.use_amp:
self._runtime_metrics["amp/grad_scale"] = float(self._scaler.get_scale())
self._runtime_metrics.setdefault("amp/nonfinite_loss_count", 0)
eval_started = time.perf_counter() eval_started = time.perf_counter()
eval_metrics = self._evaluate(iteration) eval_metrics = self._evaluate(iteration)
@@ -953,40 +961,52 @@ class DeepCFRTrainer:
- sample_started - sample_started
) )
x, y, legal, sample_iterations = self._batch_tensors(batch) x, y, legal, sample_iterations = self._batch_tensors(batch)
pred = network(x) with torch.autocast(
diff = (pred - y).masked_fill(~legal, 0.0) device_type="cuda",
if self.config.training_weighting.mode == "none": dtype=self._amp_dtype,
loss = diff.square().sum() / legal.sum().clamp_min(1) enabled=self._amp_enabled,
elif self.config.training_weighting.mode == "lcfr": ):
sample_weights = self._iteration_weights( pred = network(x)
sample_iterations, self.config.training_weighting.lcfr_alpha diff = (pred - y).masked_fill(~legal, 0.0).float()
) squared = diff.square()
action_weights = sample_weights[:, None] * legal.float() if self.config.training_weighting.mode == "none":
loss = (diff.square() * action_weights).sum() / action_weights.sum().clamp_min( loss = squared.sum() / legal.sum().clamp_min(1)
1.0e-12 elif self.config.training_weighting.mode == "lcfr":
) sample_weights = self._iteration_weights(
else: sample_iterations, self.config.training_weighting.lcfr_alpha
positive_weights = self._iteration_weights( )
sample_iterations, self.config.training_weighting.dcfr_alpha action_weights = sample_weights[:, None] * legal.float()
) loss = (squared * action_weights).sum() / action_weights.sum().clamp_min(
negative_weights = self._iteration_weights( 1.0e-12
sample_iterations, self.config.training_weighting.dcfr_beta )
) else:
target_weights = torch.where( positive_weights = self._iteration_weights(
y >= 0.0, positive_weights[:, None], negative_weights[:, None] sample_iterations, self.config.training_weighting.dcfr_alpha
) )
action_weights = target_weights * legal.float() negative_weights = self._iteration_weights(
loss = (diff.square() * action_weights).sum() / action_weights.sum().clamp_min( sample_iterations, self.config.training_weighting.dcfr_beta
1.0e-12 )
) target_weights = torch.where(
y >= 0.0, positive_weights[:, None], negative_weights[:, None]
)
action_weights = target_weights * legal.float()
loss = (squared * action_weights).sum() / action_weights.sum().clamp_min(
1.0e-12
)
if not torch.isfinite(loss):
self._record_nonfinite_loss()
optimizer.zero_grad(set_to_none=True)
continue
optimizer.zero_grad(set_to_none=True) optimizer.zero_grad(set_to_none=True)
loss.backward() self._scaler.scale(loss).backward()
if self.config.optimization.grad_clip > 0.0: if self.config.optimization.grad_clip > 0.0:
self._scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_( torch.nn.utils.clip_grad_norm_(
network.parameters(), self.config.optimization.grad_clip network.parameters(), self.config.optimization.grad_clip
) )
optimizer.step() self._scaler.step(optimizer)
losses.append(float(loss.detach().cpu())) self._scaler.update()
losses.append(float(loss.detach().float().cpu()))
return float(np.mean(losses)) if losses else 0.0 return float(np.mean(losses)) if losses else 0.0
def _train_strategy( def _train_strategy(
@@ -1007,7 +1027,13 @@ class DeepCFRTrainer:
- sample_started - sample_started
) )
x, y, legal, sample_iterations = self._batch_tensors(batch) x, y, legal, sample_iterations = self._batch_tensors(batch)
logits = network(x).masked_fill(~legal, torch.finfo(torch.float32).min) with torch.autocast(
device_type="cuda",
dtype=self._amp_dtype,
enabled=self._amp_enabled,
):
logits = network(x)
logits = logits.float().masked_fill(~legal, torch.finfo(torch.float32).min)
log_probs = nn.functional.log_softmax(logits, dim=-1).masked_fill(~legal, 0.0) log_probs = nn.functional.log_softmax(logits, dim=-1).masked_fill(~legal, 0.0)
per_sample_loss = -(y * log_probs).sum(dim=-1) per_sample_loss = -(y * log_probs).sum(dim=-1)
if self.config.training_weighting.mode == "none": if self.config.training_weighting.mode == "none":
@@ -1026,12 +1052,23 @@ class DeepCFRTrainer:
loss = (per_sample_loss * sample_weights).sum() / sample_weights.sum().clamp_min( loss = (per_sample_loss * sample_weights).sum() / sample_weights.sum().clamp_min(
1.0e-12 1.0e-12
) )
if not torch.isfinite(loss):
self._record_nonfinite_loss()
optimizer.zero_grad(set_to_none=True)
continue
optimizer.zero_grad(set_to_none=True) optimizer.zero_grad(set_to_none=True)
loss.backward() self._scaler.scale(loss).backward()
if self.config.optimization.grad_clip > 0.0: if self.config.optimization.grad_clip > 0.0:
self._scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_( torch.nn.utils.clip_grad_norm_(
network.parameters(), self.config.optimization.grad_clip network.parameters(), self.config.optimization.grad_clip
) )
optimizer.step() self._scaler.step(optimizer)
last_loss = float(loss.detach().cpu()) self._scaler.update()
last_loss = float(loss.detach().float().cpu())
return last_loss return last_loss
def _record_nonfinite_loss(self) -> None:
self._runtime_metrics["amp/nonfinite_loss_count"] = (
int(self._runtime_metrics.get("amp/nonfinite_loss_count", 0)) + 1
)
@@ -4,6 +4,7 @@ import re
from pathlib import Path from pathlib import Path
import numpy as np import numpy as np
import pytest
import torch import torch
from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state, input_dim from coolrl_lost_cities.games.classic.deep_cfr.encoding import encode_info_state, input_dim
from coolrl_lost_cities.games.classic.deep_cfr.traversal import CythonDeepCFRTraverser from coolrl_lost_cities.games.classic.deep_cfr.traversal import CythonDeepCFRTraverser
@@ -399,6 +400,76 @@ def test_deep_cfr_trainer_supports_lcfr_and_dcfr_loss_weighting() -> None:
assert metrics[0].strategy_loss >= 0.0 assert metrics[0].strategy_loss >= 0.0
def test_deep_cfr_trainer_amp_cpu_falls_back_to_fp32() -> None:
trainer = DeepCFRTrainer(
_deep_cfr_config(
{
"run": {"max_iterations": 1, "seed": 25, "device": "cpu", "use_amp": True},
"network": {"hidden_size": 16},
"traversal": {
"traversals_per_player": 1,
"max_depth": 2,
"max_nodes_per_traversal": 32,
},
"optimization": {
"advantage_updates_per_iteration": 1,
"strategy_updates_per_iteration": 1,
"advantage_batch_size": 2,
"strategy_batch_size": 2,
},
"checkpoint": {"save_every": 0},
"evaluation": {"eval_every": 0},
}
),
LostCitiesConfig(seed=25),
device="cpu",
)
metrics = trainer.train()
assert len(metrics) == 1
assert metrics[0].runtime_metrics["amp/grad_scale"] == 1.0
assert metrics[0].runtime_metrics["amp/nonfinite_loss_count"] == 0
assert metrics[0].advantage_loss >= 0.0
assert metrics[0].strategy_loss >= 0.0
def test_deep_cfr_trainer_amp_cuda_smoke() -> None:
if not torch.cuda.is_available():
pytest.skip("CUDA is not available")
trainer = DeepCFRTrainer(
_deep_cfr_config(
{
"run": {"max_iterations": 1, "seed": 26, "device": "cuda", "use_amp": True},
"network": {"hidden_size": 16},
"traversal": {
"traversals_per_player": 1,
"max_depth": 2,
"max_nodes_per_traversal": 32,
},
"optimization": {
"advantage_updates_per_iteration": 1,
"strategy_updates_per_iteration": 1,
"advantage_batch_size": 2,
"strategy_batch_size": 2,
},
"checkpoint": {"save_every": 0},
"evaluation": {"eval_every": 0},
}
),
LostCitiesConfig(seed=26),
device="cuda",
)
metrics = trainer.train()
assert len(metrics) == 1
assert metrics[0].runtime_metrics["amp/grad_scale"] > 0.0
assert metrics[0].runtime_metrics["amp/nonfinite_loss_count"] == 0
assert np.isfinite(metrics[0].advantage_loss)
assert np.isfinite(metrics[0].strategy_loss)
def test_deep_cfr_cython_traverser_restores_state_and_collects_samples() -> None: def test_deep_cfr_cython_traverser_restores_state_and_collects_samples() -> None:
trainer = DeepCFRTrainer( trainer = DeepCFRTrainer(
_deep_cfr_config( _deep_cfr_config(