From b88876379a9ffdcbf4c454b46eefec181a026628 Mon Sep 17 00:00:00 2001 From: liuzx Date: Tue, 28 Dec 2021 15:54:33 +0800 Subject: [PATCH] update --- models/cloudbrain.go | 15 ++++++++++-- routers/api/v1/api.go | 5 +++- routers/api/v1/repo/modelarts.go | 41 ++++++++++++++++++++++++++++++++ routers/repo/modelarts.go | 3 +++ 4 files changed, 61 insertions(+), 3 deletions(-) diff --git a/models/cloudbrain.go b/models/cloudbrain.go index a19af4df3..c90e31724 100755 --- a/models/cloudbrain.go +++ b/models/cloudbrain.go @@ -19,8 +19,8 @@ type JobType string type ModelArtsJobStatus string const ( - NPUResource = "NPU" - GPUResource = "CPU/GPU" + NPUResource = "NPU" + GPUResource = "CPU/GPU" JobWaiting CloudbrainStatus = "WAITING" JobStopped CloudbrainStatus = "STOPPED" @@ -1207,3 +1207,14 @@ func GetCloudbrainTrainJobCountByUserID(userID int64) (int, error) { And("job_type = ? and user_id = ? and type = ?", JobTypeTrain, userID, TypeCloudBrainTwo).Count(new(Cloudbrain)) return int(count), err } + +func UpdateInferenceJob(job *Cloudbrain) error { + return updateJobInferenceJob(x, job) +} + +func updateJobInferenceJob(e Engine, job *Cloudbrain) error { + var sess *xorm.Session + sess = e.Where("job_id = ?", job.JobID) + _, err := sess.Cols("status", "train_job_duration").Update(job) + return err +} diff --git a/routers/api/v1/api.go b/routers/api/v1/api.go index 518c63e4f..868fa23c4 100755 --- a/routers/api/v1/api.go +++ b/routers/api/v1/api.go @@ -524,7 +524,7 @@ func RegisterRoutes(m *macaron.Macaron) { Get(notify.GetThread). Patch(notify.ReadThread) }, reqToken()) - + operationReq := context.Toggle(&context.ToggleOptions{SignInRequired: true, OperationRequired: true}) //Project board m.Group("/projectboard", func() { @@ -886,6 +886,9 @@ func RegisterRoutes(m *macaron.Macaron) { m.Get("/model_list", repo.ModelList) }) }) + m.Group("/inference-job", func() { + m.Get("/:jobid", repo.GetModelArtsInferenceJob) + }) }, reqRepoReader(models.UnitTypeCloudBrain)) }, repoAssignment()) }) diff --git a/routers/api/v1/repo/modelarts.go b/routers/api/v1/repo/modelarts.go index 676097465..2b6e45461 100755 --- a/routers/api/v1/repo/modelarts.go +++ b/routers/api/v1/repo/modelarts.go @@ -343,3 +343,44 @@ func deleteJobStorage(jobName string) error { return nil } + +func GetModelArtsInferenceJob(ctx *context.APIContext) { + var ( + err error + ) + + jobID := ctx.Params(":jobid") + job, err := models.GetCloudbrainByJobID(jobID) + if err != nil { + ctx.NotFound(err) + return + } + result, err := modelarts.GetTrainJob(jobID, strconv.FormatInt(job.VersionID, 10)) + if err != nil { + ctx.NotFound(err) + return + } + + job.Status = modelarts.TransTrainJobStatus(result.IntStatus) + job.Duration = result.Duration + job.TrainJobDuration = result.TrainJobDuration + + if result.Duration != 0 { + job.TrainJobDuration = addZero(result.Duration/3600000) + ":" + addZero(result.Duration%3600000/60000) + ":" + addZero(result.Duration%60000/1000) + + } else { + job.TrainJobDuration = "00:00:00" + } + + err = models.UpdateInferenceJob(job) + if err != nil { + log.Error("UpdateJob failed:", err) + } + + ctx.JSON(http.StatusOK, map[string]interface{}{ + "JobID": jobID, + "JobStatus": job.Status, + "JobDuration": job.TrainJobDuration, + }) + +} diff --git a/routers/repo/modelarts.go b/routers/repo/modelarts.go index 61a60d8d0..d370bd5d5 100755 --- a/routers/repo/modelarts.go +++ b/routers/repo/modelarts.go @@ -1700,6 +1700,8 @@ func InferenceJobCreate(ctx *context.Context, form auth.CreateModelArtsInference EngineName: EngineName, ModelName: modelName, ModelVersion: modelVersion, + CkptName: ckptName, + ResultUrl: resultObsPath, } //将params转换Parameters.Parameter,出错时返回给前端 @@ -1744,6 +1746,7 @@ func InferenceJobIndex(ctx *context.Context) { for i, task := range tasks { tasks[i].CanDel = cloudbrain.CanDeleteJob(ctx, &task.Cloudbrain) tasks[i].CanModify = cloudbrain.CanModifyJob(ctx, &task.Cloudbrain) + tasks[i].ComputeResource = models.NPUResource } pager := context.NewPagination(int(count), setting.UI.IssuePagingNum, page, 5)