推理接口返回请求体修改
This commit is contained in:
parent
e98bc40203
commit
3b7a8054c1
|
|
@ -285,7 +285,7 @@ type InferenceSubmitTaskResponse struct {
|
|||
|
||||
// InferenceSubmitTaskResponseData 推理任务提交响应数据
|
||||
type InferenceSubmitTaskResponseData struct {
|
||||
TaskInfo InferenceTaskInfo `json:"taskInfo"` // 任务信息
|
||||
TaskInfo SubtaskRequest `json:"taskInfo"` // 任务信息
|
||||
ResultInfo InferenceResultInfo `json:"resultInfo"` // 结果信息
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -22,9 +22,9 @@ func NewTaskSubmitter(httpClient *client.HTTPClient, taskType string) *TaskSubmi
|
|||
}
|
||||
}
|
||||
|
||||
// SubmitTask 提交任务的通用方法
|
||||
// SubmitTask 提交任务的通用方法,返回 jobSetID、提交请求和错误
|
||||
func (ts *TaskSubmitter) SubmitTask(ctx context.Context, authData *models.AuthData, config *models.RunConfig,
|
||||
clusterID string, bindResultSet *models.BindResultSet) (string, error) {
|
||||
clusterID string, bindResultSet *models.BindResultSet) (string, *models.SubtaskRequest, error) {
|
||||
|
||||
// 构建任务请求
|
||||
submitTaskReq := BuildSubmitTaskRequest(ctx, authData, config, clusterID, bindResultSet, ts.taskType)
|
||||
|
|
@ -39,7 +39,7 @@ func (ts *TaskSubmitter) SubmitTask(ctx context.Context, authData *models.AuthDa
|
|||
resp, err := ts.httpClient.PostJSON("/jsm/v2/jobs/submit", submitTaskReq)
|
||||
if err != nil {
|
||||
fmt.Printf("Submit %s task failed: %v\n", getTaskTypeName(ts.taskType), err)
|
||||
return "", err
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
fmt.Printf("提交%s任务结果:%s\n", getTaskTypeName(ts.taskType), string(resp))
|
||||
|
|
@ -48,16 +48,16 @@ func (ts *TaskSubmitter) SubmitTask(ctx context.Context, authData *models.AuthDa
|
|||
var submitTaskResp models.SubmitTaskResponse
|
||||
if err := json.Unmarshal(resp, &submitTaskResp); err != nil {
|
||||
fmt.Printf("Submit %s task response unmarshal failed: %v\n", getTaskTypeName(ts.taskType), err)
|
||||
return "", err
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
if submitTaskResp.Code != models.ResponseOK {
|
||||
fmt.Printf("Submit %s task failed: %s\n", getTaskTypeName(ts.taskType), submitTaskResp.Code)
|
||||
return "", fmt.Errorf("submit %s task failed: %s", getTaskTypeName(ts.taskType), submitTaskResp.Code)
|
||||
return "", nil, fmt.Errorf("submit %s task failed: %s", getTaskTypeName(ts.taskType), submitTaskResp.Code)
|
||||
}
|
||||
|
||||
fmt.Printf("Submit %s task result: %s\n", getTaskTypeName(ts.taskType), string(resp))
|
||||
return submitTaskResp.Data.JobSetID, nil
|
||||
return submitTaskResp.Data.JobSetID, &submitTaskReq, nil
|
||||
}
|
||||
|
||||
// getTaskTypeName 获取任务类型的中文名称
|
||||
|
|
|
|||
|
|
@ -37,7 +37,8 @@ func NewInferenceService(
|
|||
// SubmitTask 提交推理任务
|
||||
func (s *inferenceService) SubmitTask(ctx context.Context, authData *models.AuthData, config *models.RunConfig, clusterID string,
|
||||
bindResultSet *models.BindResultSet) (string, error) {
|
||||
return s.taskSubmitter.SubmitTask(ctx, authData, config, clusterID, bindResultSet)
|
||||
jobSetID, _, err := s.taskSubmitter.SubmitTask(ctx, authData, config, clusterID, bindResultSet)
|
||||
return jobSetID, err
|
||||
}
|
||||
|
||||
// GetTaskStatus 查询任务状态
|
||||
|
|
@ -105,22 +106,10 @@ func (s *inferenceService) SubmitInferenceTask(ctx context.Context, authData *mo
|
|||
fmt.Println("增量模型准备成功,ID:", subModelID)
|
||||
}
|
||||
|
||||
// 4. 构建推理任务响应数据(用于返回)
|
||||
taskData := common.BuildInferenceSubmitTaskRequest(
|
||||
ctx, authData,
|
||||
codeID, modelID, subModelID,
|
||||
request.Image.ImageID,
|
||||
clusterID,
|
||||
request.ResourceType,
|
||||
request.Resource.Resources,
|
||||
request.Command,
|
||||
request.Description,
|
||||
)
|
||||
|
||||
// 5. 使用 taskSubmitter 提交任务
|
||||
// 4. 使用 taskSubmitter 提交任务
|
||||
// 注意:taskSubmitter 会使用 BuildSubmitTaskRequest,它构建的是训练任务格式
|
||||
// 对于推理任务,我们需要传入 BindDatasetID=0,这样 Files 中的 Dataset 会是空的
|
||||
jobSetID, err := s.taskSubmitter.SubmitTask(ctx, authData, runConfig, clusterID, &models.BindResultSet{
|
||||
jobSetID, submitTaskReq, err := s.taskSubmitter.SubmitTask(ctx, authData, runConfig, clusterID, &models.BindResultSet{
|
||||
BindCodeID: codeID,
|
||||
BindModelID: modelID,
|
||||
BindDatasetID: 0, // 推理任务不需要数据集
|
||||
|
|
@ -129,8 +118,14 @@ func (s *inferenceService) SubmitInferenceTask(ctx context.Context, authData *mo
|
|||
return nil, fmt.Errorf("提交任务失败: %w", err)
|
||||
}
|
||||
|
||||
// 更新结果信息
|
||||
taskData.ResultInfo.JobSetID = jobSetID
|
||||
// 5. 构建返回响应数据,使用提交请求中的信息
|
||||
taskData := models.InferenceSubmitTaskResponseData{
|
||||
TaskInfo: *submitTaskReq,
|
||||
ResultInfo: models.InferenceResultInfo{
|
||||
LocalJobID: models.MainTaskID,
|
||||
JobSetID: jobSetID,
|
||||
},
|
||||
}
|
||||
|
||||
return &models.InferenceSubmitTaskResponse{
|
||||
Code: 200,
|
||||
|
|
|
|||
|
|
@ -32,7 +32,8 @@ func NewTaskService(baseURL string, timeout time.Duration) service.TaskService {
|
|||
|
||||
func (s *taskService) SubmitTask(ctx context.Context, authData *models.AuthData, config *models.RunConfig,
|
||||
clusterID string, bindResultSet *models.BindResultSet) (string, error) {
|
||||
return s.taskSubmitter.SubmitTask(ctx, authData, config, clusterID, bindResultSet)
|
||||
jobSetID, _, err := s.taskSubmitter.SubmitTask(ctx, authData, config, clusterID, bindResultSet)
|
||||
return jobSetID, err
|
||||
}
|
||||
|
||||
func (s *taskService) GetTaskStatus(ctx context.Context, authData *models.AuthData, jobSetID string) (*models.TaskDetailResponse, error) {
|
||||
|
|
|
|||
Loading…
Reference in New Issue