Skip to content

Commit 39781c4

Browse files
committed
feat: add 2D and fft model
1 parent 448311c commit 39781c4

10 files changed

Lines changed: 1384 additions & 1158 deletions

File tree

‎src/attend.rs‎

Lines changed: 73 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,73 @@
1+
use candle_core::{DType, Result, Tensor, D};
2+
3+
#[derive(Debug)]
4+
pub struct Attend {
5+
scale: Option<f64>,
6+
drop_p: f64,
7+
#[allow(dead_code)]
8+
heads: usize,
9+
is_flash: bool,
10+
is_causal: bool,
11+
}
12+
13+
impl Attend {
14+
pub fn new(
15+
drop_p: Option<f64>,
16+
heads: Option<usize>,
17+
scale: Option<f64>,
18+
is_flash: Option<bool>,
19+
is_causal: Option<bool>,
20+
) -> Self {
21+
Self {
22+
scale,
23+
drop_p: drop_p.unwrap_or(0.0),
24+
heads: heads.unwrap_or(8),
25+
is_flash: is_flash.unwrap_or(false),
26+
is_causal: is_causal.unwrap_or(false),
27+
}
28+
}
29+
30+
pub fn forward_t(&self, q: &Tensor, k: &Tensor, v: &Tensor, train: bool) -> Result<Tensor> {
31+
let q_dims = q.dims();
32+
let scale = self.scale.unwrap_or((q_dims[q_dims.len() - 1] as f64).sqrt());
33+
34+
if self.is_flash {
35+
// Fallback to standard attention for now
36+
}
37+
38+
let q = q.contiguous()?;
39+
let k_t = k.transpose(D::Minus2, D::Minus1)?.contiguous()?;
40+
let mut sim = q.matmul(&k_t)?;
41+
sim = (sim * (1.0 / scale))?;
42+
43+
if self.is_causal {
44+
let sim_dims = sim.dims();
45+
let (i, j) = (sim_dims[sim_dims.len() - 2], sim_dims[sim_dims.len() - 1]);
46+
// upper triangular mask
47+
let mut mask_data = vec![0u8; i * j];
48+
for row in 0..i {
49+
for col in 0..j {
50+
if col as i64 > row as i64 + (j as i64 - i as i64) {
51+
mask_data[row * j + col] = 1;
52+
}
53+
}
54+
}
55+
let mask = Tensor::from_vec(mask_data, (i, j), sim.device())?
56+
.to_dtype(DType::F32)?;
57+
let mask_value = f32::NEG_INFINITY;
58+
let mask = mask.broadcast_as(sim.dims())?;
59+
sim = sim.where_cond(&mask.eq(0.0)?, &Tensor::new(mask_value, sim.device())?)?;
60+
}
61+
62+
let attn = candle_nn::ops::softmax_last_dim(&sim)?;
63+
let attn = if train && self.drop_p > 0.0 {
64+
candle_nn::ops::dropout(&attn, self.drop_p as f32)?
65+
} else {
66+
attn
67+
};
68+
69+
let attn = attn.contiguous()?;
70+
let v = v.contiguous()?;
71+
attn.matmul(&v)
72+
}
73+
}

‎src/attention.rs‎

