| 115 | |
| 116 | |
| 117 | def 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 |