1
This commit is contained in:
@@ -84,15 +84,7 @@ func (s *trainingService) pollOne(ctx context.Context, runner common.TrainingRun
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
if result != "" {
|
if result != "" {
|
||||||
// result.json 带 error 字段 = 脚本异常退出(如 tflite 导出失败),置失败并带出原因
|
s.handleResult(ctx, runner, job, t, result, tail)
|
||||||
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)
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
alive := true
|
alive := true
|
||||||
@@ -117,6 +109,13 @@ func (s *trainingService) pollOne(ctx context.Context, runner common.TrainingRun
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
if !alive {
|
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, "训练进程已退出(无结果文件)")
|
_ = s.finishFailed(ctx, t, "训练进程已退出(无结果文件)")
|
||||||
return
|
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 → 更新任务。
|
// finishSuccess 训练成功:解析 result.json(最终指标 + 类别名)→ 拉取 tflite → 更新任务。
|
||||||
// tflite 直写 trainings/<数据集名>.tflite(当前生效模型唯一位,无 per-task 存档、无 zip)。
|
// 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) {
|
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))
|
out, err := r.runCmd(ctx, job, fmt.Sprintf("kill -0 %d 2>/dev/null && echo alive || echo dead", pid))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return false, nil // 连接故障按"状态未知"处理,不误判任务结束
|
return false, err // 连接故障 = 状态未知,交上层处理,不吞错
|
||||||
}
|
}
|
||||||
return strings.TrimSpace(out) == "alive", nil
|
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) {
|
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 -f 区分「结果文件未就绪」(("", nil))与 SSH 连接/命令真实失败(("", err)):
|
||||||
if err != nil || out == "" {
|
// 原 cat 2>/dev/null 两者混为一谈,连接故障会被当成"文件不存在=仍在跑"而吞掉错误
|
||||||
return "", nil // 文件不存在 = 仍在跑
|
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
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|||||||
Binary file not shown.
@@ -55,11 +55,11 @@ def on_fit_epoch_end(trainer):
|
|||||||
}, ensure_ascii=False) + "\n")
|
}, ensure_ascii=False) + "\n")
|
||||||
|
|
||||||
|
|
||||||
def register_callback():
|
def register_callback(model):
|
||||||
# ultralytics ≥8.4 中 callbacks 字典改由 get_default_callbacks() 取得
|
# ultralytics ≥8.4 中 trainer 使用 model 实例的 callbacks 字典(model.train 前以
|
||||||
from ultralytics.utils.callbacks import get_default_callbacks
|
# _callbacks=self.callbacks 传入),模块级 get_default_callbacks() 的 append 不生效;
|
||||||
callbacks = get_default_callbacks()
|
# add_callback 追加到 model 字典即随训练触发,进度 jsonl 才有输出
|
||||||
callbacks["on_fit_epoch_end"].append(on_fit_epoch_end)
|
model.add_callback("on_fit_epoch_end", on_fit_epoch_end)
|
||||||
|
|
||||||
|
|
||||||
def write_result(result_file, payload):
|
def write_result(result_file, payload):
|
||||||
@@ -245,8 +245,8 @@ def main():
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
from ultralytics import YOLO
|
from ultralytics import YOLO
|
||||||
register_callback()
|
|
||||||
model = YOLO("yolov8n.pt")
|
model = YOLO("yolov8n.pt")
|
||||||
|
register_callback(model)
|
||||||
model.train(
|
model.train(
|
||||||
data=data, imgsz=imgsz, epochs=epochs,
|
data=data, imgsz=imgsz, epochs=epochs,
|
||||||
patience=PATIENCE, batch=batch, device=device, workers=4,
|
patience=PATIENCE, batch=batch, device=device, workers=4,
|
||||||
|
|||||||
Binary file not shown.
Reference in New Issue
Block a user