rust-candle学习笔记12-实现因果注意力

参考:about-pytorch

定义结构体:

rust 复制代码
struct CausalAttention {
    w_qkv: Linear,
    dropout: Dropout, 
    d_model: Tensor,
    mask: Tensor,
    device: Device,   
}

定义new方法:

rust 复制代码
impl CausalAttention {
    fn new(vb: VarBuilder, embedding_dim: usize, out_dim: usize, seq_len: usize, dropout: f32, device: Device) -> Result<Self> {
        Ok(Self { 
            w_qkv: linear_no_bias(embedding_dim, 3*out_dim, vb.pp("w_qkv"))?,
            d_model: Tensor::new(embedding_dim as f32, &device)?,
            mask: Tensor::tril2(seq_len, DType::U32, &device)?,
            dropout: Dropout::new(dropout),
            device
        })
    }
}

定义forward方法:

rust 复制代码
    fn forward(&self, x: &Tensor, train: bool) -> Result<Tensor> { 
        let qkv = self.w_qkv.forward(x)?;
        let (batch_size, seq_len, _) = qkv.dims3()?;
        let qkv = qkv.reshape((batch_size, seq_len, 3, ()))?;
        let q = qkv.get_on_dim(2, 0)?;
        let q = q.reshape((batch_size, seq_len, ()))?;
        let k = qkv.get_on_dim(2, 1)?;
        let k = k.reshape((batch_size, seq_len, ()))?;
        let v = qkv.get_on_dim(2, 2)?;
        let v = v.reshape((batch_size, seq_len, ()))?;
        let mut attn_score = q.matmul(&k.t()?)?;
        // println!("attn_score: {:?}\n", attn_score.to_vec3::<f32>()?);
        let dim = attn_score.rank() - 1;
        let mask_dim = attn_score.dims()[dim];
        let mask = self.mask.broadcast_as(attn_score.shape())?;
        // println!("mask: {:?}\n", mask);
        // println!("mask: {:?}\n", mask.to_vec3::<u32>()?);
        attn_score = masked_fill(&attn_score, &mask, f32::NEG_INFINITY)?;
        // println!("attn_score: {:?}\n", attn_score);
        // println!("attn_score: {:?}\n", attn_score.to_vec3::<f32>()?);
        let attn_score = attn_score.broadcast_div(&self.d_model.sqrt()?)?; 
        let attn_weights = ops::softmax(&attn_score, dim)?;
        // println!("attn_weights: {:?}\n", attn_weights);
        // println!("attn_weights: {:?}\n", attn_weights.to_vec3::<f32>()?); 
        let attn_weights = self.dropout.forward(&attn_weights, train)?;
        // println!("dropout attn_weights: {:?}\n", attn_weights);
        // println!("dropout attn_weights: {:?}\n", attn_weights.to_vec3::<f32>()?); 
        let attn_output = attn_weights.matmul(&v)?;
        Ok(attn_output)
    }

测试:

rust 复制代码
fn main() -> Result<()> {
    let device = Device::cuda_if_available(0)?;
    let varmap = VarMap::new();
    let vb = VarBuilder::from_varmap(&varmap, candle_core::DType::F32, &device);
    
    let input = Tensor::from_vec(vec![0.43f32, 0.15, 0.89, 
                                                    0.55, 0.87, 0.66,
                                                    0.57, 0.85, 0.64,
                                                    0.22, 0.58, 0.33,
                                                    0.77, 0.25, 0.10,
                                                    0.05, 0.80, 0.55, 
                                                    0.43, 0.15, 0.89, 
                                                    0.55, 0.87, 0.66,
                                                    0.57, 0.85, 0.64,
                                                    0.22, 0.58, 0.33,
                                                    0.77, 0.25, 0.10,
                                                    0.05, 0.80, 0.55], (2, 6, 3), &device)?;
    let model = CausalAttention::new(vb.clone(), 3, 2, 6, 0.5, device.clone())?;
    let output = model.forward(&input, true)?;
    println!("output: {:?}\n", output);
    println!("output: {:?}\n", output.to_vec3::<f32>()?);
    Ok(())
}
相关推荐
.小小陈.10 分钟前
数据结构2:单链表
c语言·开发语言·数据结构·笔记·学习方法
立志成为大牛的小牛18 分钟前
数据结构——二十三、并查集的终极优化(王道408)
开发语言·数据结构·笔记·学习·程序人生·考研
全栈游侠31 分钟前
04-优先级与延时链表
笔记
im_AMBER1 小时前
React 01
前端·javascript·笔记·react.js·前端框架·web
稻草猫.1 小时前
文件 IO
java·笔记·后端·java-ee·idea
QT 小鲜肉1 小时前
【个人成长笔记】Qt Creator快捷键终极指南:从入门到精通
开发语言·c++·笔记·qt·学习·学习方法
takashi_void2 小时前
本地实现斯坦福小镇(利用大语言模型使虚拟角色自主发展剧情)类似项目“Microverse”
人工智能·语言模型·自然语言处理·godot·游戏程序·斯坦福小镇
lkbhua莱克瓦242 小时前
Java基础——面向对象进阶复习知识点8
java·笔记·github·学习方法
Costrict2 小时前
解锁新阵地!CoStrict 现已支持 JetBrains 系列 IDE
大数据·ide·人工智能·深度学习·自然语言处理·ai编程·visual studio
QT 小鲜肉4 小时前
【数据结构与算法基础】05. 栈详解(C++ 实战)
开发语言·数据结构·c++·笔记·学习·算法·学习方法