guanghulab/homepage/fetch_train.py

58 lines
1.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/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_BASE = ['sshpass','-p','HkM43lFVUIsc','ssh','-o','StrictHostKeyChecking=no',
'-o','ConnectTimeout=10','-p','23647','root@connect.westd.seetacloud.com']
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
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()