MCPcopy Create free account
hub / github.com/OpenBitSys/BitDistiller / make_fused_mlp

Function make_fused_mlp

inference/modules/fused_mlp.py:117–132  ·  view source on GitHub ↗
(m, parent_name="")

Source from the content-addressed store, hash-verified

115
116
117def make_fused_mlp(m, parent_name=""):
118 if not hasattr(make_fused_mlp, "called"):
119 # print("[Warning] Calling a fake MLP fusion. But still faster than Huggingface Implimentation.")
120 make_fused_mlp.called = True
121 """
122 Replace all LlamaMLP modules with QuantLlamaMLP modules, which fuses many of the operations.
123 """
124 if m.__class__.__name__ in ["LlamaMLP"]:
125 return QuantLlamaMLP(m.gate_proj, m.down_proj, m.up_proj)
126
127 for name, child in m.named_children():
128 child = make_fused_mlp(child, parent_name=f"{parent_name}.{name}")
129
130 if isinstance(child, QuantLlamaMLP):
131 setattr(m, name, child)
132 return m

Callers 1

demo.pyFile · 0.90

Calls 1

QuantLlamaMLPClass · 0.85

Tested by

no test coverage detected