fix: fetch_train.py v2 - fix wrong 0/1 step from checkpoint sharding, parse loss from trainer_state.json

This commit is contained in:
bingshuo 2026-05-20 02:45:54 +08:00
parent b44ae07711
commit 0da8b1f3e2

View File

@ -1,25 +1,57 @@
#!/usr/bin/env python3
"""fetch_train.py v2 — 从AutoDL实时同步代码模型训练进度
修复: 排除checkpoint保存时的"0/1"假进度从trainer_state.json解析loss
"""
import subprocess, json, re
from datetime import datetime, timezone
DATA_FILE = '/opt/guanghulab-repo/homepage/training-status.json'
SSH = ['sshpass','-p','HkM43lFVUIsc','ssh','-o','StrictHostKeyChecking=no',
'-o','ConnectTimeout=10','root@connect.westd.seetacloud.com','-p','23647']
SSH_BASE = ['sshpass','-p','HkM43lFVUIsc','ssh','-o','StrictHostKeyChecking=no',
'-o','ConnectTimeout=10','-p','23647','root@connect.westd.seetacloud.com']
try:
r = subprocess.run(SSH + ['tail','-5','/root/autodl-tmp/train_coder.log'],
capture_output=True, text=True, timeout=25).stdout
sm = re.search(r'(\d+)/(\d+)', r.split(chr(10))[-1] if chr(10) in r[-5:] else r)
lm = re.search(r"'loss':\s*'([\d.]+)'", r)
step = int(sm.group(1)) if sm else 0
total = int(sm.group(2)) if sm else 11838
loss = lm.group(1) if lm else '--'
except:
step, total, loss = 0, 11838, '--'
def fetch_step():
"""从train_coder.log解析最后几步的step/total"""
try:
r = subprocess.run(SSH_BASE + ['tail','-50','/root/autodl-tmp/train_coder.log'],
capture_output=True, text=True, timeout=30).stdout
except:
return 0, 11838
matches = list(re.finditer(r'(\d+)/(\d+)', r))
for m in reversed(matches):
s, t = int(m.group(1)), int(m.group(2))
if t == 11838 and 0 < s <= t:
return s, t
return 0, 11838
data = {'step': step, 'total': total, 'loss': loss,
'updated': datetime.now(timezone.utc).isoformat()+'Z'}
with open(DATA_FILE, 'w') as f:
json.dump(data, f)
print(f"Step {step} / {total}, Loss: {loss}")
def fetch_loss():
"""从trainer_state.json读取最新loss"""
try:
cmd = 'cat /root/autodl-tmp/output/qwen25-coder-7b-sft/checkpoint-*/trainer_state.json'
r = subprocess.run(SSH_BASE + [cmd], capture_output=True, text=True, timeout=15).stdout
state = json.loads(r)
for entry in reversed(state.get('log_history', [])):
if 'loss' in entry:
return f"{entry['loss']:.3f}"
except:
pass
return '--'
def main():
step, total = fetch_step()
loss = fetch_loss()
data = {
'step': step,
'total': total,
'loss': loss,
'updated': datetime.now(timezone.utc).strftime('%Y-%m-%dT%H:%M:%S.000000+00:00Z')
}
with open(DATA_FILE, 'w') as f:
json.dump(data, f)
print(f'Step {step} / {total}, Loss: {loss}')
if __name__ == '__main__':
main()