Trim redundant comments and docstrings on the Gemma-4 MoE swap path
Shorten WHAT-style narrative on private helpers (_per_expert_forward, _quantize_one_expert_to_linear4bit) to one-line WHY statements; collapse the three-line per-module peak-VRAM note to a single line; drop the three-line opt-in description at the swap call site since the closure name already conveys the intent.
This commit is contained in:
parent
09ed2b963d
commit
6746f4d100
2 changed files with 6 additions and 18 deletions
|
|
@ -59,12 +59,8 @@ def _per_expert_forward(
|
|||
top_k_index: torch.Tensor,
|
||||
top_k_weights: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Replacement Gemma4TextExperts.forward using per-expert Linear4bit.
|
||||
|
||||
Mirrors the reference forward in transformers.models.gemma4.modeling_gemma4
|
||||
but dispatches through swapped nn.ModuleList[Linear4bit] modules instead
|
||||
of nn.functional.linear on the fused 3D Parameters.
|
||||
"""
|
||||
"""Replacement Gemma4TextExperts.forward dispatching through swapped
|
||||
nn.ModuleList[Linear4bit] instead of fused 3D Parameters."""
|
||||
final_hidden_states = torch.zeros_like(hidden_states)
|
||||
with torch.no_grad():
|
||||
expert_mask = torch.nn.functional.one_hot(
|
||||
|
|
@ -101,10 +97,8 @@ def _quantize_one_expert_to_linear4bit(
|
|||
compute_dtype: torch.dtype,
|
||||
quant_type: str = "nf4",
|
||||
):
|
||||
"""Build a bnb.nn.Linear4bit from a single (out, in) BF16 weight slice.
|
||||
|
||||
Params4bit triggers on-the-fly quantization on .to(device).
|
||||
"""
|
||||
"""Build a bnb.nn.Linear4bit from a single (out, in) weight slice.
|
||||
Params4bit triggers on-the-fly quantization on .to(device)."""
|
||||
import bitsandbytes as bnb
|
||||
|
||||
out_features, in_features = weight_2d.shape
|
||||
|
|
@ -193,17 +187,14 @@ def swap_gemma4_experts_to_per_expert_linear4bit(
|
|||
gate_up_list.append(gu.to(device))
|
||||
down_list.append(dp.to(device))
|
||||
|
||||
# The fused BF16 gate_up_proj / down_proj stay live through the loop
|
||||
# above; per-module peak is fused BF16 + accumulated per-expert nf4.
|
||||
# They are released here, before attaching the ModuleLists.
|
||||
# Per-module peak = fused BF16 + accumulated per-expert nf4; released here.
|
||||
del module.gate_up_proj
|
||||
del module.down_proj
|
||||
|
||||
module.gate_up_proj_4bit = gate_up_list
|
||||
module.down_proj_4bit = down_list
|
||||
|
||||
# Per-instance forward bind so other Gemma4TextExperts instances
|
||||
# (e.g. in a sibling model) keep the class-level method.
|
||||
# Per-instance bind so sibling Gemma4TextExperts keep the class method.
|
||||
module.forward = MethodType(_per_expert_forward, module)
|
||||
module._unsloth_gemma4_moe_4bit_swapped = True
|
||||
|
||||
|
|
|
|||
|
|
@ -1141,9 +1141,6 @@ class FastBaseModel:
|
|||
# attn_implementation = attn_implementation,
|
||||
**kwargs,
|
||||
)
|
||||
# Opt-in per-expert Linear4bit swap for Gemma-4 MoE checkpoints
|
||||
# whose fused 3D expert weights bnb cannot quantize (#5344).
|
||||
# Off by default; users enable via UNSLOTH_GEMMA4_MOE_4BIT=1.
|
||||
_maybe_swap_gemma4_moe_4bit(model)
|
||||
|
||||
# Guardrail: see _warn_if_quantization_silently_dropped + #5344.
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue