Pytorch:torch.diag()創(chuàng)建對角線張量方式
更新時間:2024年06月27日 15:06:06 作者:湫兮之風
這篇文章主要介紹了Pytorch:torch.diag()創(chuàng)建對角線張量方式,具有很好的參考價值,希望對大家有所幫助,如有錯誤或未考慮完全的地方,望不吝賜教
Pytorch torch.diag()創(chuàng)建對角線張量
torch.diag()
torch.diag是PyTorch中的一個函數,用于從給定的矩陣中提取對角線元素,或者構造一個以給定對角線元素為值的對角矩陣。這個函數對于矩陣分解和轉換等操作非常重要。
如果輸入是一個向量(1D張量),torch.diag會返回一個以該向量為對角線元素的2D方陣。如果輸入是一個矩陣(2D張量),則返回一個包含輸入矩陣對角線元素的1D張量。
torch.diag還允許你指定對角線的位置,通過參數diagonal實現。如果diagonal=0,則為主對角線;如果diagonal>0,則為位于主對角線之上的對角線;如果diagonal<0,則為位于主對角線之下的對角線。
語法:
input (Tensor):輸入張量。diagonal (int, optional):指定的對角線。out (Tensor, optional):輸出張量。
舉例一:
import torch data = torch.tensor([1,2,3,4]) data_two = torch.diag(data,0) print(data_two)
結果:

舉例二:
import torch
data = torch.tensor(float('inf')).cuda().repeat(3)
data_two = torch.diag(data,0)
print(data_two)結果:

torch.diag()取矩陣對角線元素,torch.diag_embed()指定值變成對角矩陣
1、torch.diag()
import torch
a = torch.randn(3, 3)
print(a)
tensor([[ 0.7594, 0.8073, -0.1344],
[-1.7335, -0.4356, -0.0055],
[ 1.8326, 0.3900, -0.9933]])
diag = torch.diag(a) # 取 a 對角線元素,輸出為 1*3
print(diag)
tensor([ 0.7594, -0.4356, -0.9933])2、torch.diag_embed()
import torch
tensor([ 0.7594, -0.4356, -0.9933])
a_diag = torch.diag_embed(diag) # 由 diag 變?yōu)槿S 3*3
tensor([[ 0.7594, 0.0000, 0.0000],
[ 0.0000, -0.4356, 0.0000],
[ 0.0000, 0.0000, -0.9933]])總結
以上為個人經驗,希望能給大家一個參考,也希望大家多多支持腳本之家。
相關文章
使用python搭建服務器并實現Android端與之通信的方法
今天小編就為大家分享一篇使用python搭建服務器并實現Android端與之通信的方法,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧2019-06-06
Python使用enum模塊獲取枚舉成員索引號的四種方法詳解
在 Python 中,可以使用 enum 模塊創(chuàng)建枚舉類型,并通過遍歷枚舉成員來獲取其索引號,本文介紹了常用的四種方法,大家可以根據自己的需要進行選擇2026-04-04

