This commit is contained in:
2026-08-27 11:14:38 +08:00
parent bb7bac3dee
commit 589b79c8cd
5 changed files with 36 additions and 19 deletions
+20 -9
View File
@@ -84,15 +84,7 @@ func (s *trainingService) pollOne(ctx context.Context, runner common.TrainingRun
return
}
if result != "" {
// result.json 带 error 字段 = 脚本异常退出(如 tflite 导出失败),置失败并带出原因
var resErr struct {
Error string `json:"error"`
}
if json.Unmarshal([]byte(result), &resErr) == nil && resErr.Error != "" {
_ = s.finishFailed(ctx, t, "训练失败: %s", resErr.Error)
return
}
s.finishSuccess(ctx, runner, job, t, result, tail)
s.handleResult(ctx, runner, job, t, result, tail)
return
}
alive := true
@@ -117,6 +109,13 @@ func (s *trainingService) pollOne(ctx context.Context, runner common.TrainingRun
return
}
if !alive {
// 竞态防护:脚本原子写结果文件后进程随即退出,轮询可能命中「结果未就绪 + 进程已死」窗口;
// 判死前延迟重试一次,确认文件确实缺席(结果文件迟到属正常时序,非训练失败)
time.Sleep(2 * time.Second)
if r2, e2 := runner.FetchResult(ctx, job); e2 == nil && r2 != "" {
s.handleResult(ctx, runner, job, t, r2, tail)
return
}
_ = s.finishFailed(ctx, t, "训练进程已退出(无结果文件)")
return
}
@@ -127,6 +126,18 @@ func (s *trainingService) pollOne(ctx context.Context, runner common.TrainingRun
}
}
// handleResult 处理已就绪的 result.jsonerror 字段 = 脚本异常退出 → 置失败带出原因;否则按成功收尾
func (s *trainingService) handleResult(ctx context.Context, runner common.TrainingRunner, job *common.TrainingJob, t *entity.ModelTraining, result, tail string) {
var resErr struct {
Error string `json:"error"`
}
if json.Unmarshal([]byte(result), &resErr) == nil && resErr.Error != "" {
_ = s.finishFailed(ctx, t, "训练失败: %s", resErr.Error)
return
}
s.finishSuccess(ctx, runner, job, t, result, tail)
}
// finishSuccess 训练成功:解析 result.json(最终指标 + 类别名)→ 拉取 tflite → 更新任务。
// tflite 直写 trainings/<数据集名>.tflite(当前生效模型唯一位,无 per-task 存档、无 zip)。
func (s *trainingService) finishSuccess(ctx context.Context, runner common.TrainingRunner, job *common.TrainingJob, t *entity.ModelTraining, result, tail string) {
+10 -4
View File
@@ -309,7 +309,7 @@ func (r *sshRunner) IsAlive(ctx context.Context, job *TrainingJob) (bool, error)
}
out, err := r.runCmd(ctx, job, fmt.Sprintf("kill -0 %d 2>/dev/null && echo alive || echo dead", pid))
if err != nil {
return false, nil // 连接故障按"状态未知"处理,不误判任务结束
return false, err // 连接故障 = 状态未知,交上层处理,不吞错
}
return strings.TrimSpace(out) == "alive", nil
}
@@ -330,9 +330,15 @@ func (r *sshRunner) FetchLogTail(ctx context.Context, job *TrainingJob) (string,
}
func (r *sshRunner) FetchResult(ctx context.Context, job *TrainingJob) (string, error) {
out, err := r.runCmd(ctx, job, fmt.Sprintf("cat %s/results/%d.json 2>/dev/null", job.Workdir, job.TaskId))
if err != nil || out == "" {
return "", nil // 文件不存在 = 仍在跑
// if -f 区分「结果文件未就绪」(("", nil))与 SSH 连接/命令真实失败(("", err)):
// 原 cat 2>/dev/null 两者混为一谈,连接故障会被当成"文件不存在=仍在跑"而吞掉错误
path := filepath.Join(job.Workdir, "results", fmt.Sprintf("%d.json", job.TaskId))
out, err := r.runCmd(ctx, job, fmt.Sprintf("if [ -f %s ]; then cat %s; else echo __NO_RESULT__; fi", path, path))
if err != nil {
return "", err
}
if strings.TrimSpace(out) == "__NO_RESULT__" {
return "", nil
}
return out, nil
}
Binary file not shown.
+6 -6
View File
@@ -55,11 +55,11 @@ def on_fit_epoch_end(trainer):
}, ensure_ascii=False) + "\n")
def register_callback():
# ultralytics ≥8.4 中 callbacks 字典改由 get_default_callbacks() 取得
from ultralytics.utils.callbacks import get_default_callbacks
callbacks = get_default_callbacks()
callbacks["on_fit_epoch_end"].append(on_fit_epoch_end)
def register_callback(model):
# ultralytics ≥8.4 中 trainer 使用 model 实例的 callbacks 字典(model.train 前以
# _callbacks=self.callbacks 传入),模块级 get_default_callbacks() 的 append 不生效;
# add_callback 追加到 model 字典即随训练触发,进度 jsonl 才有输出
model.add_callback("on_fit_epoch_end", on_fit_epoch_end)
def write_result(result_file, payload):
@@ -245,8 +245,8 @@ def main():
try:
from ultralytics import YOLO
register_callback()
model = YOLO("yolov8n.pt")
register_callback(model)
model.train(
data=data, imgsz=imgsz, epochs=epochs,
patience=PATIENCE, batch=batch, device=device, workers=4,
Binary file not shown.