Матрично-векторные произведения через оператор @
Содержание курса
Диагностика и исправление ошибки несовпадения размеров
Рано или поздно вы напишете что-то вроде A @ B и получите стену красного текста. Разберём, как её читать.
Типичное сообщение об ошибке выглядит так:
ValueError: matmul: Input operand 1 has a mismatch in its core dimension 0
with gufunc signature (n?,k),(k,m?)->(n?,m?)
Это не очень дружелюбно, но сообщение говорит ровно одно: внутренние размеры операндов не совпали. «Core dimension 0» правого операнда — это его первое измерение, которое должно равняться последнему измерению левого.
Воспроизведём ошибку:
import numpy as np
A = np.array([[1, 2, 3],
[4, 5, 6]]) # shape (2, 3)
B = np.array([[1, 2, 3],
[4, 5, 6]]) # shape (2, 3)
result = A @ B # ValueError!
Обе матрицы имеют одинаковый shape (2, 3) — казалось бы, должно работать. Но внутренние размеры здесь 3 (последний у A) и 2 (первый у B). Они не совпадают, отсюда и ошибка.
Систематическая диагностика занимает три шага:
- Напечатать shape обоих операндов прямо перед вызовом
@:
print(A.shape) # (2, 3)
print(B.shape) # (2, 3)
- Выписать пару «последний размер левого» и «первый размер правого» и сравнить:
(2, 3) @ (2, 3)
^ ^
3 ≠ 2 → ValueError
- Решить, что нужно исправить: транспонировать один из операндов, изменить форму через
reshape, или пересмотреть логику.
Для случая выше, если нужно перемножить A на A по правилам матричной алгебры, правильный вариант — A @ A.T или A.T @ A:
result1 = A @ A.T # (2, 3) @ (3, 2) → (2, 2)
result2 = A.T @ A # (3, 2) @ (2, 3) → (3, 3)
print(result1.shape) # (2, 2)
print(result2.shape) # (3, 3)
Оба варианта корректны с точки зрения @, но дают разные результаты — какой нужен, определяет задача.
Ещё один частый сценарий: перепутать порядок операндов. Матричное умножение не коммутативно, и A @ B в общем случае не равно B @ A. Более того, если A @ B работает, B @ A может выбросить ValueError из-за несовместимых внутренних размеров:
C = np.ones((3, 4)) # shape (3, 4)
D = np.ones((4, 2)) # shape (4, 2)
print((C @ D).shape) # (3, 2) — OK
print((D @ C).shape) # ValueError: 2 ≠ 3
Если ValueError возник неожиданно, первое действие — print(A.shape, B.shape). Почти всегда проблема сразу становится видна без дополнительной отладки.
