Files
unslothai--unsloth/unsloth/kernels/moe/grouped_gemm/reference/moe_ops.py
T
wehub-resource-sync e93507a09c
Lockfile supply-chain audit / lockfile supply-chain audit (push) Has been cancelled
Windows Studio GGUF CI / GPU prebuilt resolves without Visual Studio (push) Has been cancelled
Windows Studio GGUF CI / setup.ps1 unit tests (VS 2026 / CMake guard) (push) Has been cancelled
Windows Studio GGUF CI / real-VS detection (VS 2022) (push) Has been cancelled
Windows Studio GGUF CI / real-VS detection (VS 2026) (push) Has been cancelled
Windows Studio GGUF CI / VC++ runtime detect + install round-trip (windows-2025-vs2026) (push) Has been cancelled
Windows Studio GGUF CI / VC++ runtime detect + install round-trip (windows-latest) (push) Has been cancelled
Windows Studio Update CI / Studio Updating Tests (push) Has been cancelled
Wheel CI / Wheel build + content sanity + import smoke (push) Has been cancelled
Lint CI / Source lint (Python + shell + YAML + JSON + safety nets) (push) Has been cancelled
MLX CI on Mac M1 / dispatch (push) Has been cancelled
Security audit / advisory audit (pip + npm + cargo) (push) Has been cancelled
Security audit / pip scan-packages :: extras (push) Has been cancelled
Security audit / pip scan-packages :: studio (push) Has been cancelled
Security audit / pip scan-packages :: hf-stack (push) Has been cancelled
Security audit / npm scan-packages (Studio frontend tarballs) (push) Has been cancelled
Security audit / workflow-trigger lint (pull_request_target / cache-poisoning) (push) Has been cancelled
Security audit / pytest tests/security (push) Has been cancelled
Security audit / npm provenance + new install-script diff (push) Has been cancelled
Studio API CI / Studio API & Auth Tests (push) Has been cancelled
Backend CI / (Python 3.10) (push) Has been cancelled
Backend CI / (Python 3.11) (push) Has been cancelled
Backend CI / (Python 3.12) (push) Has been cancelled
Backend CI / (Python 3.13) (push) Has been cancelled
Backend CI / Repo tests (CPU) (push) Has been cancelled
Frontend CI / Frontend build + bundle sanity (push) Has been cancelled
Studio GGUF CI / OpenAI, Anthropic API tests (push) Has been cancelled
Studio GGUF CI / Tool calling Tests (push) Has been cancelled
Studio GGUF CI / JSON, images (push) Has been cancelled
Mac Studio GGUF CI / OpenAI, Anthropic API tests (push) Has been cancelled
Mac Studio GGUF CI / Tool calling Tests (push) Has been cancelled
Mac Studio GGUF CI / JSON, images (push) Has been cancelled
Mac Studio Install Matrix CI / Install + load (macos-14) (push) Has been cancelled
Mac Studio Install Matrix CI / Install + load (macos-15) (push) Has been cancelled
Mac Studio Install Matrix CI / Install + load (macos-26) (push) Has been cancelled
Mac Studio Install Matrix CI / Install + load (macos-15-intel) (push) Has been cancelled
Mac Studio API CI / Studio API & Auth Tests (push) Has been cancelled
Mac Studio Install Matrix CI / Install + load (macos-26-intel) (push) Has been cancelled
Mac Studio UI CI / Chat UI Tests (push) Has been cancelled
Studio Tauri CI / Tauri Linux debug build (no codesign) (push) Has been cancelled
Mac Studio Update CI / Studio Updating Tests (push) Has been cancelled
Studio UI CI / Chat UI Tests (push) Has been cancelled
Windows Studio API CI / Studio API & Auth Tests (push) Has been cancelled
Windows Studio UI CI / Chat UI Tests (push) Has been cancelled
Studio Update CI / Studio Updating Tests (push) Has been cancelled
Core / Core (HF=default + TRL=default) (push) Has been cancelled
Core / Core (HF=4.57.6 + TRL<1) (push) Has been cancelled
Core / Core (HF=latest + TRL=latest) (push) Has been cancelled
Core / llama.cpp build + smoke (push) Has been cancelled
Windows Studio GGUF CI / OpenAI, Anthropic API tests (push) Has been cancelled
Windows Studio GGUF CI / Tool calling Tests (push) Has been cancelled
Windows Studio GGUF CI / JSON, images (push) Has been cancelled
Windows Studio GGUF CI / Studio install + inference without Visual Studio (push) Has been cancelled
Studio export capability / capability (macos-latest) (push) Has been cancelled
Studio export capability / capability (ubuntu-latest) (push) Has been cancelled
Studio export capability / capability (windows-latest) (push) Has been cancelled
Cross-platform parity / parity (macos-latest) (push) Has been cancelled
Cross-platform parity / parity (windows-latest) (push) Has been cancelled
Scorecard supply-chain security / Scorecard analysis (push) Has been cancelled
Studio load-orchestrator CI / test (push) Has been cancelled
chore: import upstream snapshot with attribution
2026-07-13 12:59:56 +08:00

