()
| 284 | |
| 285 | |
| 286 | def get_expected_3(): |
| 287 | # fmt: off |
| 288 | @I.ir_module(s_tir=True) |
| 289 | class Expected: |
| 290 | @T.prim_func(private=True, s_tir=True) |
| 291 | def f_mul(var_A: T.handle, var_B: T.handle, var_f_mul: T.handle): |
| 292 | T.func_attr({"tirx.noalias": True}) |
| 293 | n = T.int64() |
| 294 | A = T.match_buffer(var_A, (n, n)) |
| 295 | B = T.match_buffer(var_B, (n, n)) |
| 296 | f_mul_1 = T.match_buffer(var_f_mul, (n, n)) |
| 297 | # with T.sblock("root"): |
| 298 | for i0, i1 in T.grid(n, n): |
| 299 | with T.sblock("f_mul"): |
| 300 | v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) |
| 301 | T.reads(A[v_i0, v_i1], B[v_i0, v_i1]) |
| 302 | T.writes(f_mul_1[v_i0, v_i1]) |
| 303 | f_mul_1[v_i0, v_i1] = A[v_i0, v_i1] * B[v_i0, v_i1] |
| 304 | |
| 305 | @T.prim_func(private=True, s_tir=True) |
| 306 | def f_mul_grad(var_A: T.handle, var_B: T.handle, var_C: T.handle, var_f_mul_grad_1: T.handle, var_f_mul_grad_2: T.handle): |
| 307 | T.func_attr({"tirx.noalias": True}) |
| 308 | n = T.int64() |
| 309 | A = T.match_buffer(var_A, (n, n)) |
| 310 | B = T.match_buffer(var_B, (n, n)) |
| 311 | C = T.match_buffer(var_C, (n, n)) |
| 312 | f_mul_grad_1 = T.match_buffer(var_f_mul_grad_1, (n, n)) |
| 313 | f_mul_grad_2 = T.match_buffer(var_f_mul_grad_2, (n, n)) |
| 314 | # with T.sblock("root"): |
| 315 | for i0, i1 in T.grid(n, n): |
| 316 | with T.sblock("f_mul_grad_1"): |
| 317 | v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) |
| 318 | T.reads(C[v_i0, v_i1], A[v_i0, v_i1]) |
| 319 | T.writes(f_mul_grad_1[v_i0, v_i1]) |
| 320 | f_mul_grad_1[v_i0, v_i1] = C[v_i0, v_i1] * A[v_i0, v_i1] |
| 321 | for i0, i1 in T.grid(n, n): |
| 322 | with T.sblock("f_mul_grad_2"): |
| 323 | v_i0, v_i1 = T.axis.remap("SS", [i0, i1]) |
| 324 | T.reads(B[v_i0, v_i1], A[v_i0, v_i1]) |
| 325 | T.writes(f_mul_grad_2[v_i0, v_i1]) |
| 326 | f_mul_grad_2[v_i0, v_i1] = B[v_i0, v_i1] * A[v_i0, v_i1] |
| 327 | |
| 328 | @R.function |
| 329 | def main_adjoint(a: R.Tensor(("n", "n"), dtype="float32"), b: R.Tensor(("n", "n"), dtype="float32")) -> R.Tuple(R.Tensor((), dtype="float32"), R.Tuple(R.Tensor(("n", "n"), dtype="float32"), R.Tensor(("n", "n"), dtype="float32"))): |
| 330 | n = T.int64() |
| 331 | cls = Expected |
| 332 | with R.dataflow(): |
| 333 | lv = R.call_tir(cls.f_mul, (a, b), out_sinfo=R.Tensor((n, n), dtype="float32")) |
| 334 | gv: R.Tensor((), dtype="float32") = R.sum(lv, axis=None, keepdims=False) |
| 335 | gv_adjoint: R.Tensor((), dtype="float32") = R.ones(R.shape([]), dtype="float32") |
| 336 | lv_adjoint: R.Tensor((n, n), dtype="float32") = R.broadcast_to(gv_adjoint, R.shape([n, n])) |
| 337 | lv_1 = R.call_tir(cls.f_mul_grad, (lv_adjoint, a, b), out_sinfo=[R.Tensor((n, n), dtype="float32"), R.Tensor((n, n), dtype="float32")]) |
| 338 | a_adjoint: R.Tensor((n, n), dtype="float32") = lv_1[0] |
| 339 | b_adjoint: R.Tensor((n, n), dtype="float32") = lv_1[1] |
| 340 | a_adjoint_out: R.Tensor((n, n), dtype="float32") = a_adjoint |
| 341 | b_adjoint_out: R.Tensor((n, n), dtype="float32") = b_adjoint |
| 342 | R.output(gv, a_adjoint_out, b_adjoint_out) |
| 343 | return (gv, (a_adjoint_out, b_adjoint_out)) |
no outgoing calls
no test coverage detected
searching dependent graphs…