| 10 | import os |
| 11 | |
| 12 | class 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() |
no outgoing calls