153 lines
4.1 KiB
Python

# SPDX-License-Identifier: GNU Affero General Public License v3.0
# Copyright 2023-present the Unsloth team. All rights reserved.
import torch
import torch.nn.functional as F
def permute(X: torch.Tensor, gather_indices: torch.Tensor, topk: int):
"""
Scatters X to a new tensor with shape [total_tokens, hidden_dim] where total_tokens is num_tokens * topk,
permuting the tokens according to sorted_token_idx.
Helper for grouped gemm where hidden states need be ordered by expert.
X: [num_tokens, hidden_dim]
sorted_token_idx: [num_tokens * topk]
topk: int
Returns:
[total_tokens, hidden_dim]
"""
assert gather_indices.ndim == 1
X = X.view(-1, X.shape[-1])
# Shortcut for topk == 1
if topk == 1:
return X[gather_indices]
return X[gather_indices // topk]
def unpermute(X: torch.Tensor, gather_indices: torch.Tensor):
X = X.view(-1, X.shape[-1]) if X.ndim > 2 else X
unpermuted = torch.empty_like(X)
unpermuted.index_copy_(0, gather_indices, X)
return unpermuted.view_as(X)
def calculate_topk(
gating_output: torch.Tensor,
top_k: int,
use_sigmoid: bool,
renormalize: bool,
pre_act: bool = True,
post_act: bool = False,
):
"""
If post_act is True, then activation function is run AFTER topk
If post_act is False, then activation function is run BEFORE topk
This is to align with triton_bench implementation (post_act) whereas most models use pre_act (e.g. llama4, deepseek)
"""
assert pre_act ^ post_act, "only one of pre_act or post_act can be True"
def _activation(gating_output: torch.Tensor):
if use_sigmoid:
scores = torch.sigmoid(gating_output.to(torch.float32)).to(gating_output.dtype)
else:
scores = F.softmax(gating_output.to(torch.float32), dim = 1).to(gating_output.dtype)
return scores
if pre_act:
scores = _activation(gating_output)
else:
scores = gating_output
topk_weights, topk_ids = torch.topk(scores, k = top_k, dim = 1)
if post_act:
topk_weights = _activation(topk_weights)
if renormalize:
topk_weights /= torch.sum(topk_weights, dim = -1, keepdim = True).to(gating_output.dtype)
return topk_weights, topk_ids
@torch.no_grad()
def get_routing_indices(
selected_experts,
num_experts,
return_scatter_indices: bool = False,
):
"""
Returns:
token_counts_by_expert: [num_experts]
gather_indices: [num_tokens]
scatter_indices [Optional] (torch.Tensor):
Indices for unpermuting gathered inputs back to token order, shape ``(bs * seqlen * top_k,)``.
"""
# group tokens together by expert indices from 0 to num_experts and pass that to experts forward
token_counts_by_expert = torch.histc(
selected_experts.view(-1),
bins = num_experts,
min = 0,
max = num_experts,
)
# token_indices_experts_sorted shape (bs*slen*top_k,)
gather_indices = torch.argsort(selected_experts.view(-1), stable = True)
if return_scatter_indices:
scatter_indices = gather_indices.argsort()
return token_counts_by_expert, gather_indices, scatter_indices
else:
return token_counts_by_expert, gather_indices
def torch_grouped_gemm(
X,
W,
m_sizes,
transpose = True,
):
"""
X: [M, K] if forward, else [M, N]
W: [E, N, K]
m_sizes: [E]
Returns:
Y: [M, N] if forward, else [M, K]
"""
X = X.view(-1, X.shape[-1])
M, K = X.shape
assert m_sizes.ndim == 1
E = m_sizes.shape[0]
assert W.ndim == 3
assert W.shape[0] == E
N = W.shape[1]
result = torch.zeros((M, N), dtype = X.dtype, device = X.device)
m_start = 0
for g in range(E):
m_size = m_sizes[g]
if m_size > 0:
m_end = m_start + m_size
# Extract group input
# m_size x K
X_g = X[m_start:m_end]
# N x K
W_g = W[g]
# Y_g = X_g @ W_g.T -> [m_size, N]
W_g = W_g.T if transpose else W_g
Y_g = X_g @ W_g
result[m_start:m_end] = Y_g
m_start = m_end
return result