Spaces:
Sleeping
Sleeping
File size: 8,410 Bytes
78738de | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 | #!/usr/bin/env python3
"""Train PPO/SAC trên SinkOneBallEnv (curriculum stage 1).
Chạy từ gốc repo (venv đã kích hoạt):
python scripts/train_sink.py # PPO, 300k bước, 8 env
python scripts/train_sink.py --algo sac # thử SAC (off-policy)
python scripts/train_sink.py --total-steps 1000000 # chạy qua đêm
Lưu ý lần chạy đầu: mỗi subprocess phải load Numba cache (~vài giây),
sau đó tốc độ ~150-250 steps/s với 8 env.
Output:
models/<run>/final_model.zip — model đã train
logs/<run>/monitor.csv — reward từng episode (nguồn learning curve)
logs/<run>/learning_curve.png — đường cong học tập (gửi GVHD)
+ đánh giá cuối: pot rate / scratch rate so với random baseline
"""
from __future__ import annotations
import argparse
import sys
import time
from pathlib import Path
sys.stdout.reconfigure(encoding="utf-8")
sys.stderr.reconfigure(encoding="utf-8")
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / "src"))
def make_env(scratch_penalty: float | None = None):
# import trong hàm để subprocess (Windows spawn) tự import lại được
from poolcoach_rl.envs import SinkOneBallEnv
if scratch_penalty is None:
return SinkOneBallEnv()
return SinkOneBallEnv(scratch_penalty=scratch_penalty)
def plot_learning_curve(monitor_csv: Path, out_png: Path, window: int = 500):
"""Vẽ rolling mean reward + pot rate từ monitor.csv (chỉ cần numpy)."""
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
# dòng 1 là metadata JSON, dòng 2 là header r,l,t,potted,scratch,contact
data = np.genfromtxt(monitor_csv, delimiter=",", names=True, skip_header=1)
rewards = np.atleast_1d(data["r"])
if len(rewards) < 2 * window:
window = max(10, len(rewards) // 10)
window = max(1, min(window, len(rewards)))
def roll(x: np.ndarray) -> np.ndarray:
return np.convolve(x, np.ones(window) / window, mode="valid")
# đọc trực tiếp từ cột info; fallback về ngưỡng reward nếu là run cũ
def col(name: str, fallback: np.ndarray) -> np.ndarray:
return np.atleast_1d(data[name]) if name in data.dtype.names else fallback
episodes = np.arange(window, len(rewards) + 1)
pot = roll(col("potted", (rewards >= 0.99).astype(float)))
scratch = roll(col("scratch", (rewards <= -0.99).astype(float)))
contact = roll(col("contact", np.zeros_like(rewards)))
aim = roll(col("aim_cos", np.zeros_like(rewards)))
fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(9, 7), sharex=True)
ax1.plot(episodes, roll(rewards), lw=1.5)
ax1.set_ylabel(f"Reward (rolling mean, w={window})")
ax1.set_title("PoolCoach stage 1 — SinkOneBallEnv learning curve")
ax1.grid(alpha=0.3)
ax2.plot(episodes, aim, lw=1.2, c="tab:purple", ls="--", alpha=0.8,
label="Aim-ghost cos (kỳ vọng → 1.0)")
ax2.plot(episodes, contact, lw=1.2, c="tab:green", alpha=0.7, label="Contact rate")
ax2.plot(episodes, pot, lw=1.5, label="Pot rate")
ax2.plot(episodes, scratch, lw=1.5, label="Scratch rate")
ax2.axhline(0.04, ls="--", c="gray", lw=1, label="Random pot (~4%)")
ax2.set_xlabel("Episode (= số cú đánh)")
ax2.set_ylabel("Tỉ lệ")
ax2.legend(loc="upper right") # loc cố định: tránh warning "best" chậm với 1M điểm
ax2.grid(alpha=0.3)
fig.tight_layout()
fig.savefig(out_png, dpi=150)
print(f"Learning curve -> {out_png}")
def evaluate(model, n_episodes: int = 200,
scratch_penalty: float | None = None) -> dict:
"""Đánh giá policy deterministic, so với random baseline."""
env = make_env(scratch_penalty)
potted = scratched = contacted = 0
rewards, aims, disps = [], [], []
for _ in range(n_episodes):
obs, _ = env.reset()
action, _ = model.predict(obs, deterministic=True)
_, r, _, _, info = env.step(action)
rewards.append(r)
potted += info.get("potted", False)
scratched += info.get("scratch", False)
contacted += info.get("contact", False)
aims.append(info.get("aim_cos", 0.0))
disps.append(info.get("tgt_disp", 0.0))
import numpy as np
return {
"pot_rate": potted / n_episodes,
"scratch_rate": scratched / n_episodes,
"contact_rate": contacted / n_episodes,
"aim_cos_mean": float(np.mean(aims)),
"tgt_disp_mean": float(np.mean(disps)),
"reward_mean": float(np.mean(rewards)),
}
def main():
p = argparse.ArgumentParser()
p.add_argument("--algo", choices=["ppo", "sac"], default="ppo")
p.add_argument("--total-steps", type=int, default=300_000)
p.add_argument("--n-envs", type=int, default=8)
p.add_argument("--seed", type=int, default=0)
p.add_argument("--ent-coef", type=float, default=0.01,
help="PPO entropy coef (thử 0.03-0.05 nếu còn collapse)")
p.add_argument("--scratch-penalty", type=float, default=None,
help="ghi đè SCRATCH_PENALTY của env; nhập 0.5 hay -0.5 "
"đều hiểu là -0.5 (mặc định: giữ -0.3 trong env)")
p.add_argument("--run-name", default=None)
p.add_argument("--plot-only", metavar="MONITOR_CSV",
help="chỉ vẽ lại curve từ monitor.csv có sẵn rồi thoát")
args = p.parse_args()
if args.plot_only:
csv = Path(args.plot_only)
plot_learning_curve(csv, csv.parent / "learning_curve.png")
return
from functools import partial
from stable_baselines3 import PPO, SAC
from stable_baselines3.common.vec_env import SubprocVecEnv, VecMonitor
if args.scratch_penalty is not None:
args.scratch_penalty = -abs(args.scratch_penalty)
print(f"[config] scratch_penalty override: {args.scratch_penalty}")
run = args.run_name or f"{args.algo}_{time.strftime('%Y%m%d_%H%M%S')}"
log_dir = ROOT / "logs" / run
model_dir = ROOT / "models" / run
log_dir.mkdir(parents=True, exist_ok=True)
model_dir.mkdir(parents=True, exist_ok=True)
venv = SubprocVecEnv(
[partial(make_env, args.scratch_penalty) for _ in range(args.n_envs)]
)
# info_keywords -> ghi thêm cột potted/scratch/contact vào monitor.csv
# (nguon chinh xac cho learning curve, thay vi suy tu nguong reward)
venv = VecMonitor(
venv,
filename=str(log_dir / "monitor.csv"),
info_keywords=("potted", "scratch", "contact", "aim_cos", "tgt_disp"),
)
common = dict(env=venv, verbose=1, seed=args.seed,
tensorboard_log=str(log_dir))
if args.algo == "ppo":
model = PPO("MlpPolicy", n_steps=128, batch_size=256,
ent_coef=args.ent_coef, **common)
else:
model = SAC("MlpPolicy", buffer_size=200_000, learning_starts=2_000,
train_freq=(8, "step"), **common)
try: # progress bar cần tqdm + rich; thiếu thì train không bar
import tqdm # noqa: F401
import rich # noqa: F401
progress = True
except ImportError:
progress = False
t0 = time.time()
model.learn(total_timesteps=args.total_steps, progress_bar=progress)
dt = time.time() - t0
print(f"\nTrain {args.total_steps:,} bước trong {dt/60:.1f} phút "
f"({args.total_steps/dt:.0f} steps/s)")
model.save(model_dir / "final_model")
print(f"Model -> {model_dir / 'final_model.zip'}")
venv.close()
plot_learning_curve(log_dir / "monitor.csv",
log_dir / "learning_curve.png")
print("\n== Đánh giá deterministic (200 cú) ==")
stats = evaluate(model, scratch_penalty=args.scratch_penalty)
print(f" Pot rate : {stats['pot_rate']:.1%} (random: ~4%)")
print(f" Scratch rate : {stats['scratch_rate']:.1%} (random: ~20-26%)")
print(f" Contact rate : {stats['contact_rate']:.1%} (random: ~18%)")
print(f" Aim cos mean : {stats['aim_cos_mean']:+.3f} (random: ~0, học tốt: →1)")
print(f" Tgt disp mean: {stats['tgt_disp_mean']:.3f} m (tap ~0.05, cú thật >=0.5)")
print(f" Reward mean : {stats['reward_mean']:+.4f}")
if __name__ == "__main__":
main()
|