diff --git a/server/biz/service/training.go b/server/biz/service/training.go index b8b0d66..9a62568 100644 --- a/server/biz/service/training.go +++ b/server/biz/service/training.go @@ -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.json:error 字段 = 脚本异常退出 → 置失败带出原因;否则按成功收尾 +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) { diff --git a/server/common/training_runner.go b/server/common/training_runner.go index df6aa76..1c0c24e 100644 --- a/server/common/training_runner.go +++ b/server/common/training_runner.go @@ -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 } diff --git a/server/data/observer.db b/server/data/observer.db index 3e1c35f..2069632 100644 Binary files a/server/data/observer.db and b/server/data/observer.db differ diff --git a/server/training/train_server.py b/server/training/train_server.py index 84c89ec..d44c028 100644 --- a/server/training/train_server.py +++ b/server/training/train_server.py @@ -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, diff --git a/server/workspace/trainings/野鸡.tflite b/server/workspace/trainings/野鸡.tflite new file mode 100644 index 0000000..771fba8 Binary files /dev/null and b/server/workspace/trainings/野鸡.tflite differ