| 108 | } |
| 109 | |
| 110 | func LoadModel( |
| 111 | errorHandler *cberrors.ErrorsContainer, |
| 112 | logger *cblog.Logger, |
| 113 | dir string, |
| 114 | sessionOptions ...*for_core_protos_go_proto.ConfigProto, |
| 115 | ) (*TfkgModel, error) { |
| 116 | var tfConfig *for_core_protos_go_proto.ConfigProto |
| 117 | if len(sessionOptions) == 1 { |
| 118 | tfConfig = sessionOptions[0] |
| 119 | } else { |
| 120 | tfConfig = &for_core_protos_go_proto.ConfigProto{} |
| 121 | } |
| 122 | tfConfigBytes, e := proto.Marshal(tfConfig) |
| 123 | if e != nil { |
| 124 | errorHandler.Error(e) |
| 125 | return nil, e |
| 126 | } |
| 127 | m, e := tf.LoadSavedModel(dir, []string{"serve"}, &tf.SessionOptions{ |
| 128 | Config: tfConfigBytes, |
| 129 | }) |
| 130 | if e != nil { |
| 131 | errorHandler.Error(e) |
| 132 | return nil, e |
| 133 | } |
| 134 | |
| 135 | tfkgSignatures := []string{ |
| 136 | "learn", |
| 137 | "evaluate", |
| 138 | "predict", |
| 139 | } |
| 140 | |
| 141 | found := 0 |
| 142 | for signatureName := range m.Signatures { |
| 143 | for _, tfkgSignature := range tfkgSignatures { |
| 144 | if signatureName == tfkgSignature { |
| 145 | found++ |
| 146 | } |
| 147 | } |
| 148 | } |
| 149 | if found != 3 { |
| 150 | e = fmt.Errorf("%s is not a TFKG model. Use model.LoadVanillaModel instead", dir) |
| 151 | errorHandler.Error(e) |
| 152 | return nil, e |
| 153 | } |
| 154 | |
| 155 | pbCache, e := ioutil.ReadFile(filepath.Join(dir, "saved_model.pb")) |
| 156 | if e != nil { |
| 157 | errorHandler.Error(e) |
| 158 | return nil, e |
| 159 | } |
| 160 | |
| 161 | return &TfkgModel{ |
| 162 | model: m, |
| 163 | pbCache: pbCache, |
| 164 | errorHandler: errorHandler, |
| 165 | logger: logger, |
| 166 | modelDefinitionSaveDir: dir, |
| 167 | }, nil |