conversion: refactor moe exp tensor handling

This commit is contained in:
Xuan Son Nguyen
2026-08-17 01:54:08 +02:00
parent 4df29be4f4
commit 051139f511
25 changed files with 191 additions and 1081 deletions
+2 -43
View File
@@ -10,7 +10,7 @@ import torch
if TYPE_CHECKING:
from torch import Tensor
from .base import MmprojModel, ModelBase, SentencePieceTokenTypes, TextModel, gguf, logger
from .base import MOE_BLOCK_SPARSE, MmprojModel, ModelBase, SentencePieceTokenTypes, TextModel, gguf, logger
@ModelBase.register("PhiForCausalLM")
@@ -339,50 +339,9 @@ class Phi4VisionMmprojModel(MmprojModel):
class PhiMoeModel(Phi3MiniModel):
model_arch = gguf.MODEL_ARCH.PHIMOE
_experts: list[dict[str, Tensor]] | None = None
moe_experts = [MOE_BLOCK_SPARSE]
def set_gguf_parameters(self):
super().set_gguf_parameters()
self.gguf_writer.add_expert_used_count(self.find_hparam(["num_experts_per_tok", "num_experts_per_token"]))
self.gguf_writer.add_expert_count(self.find_hparam(["num_local_experts", "num_experts"]))
def modify_tensors(self, data_torch: Tensor, name: str, bid: int | None) -> Iterable[tuple[str, Tensor]]:
# process the experts separately
if name.find("block_sparse_moe.experts") != -1:
n_experts = self.find_hparam(["num_local_experts", "num_experts"])
assert bid is not None
if self._experts is None:
self._experts = [{} for _ in range(self.block_count)]
self._experts[bid][name] = data_torch
if len(self._experts[bid]) >= n_experts * 3:
# merge the experts into a single 3d tensor
for w_name in ["w1", "w2", "w3"]:
datas: list[Tensor] = []
for xid in range(n_experts):
ename = f"model.layers.{bid}.block_sparse_moe.experts.{xid}.{w_name}.weight"
datas.append(self._experts[bid][ename])
del self._experts[bid][ename]
data_torch = torch.stack(datas, dim=0)
merged_name = f"model.layers.{bid}.block_sparse_moe.experts.{w_name}.weight"
yield from super().modify_tensors(data_torch, merged_name, bid)
return
else:
return
yield from super().modify_tensors(data_torch, name, bid)
def prepare_tensors(self):
super().prepare_tensors()
if self._experts is not None:
# flatten `list[dict[str, Tensor]]` into `list[str]`
experts = [k for d in self._experts for k in d.keys()]
if len(experts) > 0:
raise ValueError(f"Unprocessed experts: {experts}")