use candle_core::{Result, Tensor};
use candle_nn::{linear, linear_no_bias, ops, Linear, Module, VarBuilder};
pub struct QueryInput<'a> {
pub indices: &'a Tensor,
pub gate: &'a Tensor,
pub visible: &'a Tensor,
pub query_ids: &'a Tensor,
}
pub struct QueryRead {
pub residual: Tensor,
pub attention: Tensor,
}
pub struct QueryDecoder {
w_q: Linear,
w_k: Linear,
w_v: Linear,
e_mask: Tensor,
out: Linear,
rank: usize,
}
impl QueryDecoder {
pub fn new(embedding_dim: usize, rank: usize, vs: VarBuilder) -> Result<Self> {
let h = embedding_dim;
if rank == 0 || rank > h {
candle_core::bail!("query decoder: rank must be in 1..={h}, got {rank}");
}
let e_mask = vs.get_with_hints((1, h), "mask", candle_nn::init::DEFAULT_KAIMING_NORMAL)?;
Ok(Self {
w_q: linear_no_bias(h, rank, vs.pp("q"))?,
w_k: linear_no_bias(h, rank, vs.pp("k"))?,
w_v: linear_no_bias(h, rank, vs.pp("v"))?,
e_mask,
out: linear(rank, 1, vs.pp("out"))?,
rank,
})
}
#[must_use]
pub fn rank(&self) -> usize {
self.rank
}
pub fn forward(
&self,
features: &crate::candle::feature_embedding::FeatureEmbedding,
x: &QueryInput<'_>,
) -> Result<QueryRead> {
let (n, k) = x.indices.dims2()?;
let q = x.query_ids.dim(1)?;
let r = self.rank;
let h = features.embedding_dim();
let flat_idx = x.indices.flatten_all()?;
let gate_nk1 = x.gate.unsqueeze(2)?; let context = features.gather(&flat_idx)?; let keys = self
.w_k
.forward(&context)?
.reshape((n, k, r))?
.broadcast_mul(&gate_nk1)?; let values = self
.w_v
.forward(&context)?
.reshape((n, k, r))?
.broadcast_mul(&gate_nk1)?; let queries = features
.gather(&x.query_ids.flatten_all()?)?
.reshape((n, q, h))?
.broadcast_add(&self.e_mask)?; let qh = self.w_q.forward(&queries)?;
let scores = qh
.matmul(&keys.transpose(1, 2)?.contiguous()?)?
.affine(1.0 / (r as f64).sqrt(), 0.0)?; let neg_inf = x
.visible
.affine(-1.0, 1.0)?
.affine(-1e9, 0.0)?
.unsqueeze(1)?; let attn = ops::softmax(&scores.broadcast_add(&neg_inf)?, 2)?; let has_visible = x.visible.sum_keepdim(1)?.gt(0.0)?.to_dtype(attn.dtype())?; let attention = attn.broadcast_mul(&has_visible.unsqueeze(2)?)?;
let read = attention.matmul(&values)?; let residual = self
.out
.forward(&read)?
.squeeze(2)?
.broadcast_mul(&has_visible)?; Ok(QueryRead {
residual,
attention,
})
}
}
#[cfg(test)]
#[path = "query_decoder_tests.rs"]
mod query_decoder_tests;