PyTorch torch.unique() 基礎(chǔ)與實戰(zhàn)應(yīng)用指南
在深度學(xué)習(xí)的數(shù)據(jù)處理中經(jīng)常需要統(tǒng)計或篩選 張量(Tensor) 中的唯一值,比如去重、統(tǒng)計類別數(shù)量、計算唯一標(biāo)簽數(shù)等。
PyTorch 提供了一個非常方便的函數(shù) —— torch.unique(),可以輕松完成這些操作。
本文將帶你深入了解 torch.unique() 的用法、參數(shù)、返回值以及實際應(yīng)用場景。
一、什么是torch.unique()?
torch.unique() 是 PyTorch 中的一個去重函數(shù),用于返回張量中所有的唯一元素(unique elements)。
它類似于 Python 的 set() 或 NumPy 的 np.unique(),但專為 GPU 加速的張量操作 設(shè)計。
二、函數(shù)語法
torch.unique(input, sorted=True, return_inverse=False, return_counts=False, dim=None)
三、參數(shù)說明
| 參數(shù) | 類型 | 說明 |
|---|---|---|
input | Tensor | 輸入張量 |
sorted | bool | 是否對結(jié)果排序(默認(rèn) True) |
return_inverse | bool | 是否返回原張量中每個值在唯一值列表中的索引 |
return_counts | bool | 是否返回每個唯一值的出現(xiàn)次數(shù) |
dim | int 或 None | 按指定維度去重,默認(rèn)對整個張量去重 |
四、基本用法
?? 示例 1:基礎(chǔ)去重
import torch x = torch.tensor([1, 2, 2, 3, 3, 3]) unique_x = torch.unique(x) print(unique_x)
輸出:
tensor([1, 2, 3])
? 結(jié)果去除了重復(fù)值,并自動排序。
?? 示例 2:不排序
x = torch.tensor([3, 2, 1, 3, 2]) unique_x = torch.unique(x, sorted=False) print(unique_x)
輸出:
tensor([3, 2, 1])
當(dāng) sorted=False 時,結(jié)果的順序與首次出現(xiàn)的順序一致。
五、返回索引與計數(shù)
?? 示例 3:return_inverse
return_inverse=True 會返回一個索引張量,表示原張量中每個元素在唯一值(即新張量)中的位置。
x = torch.tensor([2, 1, 2, 3]) u, inv = torch.unique(x, return_inverse=True) print(u) print(inv)
輸出:
tensor([1, 2, 3]) tensor([1, 0, 1, 2])
解釋:
- 唯一值為
[1, 2, 3] - 原數(shù)組
[2, 1, 2, 3]中:- 第一個元素 2 → 索引 1
- 第二個元素 1 → 索引 0
- 第三個元素 2 → 索引 1
- 第四個元素 3 → 索引 2
?? 示例 4:return_counts
return_counts=True 會返回每個唯一值出現(xiàn)的次數(shù)。
x = torch.tensor([1, 2, 2, 3, 3, 3]) u, counts = torch.unique(x, return_counts=True) print(u) print(counts)
輸出:
tensor([1, 2, 3]) tensor([1, 2, 3])
表示:
- 值 1 出現(xiàn) 1 次
- 值 2 出現(xiàn) 2 次
- 值 3 出現(xiàn) 3 次
?? 示例 5:同時返回多個結(jié)果
你可以同時返回 unique 值、inverse 索引和計數(shù):
x = torch.tensor([1, 2, 2, 3, 3, 3]) u, inv, counts = torch.unique(x, return_inverse=True, return_counts=True) print(u) print(inv) print(counts)
輸出:
tensor([1, 2, 3]) tensor([0, 1, 1, 2, 2, 2]) tensor([1, 2, 3])
六、按維度去重(dim 參數(shù))
默認(rèn)情況下,torch.unique() 會將張量展開成一維后去重。
但如果你希望在特定維度上去重(如按行或按列),可以使用 dim 參數(shù)。
?? 示例 6:按行去重
x = torch.tensor([[1, 2],
[1, 2],
[3, 4]])
unique_rows = torch.unique(x, dim=0)
print(unique_rows)輸出:
tensor([[1, 2],
[3, 4]])
表示第 1、2 行重復(fù),只保留一個。
?? 示例 7:按列去重
x = torch.tensor([[1, 1, 3],
[2, 2, 4]])
unique_cols = torch.unique(x, dim=1)
print(unique_cols)輸出:
tensor([[1, 3],
[2, 4]])
七、torch.unique()與 NumPy 對比
| 功能 | PyTorch (torch.unique) | NumPy (np.unique) |
|---|---|---|
| 默認(rèn)排序 | ? 是 | ? 是 |
| 支持 GPU | ? 是 | ? 否 |
| 返回 inverse 索引 | ? 是 | ? 是 |
| 返回 counts | ? 是 | ? 是 |
| 按維度去重 | ? 是(dim) | ? 不直接支持 |
| 性能 | 高(GPU 支持) | 僅 CPU |
八、實際應(yīng)用場景
1. 分類問題中統(tǒng)計類別數(shù)量
labels = torch.tensor([0, 1, 0, 2, 2, 1, 3])
classes = torch.unique(labels)
print(f"共有 {len(classes)} 個類別: {classes.tolist()}")
輸出:
共有 4 個類別: [0, 1, 2, 3]
2. 計算樣本分布(類別頻率)
labels = torch.tensor([0, 1, 0, 2, 2, 1, 3])
u, counts = torch.unique(labels, return_counts=True)
for c, cnt in zip(u.tolist(), counts.tolist()):
print(f"類別 {c}: {cnt} 個樣本")
輸出:
類別 0: 2 個樣本 類別 1: 2 個樣本 類別 2: 2 個樣本 類別 3: 1 個樣本
3. 在圖像分割中統(tǒng)計像素類別
例如在語義分割任務(wù)中,計算 mask 圖像中有多少個不同的像素類別:
mask = torch.randint(0, 5, (256, 256)) # 隨機生成類別標(biāo)簽
num_classes = len(torch.unique(mask))
print(f"圖像中共有 {num_classes} 個類別")
?? 九、注意事項
torch.unique()** 默認(rèn)會對結(jié)果排序**,如果在意性能,可以設(shè)置sorted=False。- 對高維張量使用
dim去重時,必須保證該維度的所有元素形狀一致。 - 對大張量使用
return_counts或return_inverse時可能會消耗更多顯存。
?? 參考資料


到此這篇關(guān)于PyTorch torch.unique() 基礎(chǔ)與實戰(zhàn)應(yīng)用指南的文章就介紹到這了,更多相關(guān)PyTorch torch.unique() 使用內(nèi)容請搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!
相關(guān)文章
python實現(xiàn)逆波蘭計算表達(dá)式實例詳解
這篇文章主要介紹了python實現(xiàn)逆波蘭計算表達(dá)式的方法,較為詳細(xì)的分析了逆波蘭表達(dá)式的概念及實現(xiàn)技巧,具有一定參考借鑒價值,需要的朋友可以參考下2015-05-05
使用Python對微信好友進(jìn)行數(shù)據(jù)分析
這篇文章主要介紹了使用Python對微信好友進(jìn)行數(shù)據(jù)分析的實現(xiàn)代碼,非常不錯,具有一定的參考借鑒價值,需要的朋友可以參考下2018-06-06