Lines changed: 218 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,218 @@
1+
use candle_core::{IndexOp, Result, Tensor};
2+
use candle_nn::{
3+
layer_norm, linear_no_bias, LayerNorm, LayerNormConfig, Linear, Module, VarBuilder,
4+
};
5+
6+
use crate::attend::Attend;
7+
8+
#[derive(Debug)]
9+
pub struct ToQKV {
10+
pub(crate) linear: Linear,
11+
pub(crate) heads: usize,
12+
}
13+
14+
impl ToQKV {
15+
pub fn new(vb: VarBuilder, dim: usize, hidden_size: usize, heads: usize) -> Result<Self> {
16+
Ok(Self {
17+
linear: linear_no_bias(dim, hidden_size * 3, vb)?,
18+
heads,
19+
})
20+
}
21+
22+
pub fn rearrange(&self, xs: &Tensor) -> Result<(Tensor, Tensor, Tensor)> {
23+
let xs_dims = xs.dims();
24+
let (h, qkv) = (self.heads, 3);
25+
let b = xs_dims[0];
26+
let n = xs_dims[1];
27+
let total_dim = xs_dims[2];
28+
let dim_head = total_dim / (qkv * h);
29+
let xs = xs.reshape((b, n, qkv, h, dim_head))?;
30+
let xs = xs.permute((2, 0, 3, 1, 4))?;
31+
32+
let q = xs.i(0)?;
33+
let k = xs.i(1)?;
34+
let v = xs.i(2)?;
35+
36+
Ok((q, k, v))
37+
}
38+
39+
pub fn forward(&self, xs: &Tensor) -> Result<(Tensor, Tensor, Tensor)> {
40+
let xs = self.linear.forward(xs)?;
41+
self.rearrange(&xs)
42+
}
43+
}
44+
45+
#[derive(Debug)]
46+
pub struct ToValueResidualMix {
47+
pub(crate) linear: Linear,
48+
}
49+
50+
impl ToValueResidualMix {
51+
pub fn new(vb: VarBuilder, dim: usize, heads: usize) -> Result<Self> {
52+
Ok(Self {
53+
linear: linear_no_bias(dim, heads, vb)?,
54+
})
55+
}
56+
57+
pub fn rearrange(&self, xs: &Tensor) -> Result<Tensor> {
58+
let xs = xs.transpose(1, 2)?;
59+
xs.unsqueeze(candle_core::D::Minus1)
60+
}
61+
62+
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
63+
let xs = self.linear.forward(xs)?;
64+
let xs = self.rearrange(&xs)?;
65+
candle_nn::ops::sigmoid(&xs)
66+
}
67+
}
68+
69+
#[derive(Debug)]
70+
pub struct ToVGates {
71+
pub(crate) linear: Linear,
72+
#[allow(dead_code)]
73+
pub(crate) heads: usize,
74+
}
75+
76+
impl ToVGates {
77+
pub fn new(vb: VarBuilder, dim: usize, heads: usize) -> Result<Self> {
78+
Ok(Self {
79+
linear: linear_no_bias(dim, heads, vb)?,
80+
heads,
81+
})
82+
}
83+
84+
pub fn rearrange(&self, xs: &Tensor) -> Result<Tensor> {
85+
let xs = xs.transpose(1, 2)?;
86+
xs.unsqueeze(candle_core::D::Minus1)
87+
}
88+
89+
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
90+
let xs = self.linear.forward(xs)?;
91+
let xs = candle_nn::ops::sigmoid(&xs)?;
92+
self.rearrange(&xs)
93+
}
94+
}
95+
96+
#[derive(Debug)]
97+
pub struct ToOut {
98+
pub(crate) drop_p: f64,
99+
pub(crate) linear: Linear,
100+
}
101+
102+
impl ToOut {
103+
pub fn new(
104+
vb: VarBuilder,
105+
dim: usize,
106+
heads: usize,
107+
dim_head: usize,
108+
drop_p: Option<f64>,
109+
) -> Result<Self> {
110+
Ok(Self {
111+
drop_p: drop_p.unwrap_or(0.0),
112+
linear: linear_no_bias(dim_head * heads, dim, vb)?,
113+
})
114+
}
115+
116+
pub fn rearrange(&self, xs: &Tensor) -> Result<Tensor> {
117+
let xs_dims = xs.dims();
118+
let (b, h, n, d) = (xs_dims[0], xs_dims[1], xs_dims[2], xs_dims[3]);
119+
xs.permute((0, 2, 1, 3))?.reshape((b, n, h * d))
120+
}
121+
122+
pub fn forward_t(&self, xs: &Tensor, train: bool) -> Result<Tensor> {
123+
let xs = self.rearrange(xs)?;
124+
let xs = self.linear.forward(&xs)?;
125+
if train && self.drop_p > 0.0 {
126+
candle_nn::ops::dropout(&xs, self.drop_p as f32)
127+
} else {
128+
Ok(xs)
129+
}
130+
}
131+
}
132+
133+
#[derive(Debug)]
134+
pub struct Attention {
135+
#[allow(dead_code)]
136+
scale: f64,
137+
#[allow(dead_code)]
138+
drop_p: f64,
139+
norm: LayerNorm,
140+
pub(crate) to_qkv: ToQKV,
141+
pub(crate) to_value_residual_mix: Option<ToValueResidualMix>,
142+
pub(crate) to_v_gates: ToVGates,
143+
attend: Attend,
144+
to_out: ToOut,
145+
#[allow(dead_code)]
146+
learned_value_residual_mix: bool,
147+
}
148+
149+
impl Attention {
150+
pub fn new(
151+
vb: VarBuilder,
152+
dim: usize,
153+
dim_head: Option<usize>,
154+
heads: Option<usize>,
155+
drop_p: Option<f64>,
156+
is_flash: Option<bool>,
157+
learned_value_residual_mix: Option<bool>,
158+
) -> Result<Self> {
159+
let dim_head = dim_head.unwrap_or(32);
160+
let heads = heads.unwrap_or(4);
161+
let scale = (dim_head as f64).sqrt();
162+
163+
let norm = layer_norm(dim, LayerNormConfig::default(), vb.pp("norm"))?;
164+
let to_qkv = ToQKV::new(vb.pp("to_qkv"), dim, dim_head * heads, heads)?;
165+
let to_value_residual_mix = if learned_value_residual_mix.unwrap_or(false) {
166+
Some(ToValueResidualMix::new(
167+
vb.pp("to_value_residual_mix"),
168+
dim,
169+
heads,
170+
)?)
171+
} else {
172+
None
173+
};
174+
let to_v_gates = ToVGates::new(vb.pp("to_v_gates"), dim, heads)?;
175+
let to_out = ToOut::new(vb.pp("to_out"), dim, heads, dim_head, drop_p)?;
176+
177+
Ok(Self {
178+
scale,
179+
drop_p: drop_p.unwrap_or(0.0),
180+
norm,
181+
to_qkv,
182+
to_value_residual_mix,
183+
to_v_gates,
184+
attend: Attend::new(drop_p, None, None, is_flash, None),
185+
to_out,
186+
learned_value_residual_mix: learned_value_residual_mix.unwrap_or(false),
187+
})
188+
}
189+
190+
pub fn forward_t(
191+
&self,
192+
xs: &Tensor,
193+
value_residual: Option<&Tensor>,
194+
train: bool,
195+
) -> Result<(Tensor, Tensor)> {
196+
let xs = self.norm.forward(xs)?;
197+
let (q, k, mut v) = self.to_qkv.forward(&xs)?;
198+
let cache_v = v.clone();
199+
200+
if let Some(ref to_value_residual_mix) = self.to_value_residual_mix {
201+
if let Some(value_residual) = value_residual {
202+
let mix = to_value_residual_mix.forward(&xs)?;
203+
let diff = value_residual.sub(&v)?;
204+
let mix = mix.broadcast_as(diff.dims())?;
205+
let weighted = diff.mul(&mix)?;
206+
v = v.add(&weighted)?;
207+
}
208+
}
209+
210+
let out = self.attend.forward_t(&q, &k, &v, train)?;
211+
let gates = self.to_v_gates.forward(&xs)?;
212+
let gates = gates.broadcast_as(out.dims())?;
213+
let out = out.mul(&gates)?;
214+
let out = self.to_out.forward_t(&out, train)?;
215+
216+
Ok((out, cache_v))
217+
}
218+
}

