andreasmartin's picture
Update standalone Apertus embedding model: moe_head
8a0a543 verified
Raw
History Blame Contribute Delete
4.07 kB
from typing import Any
import torch
from torch import nn
from sentence_transformers.base.modules import Module
class LanguageMoE(Module):
config_keys = [
"hidden_size",
"embedding_dim",
"expert_names",
"top_k",
"temperature",
]
def __init__(
self,
hidden_size: int,
embedding_dim: int,
expert_names: list[str],
top_k: int = 2,
temperature: float = 1.0,
**kwargs,
) -> None:
super().__init__()
self.hidden_size = int(hidden_size)
self.embedding_dim = int(embedding_dim)
self.expert_names = list(expert_names)
self.top_k = int(top_k)
self.temperature = float(temperature)
assert 1 <= self.top_k <= len(self.expert_names)
self.router = nn.Linear(
self.hidden_size,
len(self.expert_names),
bias=True,
)
self.experts = nn.ModuleList([
nn.Linear(
self.hidden_size,
self.embedding_dim,
bias=False,
)
for _ in self.expert_names
])
self.last_router_probs = None
self.last_top_indices = None
def route(self, x: torch.Tensor):
logits = self.router(x) / self.temperature
probs = torch.softmax(logits, dim=-1)
top_probs, top_idx = torch.topk(
probs,
k=self.top_k,
dim=-1,
)
gates = top_probs / top_probs.sum(
dim=-1,
keepdim=True,
).clamp_min(1e-12)
return logits, probs, top_idx, gates
def project_sparse(
self,
x: torch.Tensor,
top_idx: torch.Tensor,
gates: torch.Tensor,
) -> torch.Tensor:
output = x.new_zeros(
(x.shape[0], self.embedding_dim)
)
for expert_id, expert in enumerate(self.experts):
selected = top_idx.eq(expert_id)
rows, slots = selected.nonzero(as_tuple=True)
if rows.numel() == 0:
continue
expert_output = expert(x[rows])
expert_gate = gates[rows, slots].unsqueeze(-1)
output.index_add_(
0,
rows,
expert_output * expert_gate,
)
return output
def forward(
self,
features: dict[str, torch.Tensor | Any],
**kwargs,
) -> dict[str, torch.Tensor | Any]:
x = features["sentence_embedding"]
_, probs, top_idx, gates = self.route(x)
features["sentence_embedding"] = self.project_sparse(
x,
top_idx,
gates,
)
self.last_router_probs = probs.detach()
self.last_top_indices = top_idx.detach()
return features
def get_embedding_dimension(self) -> int:
return self.embedding_dim
def save(
self,
output_path: str,
*args,
safe_serialization: bool = True,
**kwargs,
) -> None:
self.save_config(output_path)
self.save_torch_weights(
output_path,
safe_serialization=safe_serialization,
)
@classmethod
def load(
cls,
model_name_or_path: str,
subfolder: str = "",
token: bool | str | None = None,
cache_folder: str | None = None,
revision: str | None = None,
local_files_only: bool = False,
**kwargs,
):
config = cls.load_config(
model_name_or_path,
subfolder=subfolder,
token=token,
cache_folder=cache_folder,
revision=revision,
local_files_only=local_files_only,
)
model = cls(**config)
return cls.load_torch_weights(
model_name_or_path,
subfolder=subfolder,
token=token,
cache_folder=cache_folder,
revision=revision,
local_files_only=local_files_only,
model=model,
)