PyTorch는 텐서에서 특정 순위의 요소나 가장 큰 값을 손쉽게 추출할 수 있는 유용한 함수들을 제공합니다.
torch.kthvalue() 함수는 텐서를 오름차순으로 정렬했을 때의 k번째 요소 값을 반환하며, 해당 요소가 원본 텐서에서 위치한 인덱스도 함께 제공합니다.
torch.topk() 함수는 텐서에서 상위 "k"개, 즉 가장 큰 "k"개의 요소를 찾는 데 사용됩니다. 이 두 함수는 데이터 분석이나 모델 출력 결과 해석 시 자주 활용되는 핵심 기능입니다.
구현 단계
필요한 라이브러리를 임포트합니다. 아래의 모든 Python 예제에서 필요한 라이브러리는 torch입니다. 미리 설치되어 있는지 확인하세요.
PyTorch 텐서를 생성하고 이를 출력합니다.
torch.kthvalue(input, k)를 계산합니다. 이 함수는 두 개의 텐서를 반환하며, 반환된 텐서를 "value"와 "index"라는 두 변수에 각각 할당합니다. 여기서 input은 텐서이고, k는 정수입니다.
torch.topk(input, k)를 계산합니다. 이 함수 역시 두 개의 텐서를 반환하는데, 첫 번째 텐서에는 상위 "k"개 요소의 값이, 두 번째 텐서에는 해당 요소들이 원본 텐서에서 가지는 인덱스가 담겨 있습니다. 반환된 텐서를 각각 "values"와 "indices" 변수에 할당합니다.
텐서의 k번째 요소의 값과 인덱스, 그리고 상위 "k"개 요소들의 값과 인덱스를 출력합니다.
예제 1
다음 Python 프로그램은 텐서에서 k번째 요소를 찾는 방법을 보여줍니다.
# 텐서에서 k번째 요소를 찾는 Python 프로그램
# 필요한 라이브러리 임포트
import torch
# 1D 텐서 생성
T = torch.Tensor([2.334,4.433,-4.33,-0.433,5, 4.443])
print("Original Tensor:\n", T)
# 정렬된 텐서에서 3번째 요소 찾기.
# 먼저 텐서를 오름차순으로 정렬한 후,
# 정렬된 텐서에서 k번째 요소의 값과
# 원본 텐서에서의 인덱스를 반환합니다.
value, index = torch.kthvalue(T, 3)
# 3번째 요소의 값과 인덱스 출력
print("3rd element value:", value)
print("3rd element index:", index)출력 결과
Original Tensor:
tensor([ 2.3340, 4.4330, -4.3300, -0.4330, 5.0000, 4.4430])
3rd element value: tensor(2.3340)
3rd element index: tensor(0)예제 2
다음 Python 프로그램은 텐서에서 상위 "k"개, 즉 가장 큰 "k"개의 요소를 찾는 방법을 보여줍니다.
# 텐서에서 상위 k개 요소를 찾는 Python 프로그램
# 필요한 라이브러리 임포트
import torch
# 1D 텐서 생성
T = torch.Tensor([2.334,4.433,-4.33,-0.433,5, 4.443])
print("Original Tensor:\n", T)
# 텐서에서 top k=2, 즉 가장 큰 2개의 요소 찾기
# 가장 큰 2개의 값과 원본 텐서에서의
# 인덱스를 반환합니다.
values, indices = torch.topk(T, 2)
# 상위 2개 요소의 값과 인덱스 출력
print("Top 2 element values:", values)
print("Top 2 element indices:", indices)출력 결과
Original Tensor:
tensor([ 2.3340, 4.4330, -4.3300, -0.4330, 5.0000, 4.4430])
Top 2 element values: tensor([5.0000, 4.4430])
Top 2 element indices: tensor([4, 5])