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

Class RngKeyAdder

imperative/python/megengine/xla/ir_utils.py:61–94  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

59# because xla pass key as a tensor, while mge pass key as a param, so we need to add a
60# rng key tensor to the graph and set it as the input of the graph and rng op
61class RngKeyAdder(Pass):
62 def __call__(self, tr) -> Any:
63 for eqn in tr.eqns:
64 if _is_rng_op(eqn.op):
65 tr.has_rng_opr = True
66 break
67
68 if not tr.has_rng_opr:
69 return tr
70
71 # it should be [2, np.uint64], however, megengine donot support np.uint64/np.int64/np.uint32
72 inp_rng_state_var = AbstractVar(tr.next_vid, [2, 2], np.dtype(np.int32))
73 tr.add_input(inp_rng_state_var)
74
75 new_eqns = []
76 for eqn in tr.eqns:
77 if not _is_rng_op(eqn.op):
78 new_eqns.append(eqn)
79 continue
80
81 oup_rng_state_var = AbstractVar(tr.next_vid, [2, 2], np.dtype(np.int32))
82 tr.add_var(oup_rng_state_var)
83
84 inputs, outputs = list(eqn.inputs), list(eqn.outputs)
85 inputs.append(inp_rng_state_var.id)
86 outputs.append(oup_rng_state_var.id)
87 new_eqn = OpInfo(eqn.op, inputs, outputs, eqn.kind)
88 new_eqns.append(new_eqn)
89 inp_rng_state_var = oup_rng_state_var
90
91 tr.eqns = new_eqns
92 tr.set_var_as_oup(inp_rng_state_var)
93
94 return tr
95
96
97# in megengine, dropout return a bit-mask while xla hard to represent, so we let xla

Callers 1

build_xlaFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected