| 152 | |
| 153 | |
| 154 | def 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] |