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

Python NumPy einsum()으로 아인슈타인 합 규약 기반 텐서 수축하기


Python에서 아인슈타인 합 규약(Einstein Summation Convention)을 활용한 텐서 수축(Tensor Contraction)을 수행하려면 numpy.einsum() 메서드를 사용하면 됩니다.

첫 번째 매개변수는 subscript(첨자)로, 합산에 포함될 첨자 레이블을 쉼표로 구분한 목록 형태로 지정합니다. 두 번째 매개변수는 operands(피연산자)로, 연산에 사용될 배열들을 의미합니다.

einsum() 메서드는 피연산자에 대해 아인슈타인 합 규약을 계산합니다. 이 규약을 활용하면 다차원 선형대수학의 일반적인 배열 연산 대부분을 아주 간결한 방식으로 표현할 수 있습니다. 암시적 모드(implicit mode)에서는 einsum이 이러한 값들을 자동으로 계산해 줍니다.

반면 명시적 모드(explicit mode)에서는 특정 첨자 레이블에 대한 합산을 비활성화하거나 강제할 수 있기 때문에, 전통적인 아인슈타인 합 연산으로 분류되지 않는 다양한 배열 연산도 훨씬 유연하게 처리할 수 있습니다.

진행 단계

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

import numpy as np

np.arange()로 생성한 1차원 배열을 reshape() 메서드로 재구성하여 두 개의 3차원 NumPy 배열을 만듭니다.

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

생성된 배열을 출력해 내용을 확인합니다.

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

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

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)

아인슈타인 합 규약으로 텐서 수축을 수행하려면 Python에서 numpy.einsum() 메서드를 다음과 같이 호출합니다.

print("\nResult (Tensor contraction)...\n",np.einsum('ijk,jil->kl', arr1, arr2))

전체 예제 코드

import numpy as np

# arange()와 reshape()를 사용하여 두 개의 3차원 NumPy 배열 생성
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.einsum() 메서드로 아인슈타인 합 규약 기반 텐서 수축 수행
print("\nResult (Tensor contraction)...\n",np.einsum('ijk,jil->kl', arr1, arr2))

실행 결과

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)

Result (Tensor contraction)...
[[4400. 4730.]
 [4532. 4874.]
 [4664. 5018.]
 [4796. 5162.]
 [4928. 5306.]]