文章详情

短信预约-IT技能 免费直播动态提醒

请输入下面的图形验证码

提交验证

短信预约提醒成功

pytorch网络模型构建场景的问题如何解决

2023-07-05 11:27

关注

今天小编给大家分享一下pytorch网络模型构建场景的问题如何解决的相关知识点,内容详细,逻辑清晰,相信大部分人都还太了解这方面的知识,所以分享这篇文章给大家参考一下,希望大家阅读完这篇文章后有所收获,下面我们一起来了解一下吧。

网络模型构建中的问题

1 输入变量是Tensor张量

各个模块和网络模型的输入, 一定要是tensor 张量;

可以用一个列表存放多个张量。

如果是张量维度不够,需要升维度,

可以先使用 torch.unsqueeze(dim = expected)

然后再使用torch.cat(dim ) 进行拼接;

需要传递梯度的数据,禁止使用numpy, 也禁止先使用numpy,然后再转换成张量的这种情况出现;

这是因为pytorch的机制是只有是 Tensor 张量的类型,才会有梯度等属性值,如果是numpy这些类别,这些变量并会丢失其梯度值。

2 __init__()方法使用

class ex:    def __init__(self):        pass

__init__方法必须接受至少一个参数即self,

Python中,self是指向该对象本身的一个引用,

通过在类的内部使用self变量,

类中的方法可以访问自己的成员变量,简单来说,self.varname的意义为”访问该对象的varname属性“

当然,__init__()中可以封装任意的程序逻辑,这是允许的,init()方法还接受任意多个其他参数,允许在初始化时提供一些数据,例如,对于刚刚的worker类,可以这样写:

class worker:    def __init__(self,name,pay):        self.name=name        self.pay=pay

这样,在创建worker类的对象时,必须提供name和pay两个参数:

b=worker('Jim',5000)

Python会自动调用worker.init()方法,并传递参数。

细节参考这里init方法

3 内置函数setattr()

此时,可以使用python自带的内置函数 setattr(), 和对应的getattr()

setattr(object, name, value)

object – 对象。

name – 字符串,对象属性。

value – 属性值。

对已存在的属性进行赋值:
>>>class A(object):
...     bar = 1
... 
>>> a = A()
>>> getattr(a, 'bar')          # 获取属性 bar 值
1
>>> setattr(a, 'bar', 5)       # 设置属性 bar 值
>>> a.bar
5
如果属性不存在会创建一个新的对象属性,并对属性赋值:

>>>class A():
...     name = "runoob"
... 
>>> a = A()
>>> setattr(a, "age", 28)
>>> print(a.age)
28
>>>

setattr() 语法

setattr(object, name, value)

object – 对象。

name – 字符串,对象属性。

value – 属性值。

4 网络模型的构建

注意到, 在python的 __init__() 函数中, self 本身就是该类的对象的一个引用,即self是指向该对象本身的一个引用,

利用上述这一点,当在神经网络中,

需要给多个属性进行实例化时,

且这多个属性使用的是同一个类进行实例化.

则使用 setattr(self, string, object1) 添加属性;

class Temporal_GroupTrans(nn.Module):    def __init__(self,   num_classes=10,num_groups=35, drop_prob=0.5, pretrained= True):        super(Temporal_GroupTrans, self).__init__()        conv_block = Basic_slide_conv()        for i in range( num_groups):            setattr(self, "group" + str(i), conv_block)        # 自定义transformer模型的初始化, CustomTransformerModel() 在该类中传入初始化模型的参数,        # nip:512 输入序列中,每个列向量的编码维度, 16: 注意力头的个数        # 600: 中间mlp 隐藏层的维数,  6: 堆叠transforEncode 编码模块的个数;        self.trans_model = CustomTransformerModel(512,16,600, 6,droupout=0.5,nclass=4)

则使用 getattr(self, string, object1) 获取属性;

        trans_input_sequence = []        for i in range(0, num_groups, ):            #   每组语谱图的大小是一个 (bt, ch,96,12)的矩阵,组与组之间没有重叠;            cur_group = x[:, :, :, 12 * i:12 * (i + 1)]            # VARIABLE_fun = "self.group"   # 每一组,与之对应的卷积模块;            # cur_fun = eval(VARIABLE_fun + str(i ))            cur_fun = getattr(self, 'group'+str(i))            cur_group_out = cur_fun(cur_group).unsqueeze(dim=1)  # [bt,1, 512]            trans_input_sequence.append(cur_group_out)

以上就是“pytorch网络模型构建场景的问题如何解决”这篇文章的所有内容,感谢各位的阅读!相信大家阅读完这篇文章都有很大的收获,小编每天都会为大家更新不同的知识,如果还想学习更多的知识,请关注编程网行业资讯频道。

阅读原文内容投诉

免责声明:

① 本站未注明“稿件来源”的信息均来自网络整理。其文字、图片和音视频稿件的所属权归原作者所有。本站收集整理出于非商业性的教育和科研之目的,并不意味着本站赞同其观点或证实其内容的真实性。仅作为临时的测试数据,供内部测试之用。本站并未授权任何人以任何方式主动获取本站任何信息。

② 本站未注明“稿件来源”的临时测试数据将在测试完成后最终做删除处理。有问题或投稿请发送至: 邮箱/279061341@qq.com QQ/279061341

软考中级精品资料免费领

  • 历年真题答案解析
  • 备考技巧名师总结
  • 高频考点精准押题
  • 2024年上半年信息系统项目管理师第二批次真题及答案解析(完整版)

    难度     807人已做
    查看
  • 【考后总结】2024年5月26日信息系统项目管理师第2批次考情分析

    难度     351人已做
    查看
  • 【考后总结】2024年5月25日信息系统项目管理师第1批次考情分析

    难度     314人已做
    查看
  • 2024年上半年软考高项第一、二批次真题考点汇总(完整版)

    难度     433人已做
    查看
  • 2024年上半年系统架构设计师考试综合知识真题

    难度     221人已做
    查看

相关文章

发现更多好内容

猜你喜欢

AI推送时光机
位置:首页-资讯-后端开发
咦!没有更多了?去看看其它编程学习网 内容吧
首页课程
资料下载
问答资讯