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
153 lines
4.1 KiB
Python
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
|