PyTorch中tensor.squeeze()?使用小結(jié)
在深度學(xué)習(xí)中經(jīng)常會(huì)處理各種形狀(shape)復(fù)雜的張量(tensor)。
有時(shí)候,模型的輸入或輸出會(huì)多出一些沒用的“維度”,例如 (1, 3, 1, 224, 224)。
這時(shí),PyTorch 提供了一個(gè)非常實(shí)用的函數(shù) —— torch.squeeze(),可以幫我們輕松去除大小為 1 的維度。
本文將帶你從入門到實(shí)戰(zhàn),徹底掌握 squeeze() 的使用方法與常見坑。??
一、函數(shù)簡(jiǎn)介
官方定義
torch.squeeze(input, dim=None) → Tensor
作用:
返回一個(gè)新的張量,去掉所有大小為 1 的維度。
如果指定了 dim 參數(shù),則只會(huì)在那個(gè)維度上去除大小為 1 的維度。
squeeze 本身英文釋義如下:

二、為什么需要squeeze()?
在深度學(xué)習(xí)中,模型輸入輸出的維度往往需要嚴(yán)格匹配。
但是,在數(shù)據(jù)加載或卷積操作之后,可能會(huì)出現(xiàn)一些“冗余維度”。
舉個(gè)例子:
import torch x = torch.randn(1, 3, 1, 4) print(x.shape)
輸出:
torch.Size([1, 3, 1, 4])
可以看到,這個(gè)張量的第 0 和第 2 維都是大小為 1 的“空維度”。
這些維度不會(huì)存儲(chǔ)實(shí)際信息,但可能會(huì)導(dǎo)致維度不匹配錯(cuò)誤。
這時(shí),我們就可以用:
x.squeeze()
輸出:
torch.Size([3, 4])
? 所有大小為 1 的維度都被自動(dòng)去掉!
三、函數(shù)語法與參數(shù)說明
| 參數(shù) | 類型 | 說明 |
|---|---|---|
| input | Tensor | 輸入張量 |
| dim | int, 可選 | 指定要壓縮的維度 |
| 返回值 | Tensor | 新張量(共享存儲(chǔ),不復(fù)制數(shù)據(jù)) |
四、示例講解
1、去除所有大小為 1 的維度
x = torch.randn(1, 3, 1, 4, 1)
print("原形狀:", x.shape)
y = torch.squeeze(x)
print("壓縮后:", y.shape)
輸出:
原形狀: torch.Size([1, 3, 1, 4, 1])
壓縮后: torch.Size([3, 4])
說明:squeeze() 去掉了所有維度為 1 的軸。
2、指定某個(gè)維度壓縮
有時(shí)候我們不想去掉所有維度,只想處理特定的一個(gè)。
x = torch.randn(1, 3, 1, 4)
print("原形狀:", x.shape)
y = torch.squeeze(x, dim=0)
print("壓縮后:", y.shape)
輸出:
原形狀: torch.Size([1, 3, 1, 4])
壓縮后: torch.Size([3, 1, 4])
?? 只去掉了第 0 維,因?yàn)樗拇笮∈?1。
其他維度保持不變。
3、如果指定的維度不是 1,會(huì)怎樣?
x = torch.randn(2, 1, 3) y = torch.squeeze(x, dim=0) print(y.shape)
輸出:
torch.Size([2, 1, 3])
沒有任何變化,因?yàn)榈?0 維的大小是 2,不是 1。squeeze() 只會(huì)壓縮大小為 1 的維度,不會(huì)報(bào)錯(cuò)。
五、與unsqueeze()的關(guān)系
如果說 squeeze() 是“去掉維度”,
那 unsqueeze() 就是“增加維度”。
x = torch.tensor([1, 2, 3]) print(x.shape) # torch.Size([3]) y = x.unsqueeze(0) print(y.shape) # torch.Size([1, 3]) z = y.squeeze(0) print(z.shape) # torch.Size([3])
? unsqueeze() 與 squeeze() 是一對(duì)反操作。
一個(gè)增加維度,一個(gè)去除維度。
六、常見應(yīng)用場(chǎng)景
1、數(shù)據(jù)集加載時(shí)去掉多余維度
# 讀取圖片后通常是 (1, H, W) img = torch.randn(1, 224, 224) img = img.squeeze(0) print(img.shape) # torch.Size([224, 224])
2、模型輸出后去掉 batch 維度
# 例如分類模型輸出 [1, num_classes] output = torch.randn(1, 10) pred = output.squeeze(0) print(pred.shape) # torch.Size([10])
3、多維卷積層結(jié)果調(diào)整
在 Conv2d、LSTM 等層輸出中,有時(shí)需要將 [batch, seq_len, 1] 變成 [batch, seq_len]。
out = torch.randn(32, 100, 1) out = out.squeeze(-1) print(out.shape) # torch.Size([32, 100])
?? 七、注意事項(xiàng)與坑點(diǎn)
| 問題 | 說明 |
|---|---|
| ? 誤刪維度 | 默認(rèn)不傳 dim 會(huì)刪除所有大小為 1 的維度,可能導(dǎo)致形狀變化過多 |
| ? 建議 | 當(dāng)只想去掉某個(gè)維度時(shí),一定要寫 dim 參數(shù) |
| ?? 內(nèi)存共享 | squeeze() 返回的張量與原張量共享內(nèi)存,不會(huì)復(fù)制數(shù)據(jù) |
八、擴(kuò)展:與 NumPy 對(duì)比
PyTorch 的 squeeze() 和 NumPy 的 numpy.squeeze() 功能幾乎一致。
import numpy as np a = np.random.randn(1, 3, 1, 4) print(a.shape) # (1, 3, 1, 4) print(a.squeeze().shape) # (3, 4)
如果熟悉 NumPy 的用法,PyTorch 中也能無縫銜接。
九、總結(jié)
| 功能 | 說明 |
|---|---|
| 函數(shù) | torch.squeeze(input, dim=None) |
| 作用 | 刪除大小為 1 的維度 |
| 參數(shù) | dim:指定要壓縮的維度(可選) |
| 返回 | 新張量(共享內(nèi)存) |
| 反操作 | unsqueeze() |
| 常用場(chǎng)景 | 模型輸出處理、數(shù)據(jù)預(yù)處理、維度調(diào)整 |
十、參考資料

