MCPcopy Create free account
hub / github.com/llSourcell/YOLO_Object_Detection / TFNet

Class TFNet

darkflow/net/build.py:12–177  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

10import os
11
12class TFNet(object):
13
14 _TRAINER = dict({
15 'rmsprop': tf.train.RMSPropOptimizer,
16 'adadelta': tf.train.AdadeltaOptimizer,
17 'adagrad': tf.train.AdagradOptimizer,
18 'adagradDA': tf.train.AdagradDAOptimizer,
19 'momentum': tf.train.MomentumOptimizer,
20 'adam': tf.train.AdamOptimizer,
21 'ftrl': tf.train.FtrlOptimizer,
22 'sgd': tf.train.GradientDescentOptimizer
23 })
24
25 # imported methods
26 _get_fps = help._get_fps
27 say = help.say
28 train = flow.train
29 camera = help.camera
30 predict = flow.predict
31 return_predict = flow.return_predict
32 to_darknet = help.to_darknet
33 build_train_op = help.build_train_op
34 load_from_ckpt = help.load_from_ckpt
35
36 def __init__(self, FLAGS, darknet = None):
37 self.ntrain = 0
38
39 if isinstance(FLAGS, dict):
40 from ..defaults import argHandler
41 newFLAGS = argHandler()
42 newFLAGS.setDefaults()
43 newFLAGS.update(FLAGS)
44 FLAGS = newFLAGS
45
46 self.FLAGS = FLAGS
47 if self.FLAGS.pbLoad and self.FLAGS.metaLoad:
48 self.say('\nLoading from .pb and .meta')
49 self.graph = tf.Graph()
50 device_name = FLAGS.gpuName \
51 if FLAGS.gpu > 0.0 else None
52 with tf.device(device_name):
53 with self.graph.as_default() as g:
54 self.build_from_pb()
55 return
56
57 if darknet is None:
58 darknet = Darknet(FLAGS)
59 self.ntrain = len(darknet.layers)
60
61 self.darknet = darknet
62 args = [darknet.meta, FLAGS]
63 self.num_layer = len(darknet.layers)
64 self.framework = create_framework(*args)
65
66 self.meta = darknet.meta
67
68 self.say('\nBuilding net ...')
69 start = time.time()

Calls

no outgoing calls