mlx.core.tensordot#
- mlx.core.tensordot(a: array, b: array, /, axes: Union[int, List[List[int]]] = 2, *, stream: Union[None, Stream, Device] = None) array#
Compute the tensor dot product along the specified axes.
- Parameters:
a (array) – Input array
b (array) – Input array
axes (int or list(list(int)), optional) – The number of dimensions to sum over. If an integer is provided, then sum over the last
axesdimensions ofaand the firstaxesdimensions ofb. If a list of lists is provided, then sum over the corresponding dimensions ofaandb. (default: 2)
- Returns:
The tensor dot product.
- Return type:
result (array)