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

Function Predict

go/codegen/pai/codegen.go:143–201  ·  view source on GitHub ↗

Predict generates a Python program for train a TensorFlow model.

(ir *ir.PredictStmt, session *pb.Session, tarball, paramsFile, modelName, ossModelPath, cwd string, modelType int)

Source from the content-addressed store, hash-verified

141
142// Predict generates a Python program for train a TensorFlow model.
143func Predict(ir *ir.PredictStmt, session *pb.Session, tarball, paramsFile, modelName, ossModelPath, cwd string, modelType int) (code, paiCmd, requirements string, e error) {
144 currProject := ""
145 currProject, e = database.GetDatabaseName(session.DbConnStr)
146 if e != nil {
147 return
148 }
149 if modelType == model.PAIML {
150 if paiCmd, e = getPAIPredictCmd(ir, session); e != nil {
151 return
152 }
153 } else if modelType == model.XGBOOST {
154 requirements, e = genRequirements(true)
155 ossURI := OSSModelURL(ossModelPath)
156 var xgbPredCode bytes.Buffer
157 var tpl = template.Must(template.New("xgbPredTemplate").Parse(xgbPredTemplateText))
158 paiPredictTable := ""
159 if tensorflow.IsPAI() && ir.TmpPredictTable != "" {
160 paiPredictTable = ir.TmpPredictTable
161 }
162 filler := &xgbPredictFiller{
163 OSSModelDir: ossURI,
164 DataSource: session.DbConnStr,
165 PredSelect: ir.Select,
166 ResultTable: ir.ResultTable,
167 ResultColumn: ir.ResultColumn,
168 HDFSNameNodeAddr: session.HdfsNamenodeAddr,
169 HiveLocation: session.HiveLocation,
170 HDFSUser: session.HdfsUser,
171 HDFSPass: session.HdfsPass,
172 PAIPredictTable: paiPredictTable,
173 }
174 if e = tpl.Execute(&xgbPredCode, filler); e != nil {
175 return
176 }
177 code = xgbPredCode.String()
178
179 cc, err := GetClusterConfig(ir.Attributes)
180 if err != nil {
181 return
182 }
183 // NOTE(typhoonzero): submit a PAI TF job to install xgboost and run.
184 if paiCmd, e = getTFPAICmd(cc, tarball, paramsFile, modelName, ossModelPath, ir.TmpPredictTable, "", ir.ResultTable, currProject, cwd); e != nil {
185 return
186 }
187 } else {
188 requirements, e = genRequirements(false)
189 cc, err := GetClusterConfig(ir.Attributes)
190 if err != nil {
191 return
192 }
193 if code, e = TFLoadAndPredict(ir, session, ossModelPath); e != nil {
194 return
195 }
196 if paiCmd, e = getTFPAICmd(cc, tarball, paramsFile, modelName, ossModelPath, ir.TmpPredictTable, "", ir.ResultTable, currProject, cwd); e != nil {
197 return
198 }
199 }
200 return

Callers 3

ExecutePredictMethod · 0.92
ExecutePredictMethod · 0.92
TestPredictCodegenFunction · 0.85

Calls 10

GetDatabaseNameFunction · 0.92
IsPAIFunction · 0.92
getPAIPredictCmdFunction · 0.85
genRequirementsFunction · 0.85
OSSModelURLFunction · 0.85
GetClusterConfigFunction · 0.85
getTFPAICmdFunction · 0.85
TFLoadAndPredictFunction · 0.85
ParseMethod · 0.65
StringMethod · 0.45

Tested by 1

TestPredictCodegenFunction · 0.68