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

Function constructOptimizers

go/codegen/tensorflow/codegen.go:107–150  ·  view source on GitHub ↗

constructOptimizers generates a python optimizer object using: model.optimizer = "OptimizerName" optimizer.arg1 = 1 optimizer.arg2 = "2" To: model.optimizer = "OptimizerName(arg1=1, arg2=\"2\")"

(trainStmt *ir.TrainStmt)

Source from the content-addressed store, hash-verified

105// To:
106// model.optimizer = "OptimizerName(arg1=1, arg2=\"2\")"
107func constructOptimizers(trainStmt *ir.TrainStmt) {
108 optimizerArgs := map[string]map[string]interface{}{}
109 for k, v := range trainStmt.Attributes {
110 if attrIsOptimizer(k) {
111 if optimizerArgs[k] == nil {
112 optimizerArgs[k] = map[string]interface{}{}
113 }
114 }
115 pieces := strings.Split(k, ".")
116 if len(pieces) == 2 {
117 if attrIsOptimizer("model." + pieces[0]) { // k is like "optimizer.learning_rate"
118 if optimizerArgs["model."+pieces[0]] == nil {
119 optimizerArgs["model."+pieces[0]] = map[string]interface{}{}
120 }
121 optimizerArgs["model."+pieces[0]][pieces[1]] = v
122 // delete these attributes because they are only used to initialized the python object
123 delete(trainStmt.Attributes, k)
124 }
125 }
126 }
127 tf1OptimizerClsNames := map[string]string{
128 "Adagrad": "tf.train.AdagradOptimizer",
129 "Adam": "tf.train.AdamOptimizer",
130 "Ftrl": "tf.train.FtrlOptimizer",
131 "RMSProp": "tf.train.RMSPropOptimizer",
132 "SGD": "tf.train.GradientDescentOptimizer",
133 }
134
135 for optimizerParamName, args := range optimizerArgs {
136 if _, ok := trainStmt.Attributes[optimizerParamName]; !ok {
137 setDefaultOptimizer(trainStmt, optimizerParamName)
138 }
139 optimizerCls := fmt.Sprintf("%v", trainStmt.Attributes[optimizerParamName])
140 if cls, ok := tf1OptimizerClsNames[optimizerCls]; ok && IsPAI() {
141 optimizerCls = cls
142 }
143 optimizerInitPyCode := fmt.Sprintf("%s(", optimizerCls)
144 for k, v := range args {
145 optimizerInitPyCode += fmt.Sprintf("%s=%v, ", k, v)
146 }
147 optimizerInitPyCode += ")"
148 trainStmt.Attributes[optimizerParamName] = optimizerInitPyCode
149 }
150}
151
152// constructLosses generate a python loss function call using:
153// model.loss = "LossName"

Callers 1

InitializeAttributesFunction · 0.85

Calls 3

attrIsOptimizerFunction · 0.85
setDefaultOptimizerFunction · 0.85
IsPAIFunction · 0.85

Tested by

no test coverage detected