Spec-Zone.ru › PyTorch 2.14

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 с SM≥80SM \ge 80.

Параметры:
  • 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 остаётся недифференцируемым.
Возвращает:

Тензор, содержащий объединённые результаты GEMM для каждой группы; его форма определяется операндами и offs.

Тип возвращаемого значения:

Tensor

© 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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API