1
This commit is contained in:
@@ -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) {
|
||||
|
||||
@@ -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.
@@ -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.
Reference in New Issue
Block a user