• 企业400电话
  • 微网小程序
  • AI电话机器人
  • 电商代运营
  • 全 部 栏 目

    企业400电话 网络优化推广 AI电话机器人 呼叫中心 网站建设 商标✡知产 微网小程序 电商运营 彩铃•短信 增值拓展业务
    Pytorch 使用tensor特定条件判断索引

    torch.where() 用于将两个broadcastable的tensor组合成新的tensor,类似于c++中的三元操作符“?:”

    区别于python numpy中的where()直接可以找到特定条件元素的index

    想要实现numpy中where()的功能,可以借助nonzero()

    对应numpy中的where()操作效果:

    补充:Pytorch torch.Tensor.detach()方法的用法及修改指定模块权重的方法

    detach

    detach的中文意思是分离,官方解释是返回一个新的Tensor,从当前的计算图中分离出来

    需要注意的是,返回的Tensor和原Tensor共享相同的存储空间,但是返回的 Tensor 永远不会需要梯度

    import torch as t
    a = t.ones(10,)
    b = a.detach()
    print(b)
    tensor([1., 1., 1., 1., 1., 1., 1., 1., 1., 1.])
    

    那么这个函数有什么作用?

    –假如A网络输出了一个Tensor类型的变量a, a要作为输入传入到B网络中,如果我想通过损失函数反向传播修改B网络的参数,但是不想修改A网络的参数,这个时候就可以使用detcah()方法

    a = A(input)
    a = detach()
    b = B(a)
    loss = criterion(b, target)
    loss.backward()
    

    来看一个实际的例子:

    import torch as t
    x = t.ones(1, requires_grad=True)
    x.requires_grad   #True
    y = t.ones(1, requires_grad=True)
    y.requires_grad   #True
    x = x.detach()   #分离之后
    x.requires_grad   #False
    y = x+y         #tensor([2.])
    y.requires_grad   #我还是True
    y.retain_grad()   #y不是叶子张量,要加上这一行
    z = t.pow(y, 2)
    z.backward()    #反向传播
    y.grad        #tensor([4.])
    x.grad        #None
    

    以上代码就说明了反向传播到y就结束了,没有到达x,所以x的grad属性为None

    既然谈到了修改模型的权重问题,那么还有一种情况是:

    –假如A网络输出了一个Tensor类型的变量a, a要作为输入传入到B网络中,如果我想通过损失函数反向传播修改A网络的参数,但是不想修改B网络的参数,这个时候又应该怎么办了?

    这时可以使用Tensor.requires_grad属性,只需要将requires_grad修改为False即可.

    for param in B.parameters():
     param.requires_grad = False
    a = A(input)
    b = B(a)
    loss = criterion(b, target)
    loss.backward()
    

    以上为个人经验,希望能给大家一个参考,也希望大家多多支持脚本之家。如有错误或未考虑完全的地方,望不吝赐教。

    您可能感兴趣的文章:
    • Python深度学习之使用Pytorch搭建ShuffleNetv2
    • win10系统配置GPU版本Pytorch的详细教程
    • 浅谈pytorch中的nn.Sequential(*net[3: 5])是啥意思
    • pytorch visdom安装开启及使用方法
    • PyTorch CUDA环境配置及安装的步骤(图文教程)
    • pytorch中的nn.ZeroPad2d()零填充函数实例详解
    • 使用pytorch实现线性回归
    • pytorch实现线性回归以及多元回归
    • pytorch显存一直变大的解决方案
    • 在Windows下安装配置CPU版的PyTorch的方法
    • PyTorch两种安装方法
    • PyTorch的Debug指南
    上一篇:Python面向对象封装继承和多态示例讲解
    下一篇:python 实现简单的吃豆人游戏
  • 相关文章
  • 

    © 2016-2020 巨人网络通讯 版权所有

    《增值电信业务经营许可证》 苏ICP备15040257号-8

    Pytorch 使用tensor特定条件判断索引 Pytorch,使用,tensor,特定条件,