‎src/feedforward.rs‎

Lines changed: 60 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,60 @@
1+
use candle_core::{IndexOp, Result, Tensor};
2+
use candle_nn::{layer_norm, linear, LayerNorm, LayerNormConfig, Linear, Module, VarBuilder};
3+
4+
#[derive(Debug, Clone)]
5+
pub struct GEGLU;
6+
7+
impl GEGLU {
8+
pub fn rearrange(&self, xs: &Tensor) -> Result<(Tensor, Tensor)> {
9+
let xs_dims = xs.dims();
10+
let (b, n, d) = (xs_dims[0], xs_dims[1], xs_dims[2]);
11+
let reshaped = xs.reshape((b, n, 2, d / 2))?;
12+
let x = reshaped.i((.., .., 0, ..))?;
13+
let gate = reshaped.i((.., .., 1, ..))?;
14+
Ok((x, gate))
15+
}
16+
17+
pub fn forward(&self, xs: &Tensor) -> Result<Tensor> {
18+
let (x, gate) = self.rearrange(xs)?;
19+
let gate = gate.gelu()?;
20+
x.mul(&gate)
21+
}
22+
}
23+
24+
#[derive(Debug)]
25+
pub struct FeedForward {
26+
drop_p: f64,
27+
norm: LayerNorm,
28+
linear1: Linear,
29+
geglu: GEGLU,
30+
linear2: Linear,
31+
}
32+
33+
impl FeedForward {
34+
pub fn new(vb: VarBuilder, dim: usize, mult: usize, drop_p: Option<f64>) -> Result<Self> {
35+
let hidden_size = ((dim * mult * 2) as f64 / 3.0).trunc() as usize;
36+
let norm = layer_norm(dim, LayerNormConfig::default(), vb.pp("norm"))?;
37+
let linear1 = linear(dim, hidden_size * 2, vb.pp("linear1"))?;
38+
let linear2 = linear(hidden_size, dim, vb.pp("linear2"))?;
39+
40+
Ok(Self {
41+
drop_p: drop_p.unwrap_or(0.0),
42+
norm,
43+
linear1,
44+
geglu: GEGLU,
45+
linear2,
46+
})
47+
}
48+
49+
pub fn forward_t(&self, xs: &Tensor, train: bool) -> Result<Tensor> {
50+
let xs = self.norm.forward(xs)?;
51+
let xs = self.linear1.forward(&xs)?;
52+
let xs = self.geglu.forward(&xs)?;
53+
let xs = if train && self.drop_p > 0.0 {
54+
candle_nn::ops::dropout(&xs, self.drop_p as f32)?
55+
} else {
56+
xs
57+
};
58+
self.linear2.forward(&xs)
59+
}
60+
}

0 commit comments

Comments
 (0)