File size: 2,841 Bytes
4397e12
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/bin/bash
# Run GRPO under a liveness watchdog. On 2026-10-07 r1 hit an xe GPU job timeout (5 s limit) and
# then HUNG instead of exiting, so a restart-on-exit loop would have waited forever. Here a run
# counts as dead when its log.jsonl stops changing for STALE_S seconds or the kernel logs a
# "Timedout job" for its pid; it is then SIGKILLed and restarted from its newest rl_stepN.pt
# (numbering continues via --step0), at most MAX_RESTARTS times.
#   RL_OUT=$TA_DATA/rl/r2 RL_CKPT=.../final.pt setsid nohup bash scripts/rl_watchdog.sh [grpo args] >/dev/null 2>&1 < /dev/null &
set -u
cd "$(dirname "$0")/.."
source env.sh
OUT=${RL_OUT:?set RL_OUT}
CKPT=${RL_CKPT:?set RL_CKPT}
STEPS=${RL_STEPS:-300}
STALE_S=${STALE_S:-420}
MAX_RESTARTS=${MAX_RESTARTS:-3}
L=$TA_DATA/logs/rl_watchdog.log
log() { echo "[$(date '+%F %T')] $*" | tee -a "$L"; }
mkdir -p "$OUT"

latest_step() { ls "$OUT"/rl_step*.pt 2>/dev/null | sed -E 's/.*rl_step([0-9]+)\.pt/\1/' | sort -n | tail -1; }
gpu_ok() { timeout 120 $TA_PY -c "import torch; x=torch.randn(1024,1024,device='xpu'); (x@x).sum().item()" >/dev/null 2>&1; }

restarts=0
ck=$CKPT
step0=0
while :; do
  log "starting grpo -> $OUT from $(basename "$ck") (step0=$step0, restart $restarts/$MAX_RESTARTS) args: $*"
  since=$(date '+%F %T')
  $TA_PY scripts/grpo.py --ckpt "$ck" --out "$OUT" --steps "$STEPS" --step0 "$step0" "$@" >> "$OUT/stdout.log" 2>&1 &
  pid=$!
  reason=""
  while kill -0 "$pid" 2>/dev/null; do
    sleep 30
    [ -f "$OUT/log.jsonl" ] && age=$(( $(date +%s) - $(stat -c %Y "$OUT/log.jsonl") )) || age=0
    # before the first log line, allow for model load + compile + the step-0 eval
    started=$(( $(date +%s) - $(date -d "$since" +%s) ))
    if [ "$started" -gt "$STALE_S" ] && [ "$age" -gt "$STALE_S" ]; then reason="log stale ${age}s"; fi
    if journalctl -k --since "$since" --no-pager 2>/dev/null | grep -q "Timedout job.*\[$pid\]"; then reason="GPU job timeout"; fi
    if [ -n "$reason" ]; then
      # confirm the pid is still our grpo before killing (pids get reused)
      if ps -o args= -p "$pid" | grep -q "scripts/grpo.py"; then kill -9 "$pid"; fi
      break
    fi
  done
  wait "$pid" 2>/dev/null; rc=$?
  last=$(latest_step)
  if [ -z "$reason" ] && [ "$rc" -eq 0 ]; then log "grpo finished (rc 0), last checkpoint step ${last:-none}"; break; fi
  log "grpo died: ${reason:-exit $rc}; tail: $(grep -v '^{' "$OUT/stdout.log" | tail -2 | tr '\n' ' ' | cut -c1-300)"
  restarts=$((restarts + 1))
  if [ "$restarts" -gt "$MAX_RESTARTS" ]; then log "giving up after $MAX_RESTARTS restarts"; break; fi
  sleep 20
  if ! gpu_ok; then log "GPU not usable after the failure; waiting 2 min"; sleep 120; gpu_ok || { log "GPU still not usable; giving up"; break; }; fi
  if [ -n "$last" ]; then ck="$OUT/rl_step$last.pt"; step0=$last; else ck=$CKPT; step0=0; fi
done