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

Class DropoutMaskCanonicalizer

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

Source from the content-addressed store, hash-verified

97# in megengine, dropout return a bit-mask while xla hard to represent, so we let xla
98# return a uint8 mask, which means the mask is 8 times larger than mge
99class DropoutMaskCanonicalizer(Pass):
100 def __call__(self, tr) -> Any:
101 for eqn in tr.eqns:
102 if not isinstance(eqn.op, mops.Dropout):
103 continue
104
105 inputs, outputs = list(eqn.inputs), list(eqn.outputs)
106 mask_var = tr.vars[outputs[1]]
107 inp_shape = tr.vars[inputs[0]].shape
108 new_mask_var = AbstractVar(
109 mask_var.id, (int(np.prod(inp_shape)),), mask_var.dtype
110 )
111 tr.vars[mask_var.id] = new_mask_var
112
113 return tr
114
115
116class TraceResult:

Callers 1

build_xlaFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected