torch.nn.functional.grouped_mm
-
torch.nn.functional.grouped_mm(mat_a, mat_b, *, offs=None, bias=None, out_dtype=None)[исходный код] -
Вычисляет сгруппированное умножение матриц, при котором формы весов общие для экспертов, но допускается переменное количество токенов для каждого эксперта, что часто встречается в слоях Mixture-of-Experts (MoE). Оба тензора
mat_aиmat_bдолжны быть двумерными или трёхмерными и уже соответствовать ограничениям на физическое размещение данных в ядрах grouped GEMM (например,mat_aс построчным размещением иmat_bс размещением по столбцам для входных данных FP8). В настоящее время ожидается, что входные данные будут значениямиtorch.bfloat16на устройствах CUDA с .- Параметры:
-
-
mat_a (Tensor) – Левый операнд. Если он двумерный, его ведущая размерность разбивается на группы согласно
offs. Если он трёхмерный, его первое измерение непосредственно задаёт группы, аoffsдолжно бытьNone. -
mat_b (Tensor) – Правый операнд. Если оба операнда двумерные (например, при обновлении градиентов весов MoE), замыкающая размерность
mat_aи ведущая размерностьmat_bразбиваются на части согласно одному и тому же тензоруoffs. В типичном прямом проходе (out = input @ weight.T)mat_bявляется трёхмерным тензором формы(num_groups, K, N). Если веса экспертов хранятся в стандартном форматеnn.Linear(num_groups, N, K), передайтеweight.transpose(-2, -1)какmat_b. -
offs (Tensor | None) – Необязательный одномерный тензор монотонно возрастающих смещений
int32, задающих границы переменной размерности любого двумерного операнда.offs[i]обозначает конец группыiиoffs[-1]должно быть строго меньше общей длины срезанной размерности этого операнда; элементы послеoffs[-1]игнорируются. - bias (Tensor | None) – Необязательный тензор, который прибавляется к сгруппированным результатам. Смещение не имеет переменной размерности и должно поддерживать широковещательное распространение до формы результата каждой группы.
-
out_dtype (dtype | None) – Необязательный тип данных, определяющий тип данных для аккумуляции и выходных данных. Передача
torch.float32приводит к аккумуляции входных данных BF16 в FP32, при этом API grouped GEMM остаётся недифференцируемым.
-
mat_a (Tensor) – Левый операнд. Если он двумерный, его ведущая размерность разбивается на группы согласно
- Возвращает:
-
Тензор, содержащий объединённые результаты GEMM для каждой группы; его форма определяется операндами и
offs. - Тип возвращаемого значения:
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.nn.functional.grouped_mm.html