MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / subgraph

Function subgraph

imperative/python/megengine/core/tensor/utils.py:154–268  ·  view source on GitHub ↗
(
    name, dtype, device, nr_inputs, gopt_level=None, jit_fusion=False, custom_grad=False
)

Source from the content-addressed store, hash-verified

152
153
154def subgraph(
155 name, dtype, device, nr_inputs, gopt_level=None, jit_fusion=False, custom_grad=False
156):
157 if not device.physical_name.startswith("gpu"):
158 jit_fusion = False
159
160 if jit_fusion and not jit_supported:
161 jit_fusion = False # jit unusable, fallback to graph compile
162 gopt_level = 2
163
164 def as_op(op, nargs):
165 if isinstance(op, str):
166 assert (op, nargs) in _opr_map, "unknown operator"
167 op = _opr_map[(op, nargs)]
168 return op
169
170 def decorator(func):
171 builder = _SubgraphBuilder(name)
172
173 def apply_expr(op, *args, nr_out=None):
174 op = as_op(op, len(args))
175 results = builder.apply(op, args, 1 if nr_out is None else nr_out)
176 if nr_out is None:
177 assert len(results) == 1
178 return results[0]
179 else:
180 assert len(results) == nr_out
181 return results
182
183 def apply_const(value, dtype=dtype, device=device):
184 return builder.apply_const(value, dtype, device)
185
186 def build(builder, outputs, outputs_has_grad):
187 builder = type(builder)(builder)
188 builder.outputs(outputs)
189 builder.outputs_has_grad(outputs_has_grad)
190 if jit_fusion:
191 assert gopt_level is None
192 op = lambda: builder.jit_fuse()
193 elif gopt_level is None:
194 op = lambda: builder.get()
195 else:
196 op = lambda: builder.compile(gopt_level)
197 return op
198
199 inputs = [builder.input() for _ in range(nr_inputs)]
200 if not custom_grad:
201 outputs, outputs_has_grad = func(inputs, apply_expr, apply_const)
202 return build(builder, outputs, outputs_has_grad)
203 else:
204 gen = func(inputs, apply_expr, apply_const)
205 outputs = gen.send(None)
206 nr_outputs = len(outputs)
207 forward_fn = build(builder, outputs, [False] * nr_outputs)
208 output_grads = [builder.input() for _ in range(nr_outputs)]
209 input_grads = gen.send(output_grads)
210 assert len(input_grads) == nr_inputs
211 input_grads_mask = [input_grad is not None for input_grad in input_grads]

Callers 1

decoratorFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected