Advertisement

PyTorch基础乘法与神经网络功能

阅读量:
  • torch.mul(a, b) 表示对矩阵a与矩阵b进行逐元素相乘操作,其中a和b的维度需保持一致,运算后的输出矩阵维度与输入相同
    • torch.mm(a, b) 用于执行矩阵a与矩阵b之间的矩阵乘法运算,例如当a的维度为(1, 2),而b的维度为(2, 3)时,最终得到的矩阵维度将为(1, 3)

以实例说明:

复制代码
    import torch
    
    a = torch.rand(1, 4)
    b = torch.rand(4, 3)
    c = torch.rand(1, 3)
    
    print(a,'\n',b,'\n',c)
    ab = torch.mm(a, b)
    abc = torch.mul(ab, c)
    
    print("+"*10)
    print(ab)  # 返回1*3的tensor
    print(abc)  #返回1*3的tensor
    
    
      
      
      
      
      
      
      
      
      
      
      
      
      
    
复制代码
    tensor([[0.5638, 0.6036, 0.7021, 0.2039]]) 

全部评论 (0)

还没有任何评论哟~