use crate::candle::fast_index::{gather_rows, scatter_add_cols};
use crate::candle::feature_embedding::FeatureEmbedding;
use candle_core::{Result, Tensor};
pub fn masked_scores(scores: &Tensor, visible: &Tensor) -> Result<Tensor> {
scores + visible.affine(1e9, -1e9)?
}
pub fn attention_scores_from_vector(
gate_nk: &Tensor,
idx_nk: &Tensor,
rq_d: &Tensor,
visible_nk: &Tensor,
scale: f64,
) -> Result<Tensor> {
let (n, k) = gate_nk.dims2()?;
let rq_nk = gather_rows(rq_d, &idx_nk.flatten_all()?.contiguous()?)?.reshape((n, k))?;
let scores = (gate_nk * rq_nk)?.affine(scale, 0.0)?;
masked_scores(&scores, visible_nk)
}
pub fn pool_by_scatter(
attn_nk: &Tensor,
gate_nk: &Tensor,
idx_nk: &Tensor,
features: &FeatureEmbedding,
) -> Result<Tensor> {
Ok(pool_parts(attn_nk, gate_nk, idx_nk, features)?.2)
}
fn pool_parts(
attn_nk: &Tensor,
gate_nk: &Tensor,
idx_nk: &Tensor,
features: &FeatureEmbedding,
) -> Result<(Tensor, Tensor, Tensor)> {
let w_nk = (attn_nk * gate_nk)?; let w_nd = scatter_add_cols(idx_nk, &w_nk, features.n_features())?; let pool_nh = features.map_rows_linear(|rows| w_nd.matmul(rows))?; Ok((w_nk, w_nd, pool_nh))
}
pub fn query_over_features(features: &FeatureEmbedding, attn_query_1h: &Tensor) -> Result<Tensor> {
let h = attn_query_1h.elem_count();
let q_h1 = attn_query_1h.reshape((h, 1))?;
features.project_dims(&q_h1)?.squeeze(1)
}
#[cfg(test)]
#[path = "scatter_pool_tests.rs"]
mod scatter_pool_tests;