This commit is contained in:
2026-09-09 21:05:25 +08:00
parent 87ab06221c
commit c590f74b1e
14 changed files with 229 additions and 208 deletions
+17 -4
View File
@@ -8,7 +8,9 @@
任务参数(Go 侧写入,字段相对训练机 workdir):
workdir 训练机工作目录(脚本 / yolov8s.pt / venv 所在),启动即 chdir
yolo 训练集目录(含 dataset.yaml),相对 workdir
model 训练基座权重文件名(workdir 下,默认 yolov8s.pt
model 训练基座权重文件名(workdir 下,默认 yolov8s.pt;增量训练时为 Go 侧推到
trainings/weights/<基名>.pt 的上次成功权重,类别数不符时 ultralytics
自动重建检测头、backbone 热启动)
imgsz 训练/导出分辨率(默认 1280,与端侧推理对齐)
epochs 训练轮数
batch 批大小
@@ -19,10 +21,13 @@
产物契约:
log_file {"epoch":1,"total":150,"metrics":{"metrics/mAP50(B)":0.87,...}}
result_file {"metrics":{...},"names":["pheasant","suspect"],"best_tflite":"runs/tasks/<id>/weights/best.tflite",
result_file {"metrics":{...},"names":["pheasant","suspect"],
"best_tflite":"runs/tasks/<id>/train/weights/best.tflite",
"best_pt":"runs/tasks/<id>/train/weights/best.pt",
"tflite_check":{"ok":true,"reason":"","inputs":[...],"outputs":[...]}}
result_file 存在 = 训练完成;异常时写 {"error":"..."}Go 侧据以置失败并展示原因。
tflite_check.ok=false(产物 shape 异常)时 Go 侧置训练失败并带出 reason。
best_pt 供 Go 侧存档 workspace/trainings/weights/ 作下次增量基座(2026-09-09)。
"""
import argparse
@@ -251,7 +256,12 @@ def main():
try:
from ultralytics import YOLO
model = YOLO(task.get("model") or "yolov8s.pt")
# 基座权重相对 workdir(增量训练时为 Go 推送的 trainings/weights/<基名>.pt);
# 仅在 workdir 下确有该文件时绝对化,裸名(yolov8s.pt)缺文件仍走 ultralytics 自下载
model_path = task.get("model") or "yolov8s.pt"
if not os.path.isabs(model_path) and os.path.exists(os.path.join(base, model_path)):
model_path = os.path.join(base, model_path)
model = YOLO(model_path)
register_callback(model)
if is_mps:
def _mps_cache(_trainer):
@@ -321,11 +331,14 @@ def main():
if isinstance(v, (int, float)) and math.isfinite(v):
clean_metrics[k] = round(float(v), 5)
# best_tflite 相对 workdirGo 侧按此路径拉取(不再打包 zip,仅回传 tflite
# best_tflite/best_pt 相对 workdir,Go 侧按此路径拉取(不再打包 zip,仅回传 tflite/pt);
# best_pt 存档 trainings/weights/ 作下次增量基座(不存在不致命,缺省空串)
best_pt = save_dir / "weights" / "best.pt"
write_result(result_file, {
"metrics": clean_metrics,
"names": names,
"best_tflite": os.path.relpath(best_tflite, base),
"best_pt": os.path.relpath(best_pt, base) if best_pt.exists() else "",
"tflite_check": tflite_check,
})
except Exception: