MCPcopy Create free account
hub / github.com/apache/singa / wrapper

Method wrapper

examples/model_selection/Trails/singa_pkg_code/model.py:65–112  ·  view source on GitHub ↗
(self, *args, **kwargs)

Source from the content-addressed store, hash-verified

63
64 @wraps(func)
65 def wrapper(self, *args, **kwargs):
66 # print ("in model.py wrapper function")
67 # print ("in model.py wrapper function args[0] shape: ", args[0].shape)
68 # print ("in model.py self._buffered: ", self._buffered)
69 # print ("in model.py begin wrapper self._results: ", self._results)
70 if self.graph_mode and self.training:
71 if len(args) == 0:
72 raise ValueError('expect at least one input tensor')
73
74 if isinstance(args[0], list):
75 assert isinstance(
76 args[0][0],
77 Tensor), ('function expects PlaceHolders or Tensors')
78 dev = args[0][0].device
79 else:
80 assert isinstance(
81 args[0],
82 Tensor), ('function expects PlaceHolders or Tensors')
83 dev = args[0].device
84
85 if not self._buffered:
86 # buffer operations
87 dev.EnableGraph(True)
88 # print ("model.py wrap not self._buffered args[0].shape", args[0].shape)
89 self._results = func(self, *args, **kwargs)
90 # print ("model.py wrap not self._buffered func: ", func)
91 dev.Sync()
92 dev.EnableGraph(False)
93 self._buffered = True
94
95 # deconstruct Operations before running the entire graph
96 remove_creator(self._results)
97
98 # make sure all Operations are deallocated
99 gc.collect()
100
101 # run graph
102 # print ("in model.py before dev.RunGraph self._results[0] shape: ", self._results[0].shape)
103 # print ("in model.py before dev.RunGraph args[0] shape: ", args[0].shape)
104 # print ("in model.py before dev.RunGraph self._results: ", self._results)
105 dev.RunGraph(self.sequential)
106 # print ("in model.py after dev.RunGraph")
107 # print ("in model.py after dev.RunGraph self._results[0] shape: ", self._results[0].shape)
108 # print ("in model.py after dev.RunGraph self._results: ", self._results)
109 # print ("in model.py after dev.RunGraph args[0] shape: ", args[0].shape)
110 return self._results
111 else:
112 return func(self, *args, **kwargs)
113
114 print ("model.py return buffer_operation wrapper: ", wrapper)
115 return wrapper

Callers

nothing calls this directly

Calls 3

EnableGraphMethod · 0.80
SyncMethod · 0.45
RunGraphMethod · 0.45

Tested by

no test coverage detected