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

Method buffer_operation

python/singa/model.py:41–93  ·  view source on GitHub ↗
(func)

Source from the content-addressed store, hash-verified

39class ModelMeta(layer.LayerMeta):
40
41 def buffer_operation(func):
42
43 def remove_creator(tensors):
44 if not tensors:
45 return
46
47 if isinstance(tensors, Iterable):
48 for item in tensors:
49 if isinstance(item, Iterable):
50 remove_creator(item)
51 elif isinstance(item, tensor.Tensor):
52 item.creator = None
53 elif isinstance(tensors, tensor.Tensor):
54 tensors.creator = None
55
56 @wraps(func)
57 def wrapper(self, *args, **kwargs):
58 if self.graph_mode and self.training:
59 if len(args) == 0:
60 raise ValueError('expect at least one input tensor')
61
62 if isinstance(args[0], list):
63 assert isinstance(
64 args[0][0],
65 Tensor), ('function expects PlaceHolders or Tensors')
66 dev = args[0][0].device
67 else:
68 assert isinstance(
69 args[0],
70 Tensor), ('function expects PlaceHolders or Tensors')
71 dev = args[0].device
72
73 if not self._buffered:
74 # buffer operations
75 dev.EnableGraph(True)
76 self._results = func(self, *args, **kwargs)
77 dev.Sync()
78 dev.EnableGraph(False)
79 self._buffered = True
80
81 # deconstruct Operations before running the entire graph
82 remove_creator(self._results)
83
84 # make sure all Operations are deallocated
85 gc.collect()
86
87 # run graph
88 dev.RunGraph(self.sequential)
89 return self._results
90 else:
91 return func(self, *args, **kwargs)
92
93 return wrapper
94
95 def __new__(cls, name, bases, attr):
96 if 'train_one_batch' in attr:

Callers 1

__new__Method · 0.45

Calls

no outgoing calls

Tested by

no test coverage detected