MCPcopy Create free account
hub / github.com/apache/tvm / get_expected_3

Function get_expected_3

tests/python/relax/test_transform_gradient_te_register.py:286–355  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

284
285
286def 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))

Callers 1

test_tir_varFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…