torch.sum

exp指数函数,把所有的y都整到0以上,不用担心-0.5和0.5抵消的问题

    import torch
    import numpy as np

    data=np.array([[[0.5,-0.5],[-0.05,-0.05]]])
    x = torch.from_numpy(data.astype(np.float32))

    aaa=torch.exp(torch.Tensor([1]))
    print(aaa)
    print(x)

    print(x.sum(2))  # 按行求和

    x=torch.exp(x)
    print(x.sum(2))  # 按行求和
    print(x.sum(2)[x.sum(2)>aaa])  # 按行求和
发布了2718 篇原创文章 · 获赞 1004 · 访问量 536万+

猜你喜欢

转载自blog.csdn.net/jacke121/article/details/104709153