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

Method buffer_operation

examples/cnn_ms/pkg_model_code/model.py:41–96  ·  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 if isinstance(tensors, str):
49 return
50 else:
51 for item in tensors:
52 if isinstance(item, Iterable):
53 remove_creator(item)
54 elif isinstance(item, tensor.Tensor):
55 item.creator = None
56 elif isinstance(tensors, tensor.Tensor):
57 tensors.creator = None
58
59 @wraps(func)
60 def wrapper(self, *args, **kwargs):
61 if self.graph_mode and self.training:
62 if len(args) == 0:
63 raise ValueError('expect at least one input tensor')
64
65 if isinstance(args[0], list):
66 assert isinstance(
67 args[0][0],
68 Tensor), ('function expects PlaceHolders or Tensors')
69 dev = args[0][0].device
70 else:
71 assert isinstance(
72 args[0],
73 Tensor), ('function expects PlaceHolders or Tensors')
74 dev = args[0].device
75
76 if not self._buffered:
77 # buffer operations
78 dev.EnableGraph(True)
79 self._results = func(self, *args, **kwargs)
80 dev.Sync()
81 dev.EnableGraph(False)
82 self._buffered = True
83
84 # deconstruct Operations before running the entire graph
85 remove_creator(self._results)
86
87 # make sure all Operations are deallocated
88 gc.collect()
89
90 # run graph
91 dev.RunGraph(self.sequential)
92 return self._results
93 else:
94 return func(self, *args, **kwargs)
95
96 return wrapper
97
98 def __new__(cls, name, bases, attr):

Callers 1

__new__Method · 0.45

Calls

no outgoing calls

Tested by

no test coverage detected