| 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 |
| 61 | class 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 |