Evaluate generates a Python program to evaluate a trained model.
(stmt *ir.EvaluateStmt, session *pb.Session)
| 392 | |
| 393 | // Evaluate generates a Python program to evaluate a trained model. |
| 394 | func Evaluate(stmt *ir.EvaluateStmt, session *pb.Session) (string, error) { |
| 395 | modelParams, featureColumnsCode, fieldDescs, err := restoreModel(stmt.TrainStmt) |
| 396 | if err != nil { |
| 397 | return "", err |
| 398 | } |
| 399 | labelFM := stmt.TrainStmt.Label.GetFieldDesc()[0] |
| 400 | validationParams := resolveParams(stmt.Attributes, "validation.") |
| 401 | if len(validationParams) == 0 { |
| 402 | // add default validation.metrics = "Accuracy". |
| 403 | validationParams["metrics"] = "Accuracy" |
| 404 | } |
| 405 | |
| 406 | filler := evaluateFiller{ |
| 407 | DataSource: session.DbConnStr, |
| 408 | Select: stmt.Select, |
| 409 | Estimator: stmt.TrainStmt.Estimator, |
| 410 | FieldDescs: fieldDescs, |
| 411 | FeatureColumnCode: fmt.Sprintf("{%s}", strings.Join(featureColumnsCode, ",\n")), |
| 412 | Y: labelFM, |
| 413 | ModelParams: modelParams, |
| 414 | ValidationParams: validationParams, |
| 415 | Save: "model_save", |
| 416 | ResultTable: stmt.Into, |
| 417 | HDFSNameNodeAddr: session.HdfsNamenodeAddr, |
| 418 | HiveLocation: session.HiveLocation, |
| 419 | HDFSUser: session.HdfsUser, |
| 420 | HDFSPass: session.HdfsPass, |
| 421 | } |
| 422 | var program bytes.Buffer |
| 423 | var tmpl = template.Must(template.New("Evaluate").Funcs(template.FuncMap{ |
| 424 | "intArrayToJSONString": ir.MarshalToJSONString, |
| 425 | "attrToPythonValue": ir.AttrToPythonValue, |
| 426 | "DTypeToString": ir.DTypeToString, |
| 427 | }).Parse(tfEvaluateTemplateText)) |
| 428 | if err := tmpl.Execute(&program, filler); err != nil { |
| 429 | return "", err |
| 430 | } |
| 431 | return program.String(), nil |
| 432 | } |
| 433 | |
| 434 | // restoreModel reconstruct necessary python objects from TrainStmt |
| 435 | func restoreModel(stmt *ir.TrainStmt) (modelParams map[string]interface{}, featureColumnsCode []string, fieldDescs map[string][]*ir.FieldDesc, err error) { |
no test coverage detected