到此這篇關(guān)于PyTorch中tensor.squeeze() 使用小結(jié)的文章就介紹到這了,更多相關(guān)PyTorch tensor.squeeze()內(nèi)容請(qǐng)搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!
- pytorch中tensor.expand()和tensor.expand_as()函數(shù)詳解
- PyTorch中 tensor.detach() 和 tensor.data 的區(qū)別詳解
- pytorch中Tensor.to(device)和model.to(device)的區(qū)別及說明
- pytorch中torch.max和Tensor.view函數(shù)用法詳解
- PyTorch中tensor.backward()函數(shù)的詳細(xì)介紹及功能實(shí)現(xiàn)
- PyTorch中關(guān)于tensor.repeat()的使用
- PyTorch中 tensor.detach() 和 tensor.data 的區(qū)別解析
- pytorch中Tensor.new()的使用解析
- pytorch中函數(shù)tensor.numpy()的數(shù)據(jù)類型解析
相關(guān)文章
pandas進(jìn)行數(shù)據(jù)輸入和輸出的方法詳解
這篇文章主要為大家詳細(xì)介紹了pandas進(jìn)行數(shù)據(jù)輸入和輸出的方法,文中示例代碼介紹的非常詳細(xì),具有一定的參考價(jià)值,感興趣的小伙伴們可以參考一下,希望能夠給你帶來幫助2022-03-03
Python通過fnmatch模塊實(shí)現(xiàn)文件名匹配
這篇文章主要介紹了Python通過fnmatch模塊實(shí)現(xiàn)文件名匹配,文中通過示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下2020-09-09
Python3標(biāo)準(zhǔn)庫之threading進(jìn)程中管理并發(fā)操作方法
這篇文章主要介紹了Python3標(biāo)準(zhǔn)庫之threading進(jìn)程中管理并發(fā)操作方法,本文通過實(shí)例代碼給大家介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或工作具有一定的參考借鑒價(jià)值,需要的朋友可以參考下2020-03-03
Django 配置多站點(diǎn)多域名的實(shí)現(xiàn)步驟
這篇文章主要介紹了Django 配置多站點(diǎn)多域名的實(shí)現(xiàn)步驟,文中通過示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧2019-05-05
python將數(shù)據(jù)插入數(shù)據(jù)庫的代碼分享
在本篇文章里小編給大家整理的是關(guān)于python將數(shù)據(jù)插入數(shù)據(jù)庫的代碼內(nèi)容,有興趣的朋友們可以參考下。2020-08-08

