(cl *ir.EvaluateStmt)
| 307 | } |
| 308 | |
| 309 | func (s *pythonExecutor) ExecuteEvaluate(cl *ir.EvaluateStmt) error { |
| 310 | // NOTE(typhoonzero): model is already loaded under s.Cwd |
| 311 | var code string |
| 312 | var err error |
| 313 | if cl.TrainStmt.GetModelKind() == ir.XGBoost { |
| 314 | code, err = xgboost.Evaluate(cl, s.Session) |
| 315 | if err != nil { |
| 316 | return err |
| 317 | } |
| 318 | } else { |
| 319 | code, err = tensorflow.Evaluate(cl, s.Session) |
| 320 | if err != nil { |
| 321 | return err |
| 322 | } |
| 323 | } |
| 324 | |
| 325 | if cl.Into != "" { |
| 326 | // create evaluation result table |
| 327 | db, err := database.OpenAndConnectDB(s.Session.DbConnStr) |
| 328 | if err != nil { |
| 329 | return err |
| 330 | } |
| 331 | defer db.Close() |
| 332 | // default always output evaluation loss |
| 333 | metricNames := []string{"loss"} |
| 334 | metricsAttr, ok := cl.Attributes["validation.metrics"] |
| 335 | if ok { |
| 336 | metricsList := strings.Split(metricsAttr.(string), ",") |
| 337 | metricNames = append(metricNames, metricsList...) |
| 338 | } |
| 339 | err = createEvaluationResultTable(db, cl.Into, metricNames) |
| 340 | if err != nil { |
| 341 | return err |
| 342 | } |
| 343 | } |
| 344 | if err = s.runProgram(code, false); err != nil { |
| 345 | return err |
| 346 | } |
| 347 | return nil |
| 348 | } |
| 349 | |
| 350 | func (s *pythonExecutor) ExecuteOptimize(stmt *ir.OptimizeStmt) error { |
| 351 | db, err := database.OpenAndConnectDB(s.Session.DbConnStr) |
nothing calls this directly
no test coverage detected