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)
| 105 | // To: |
| 106 | // model.optimizer = "OptimizerName(arg1=1, arg2=\"2\")" |
| 107 | func 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" |
no test coverage detected