pytorch 在sequential中使用view來reshape的例子
更新時間:2019年08月20日 08:54:52 作者:青盞
今天小編就為大家分享一篇pytorch 在sequential中使用view來reshape的例子,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
pytorch中view是tensor方法,然而在sequential中包裝的是nn.module的子類,
因此需要自己定義一個方法:
import torch.nn as nn class Reshape(nn.Module): def __init__(self, *args): super(Reshape, self).__init__() self.shape = args def forward(self, x): # 如果數(shù)據(jù)集最后一個batch樣本數(shù)量小于定義的batch_batch大小,會出現(xiàn)mismatch問題。可以自己修改下,如只傳入后面的shape,然后通過x.szie(0),來輸入。 return x.view(self.shape)
class Reshape(nn.Module): def __init__(self, *args): super(Reshape, self).__init__() self.shape = args def forward(self, x): return x.view((x.size(0),)+self.shape)
以上這篇pytorch 在sequential中使用view來reshape的例子就是小編分享給大家的全部內容了,希望能給大家一個參考,也希望大家多多支持腳本之家。

