Computer >> 컴퓨터 >  >> 프로그래밍 >> Python

Python NumPy tensordot()로 텐서 내적(Tensor Dot Product) 계산하기

두 개의 텐서 ab, 그리고 두 개의 배열 객체(a_axes, b_axes)를 담고 있는 array_like 객체가 주어졌을 때, a_axes와 b_axes에 지정된 축(axis)을 따라 a와 b의 요소(성분)들의 곱을 합산합니다. 세 번째 인자는 단일한 음수가 아닌 정수 스칼라 N일 수도 있는데, 이 경우에는 a의 마지막 N개 차원과 b의 첫 번째 N개 차원이 합산됩니다.

Python에서 텐서 내적(tensor dot product)을 계산하려면 numpy.tensordot() 메서드를 사용하면 됩니다. 매개변수 ab는 내적(dot)할 대상이 되는 텐서입니다. axes 매개변수가 정수 N이라면, 순서대로 a의 마지막 N개 축과 b의 첫 번째 N개 축에 대해 합산이 수행되며, 이때 서로 대응하는 축의 크기는 반드시 일치해야 합니다.

텐서 내적 계산 단계

1단계: 필요한 라이브러리 임포트

먼저 필요한 라이브러리를 임포트합니다.

import numpy as np

2단계: 3차원 배열 생성

array() 메서드(여기서는 arange와 reshape 조합)를 사용하여 두 개의 NumPy 3차원 배열을 생성합니다.

arr1 = np.arange(60.).reshape(3,4,5)
arr2 = np.arange(24.).reshape(4,3,2)

3단계: 배열 출력 및 속성 확인

생성된 배열을 화면에 출력합니다.

print("Array1...\n",arr1)
print("\nArray2...\n",arr2)

두 배열의 차원(dimension)을 확인합니다.

print("\nDimensions of Array1...\n",arr1.ndim)
print("\nDimensions of Array2...\n",arr2.ndim)

두 배열의 형태(shape)도 함께 확인합니다.

print("\nShape of Array1...\n",arr1.shape)
print("\nShape of Array2...\n",arr2.shape)

4단계: tensordot()으로 텐서 내적 계산

Python에서 텐서 내적을 계산하려면 numpy.tensordot() 메서드를 사용합니다. 매개변수 a와 b는 내적할 텐서이며, axes 인자에 [1,0]과 [0,1]을 지정하여 arr1의 1번, 0번 축과 arr2의 0번, 1번 축을 각각 대응시켜 합산합니다.

print("\nTensor dot product...\n", np.tensordot(arr1,arr2, axes=([1,0],[0,1])))

전체 예제 코드

import numpy as np

# array() 메서드를 사용하여 두 개의 NumPy 3차원 배열 생성
arr1 = np.arange(60.).reshape(3,4,5)
arr2 = np.arange(24.).reshape(4,3,2)

# 배열 출력
print("Array1...\n",arr1)
print("\nArray2...\n",arr2)

# 두 배열의 차원 확인
print("\nDimensions of Array1...\n",arr1.ndim)
print("\nDimensions of Array2...\n",arr2.ndim)

# 두 배열의 형태 확인
print("\nShape of Array1...\n",arr1.shape)
print("\nShape of Array2...\n",arr2.shape)

# numpy.tensordot() 메서드로 텐서 내적 계산
# a, b 매개변수는 내적할 텐서입니다.
print("\nTensor dot product...\n", np.tensordot(arr1,arr2, axes=([1,0],[0,1])))

실행 결과

Array1...
[[[ 0. 1. 2. 3. 4.]
[ 5. 6. 7. 8. 9.]
[10. 11. 12. 13. 14.]
[15. 16. 17. 18. 19.]]

[[20. 21. 22. 23. 24.]
[25. 26. 27. 28. 29.]
[30. 31. 32. 33. 34.]
[35. 36. 37. 38. 39.]]

[[40. 41. 42. 43. 44.]
[45. 46. 47. 48. 49.]
[50. 51. 52. 53. 54.]
[55. 56. 57. 58. 59.]]]

Array2...
[[[ 0. 1.]
[ 2. 3.]
[ 4. 5.]]

[[ 6. 7.]
[ 8. 9.]
[10. 11.]]

[[12. 13.]
[14. 15.]
[16. 17.]]

[[18. 19.]
[20. 21.]
[22. 23.]]]

Dimensions of Array1...
3

Dimensions of Array2...
3

Shape of Array1...
(3, 4, 5)

Shape of Array2...
(4, 3, 2)

Tensor dot product...
[[4400. 4730.]
[4532. 4874.]
[4664. 5018.]
[4796. 5162.]
[4928. 5306.]]