MCPcopy Create free account
hub / github.com/CodingBeard/tfkg / LoadModel

Function LoadModel

model/model.go:110–168  ·  view source on GitHub ↗
(
	errorHandler *cberrors.ErrorsContainer,
	logger *cblog.Logger,
	dir string,
	sessionOptions ...*for_core_protos_go_proto.ConfigProto,
)

Source from the content-addressed store, hash-verified

108}
109
110func 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

Callers 5

mainFunction · 0.92
mainFunction · 0.92
mainFunction · 0.92
mainFunction · 0.92
mainFunction · 0.92

Calls 1

ErrorMethod · 0.65

Tested by

no test coverage detected