MCPcopy Create free account
hub / github.com/sql-machine-learning/sqlflow / Evaluate

Function Evaluate

go/codegen/tensorflow/codegen.go:394–432  ·  view source on GitHub ↗

Evaluate generates a Python program to evaluate a trained model.

(stmt *ir.EvaluateStmt, session *pb.Session)

Source from the content-addressed store, hash-verified

392
393// Evaluate generates a Python program to evaluate a trained model.
394func 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
435func restoreModel(stmt *ir.TrainStmt) (modelParams map[string]interface{}, featureColumnsCode []string, fieldDescs map[string][]*ir.FieldDesc, err error) {

Callers 1

ExecuteEvaluateMethod · 0.92

Calls 5

restoreModelFunction · 0.85
resolveParamsFunction · 0.70
GetFieldDescMethod · 0.65
ParseMethod · 0.65
StringMethod · 0.45

Tested by

no test coverage detected