1use hanzo_ml::{DType, Error, Module, Result, Tensor, D};
32
33#[derive(Debug, Clone, Copy, PartialEq)]
34pub struct LayerNormConfig {
35 pub eps: f64,
36 pub remove_mean: bool,
39 pub affine: bool,
40}
41
42impl Default for LayerNormConfig {
43 fn default() -> Self {
44 Self {
45 eps: 1e-5,
46 remove_mean: true,
47 affine: true,
48 }
49 }
50}
51
52impl From<f64> for LayerNormConfig {
53 fn from(eps: f64) -> Self {
54 Self {
55 eps,
56 remove_mean: true,
57 affine: true,
58 }
59 }
60}
61
62#[derive(Clone, Debug)]
64pub struct LayerNorm {
65 weight: Tensor,
66 bias: Option<Tensor>,
67 remove_mean: bool,
68 eps: f64,
69}
70
71impl LayerNorm {
72 pub fn new(weight: Tensor, bias: Tensor, eps: f64) -> Self {
73 Self {
74 weight,
75 bias: Some(bias),
76 remove_mean: true,
77 eps,
78 }
79 }
80
81 pub fn new_no_bias(weight: Tensor, eps: f64) -> Self {
82 Self {
83 weight,
84 bias: None,
85 remove_mean: true,
86 eps,
87 }
88 }
89
90 pub fn rms_norm(weight: Tensor, eps: f64) -> Self {
91 Self {
92 weight,
93 bias: None,
94 remove_mean: false,
95 eps,
96 }
97 }
98
99 pub fn weight(&self) -> &Tensor {
100 &self.weight
101 }
102
103 pub fn bias(&self) -> Option<&Tensor> {
104 self.bias.as_ref()
105 }
106
107 pub fn eps(&self) -> f64 {
108 self.eps
109 }
110
111 pub fn remove_mean(&self) -> bool {
112 self.remove_mean
113 }
114}
115
116impl Module for LayerNorm {
117 fn forward(&self, x: &Tensor) -> Result<Tensor> {
118 if x.is_contiguous() && self.remove_mean {
119 if let Some(bias) = self.bias.as_ref() {
120 return crate::ops::layer_norm(x, &self.weight, bias, self.eps as f32);
121 }
122 }
123 let x_dtype = x.dtype();
124 let internal_dtype = match x_dtype {
125 DType::F16 | DType::BF16 => DType::F32,
126 d => d,
127 };
128 let hidden_size = x.dim(D::Minus1)?;
129 let x = x.to_dtype(internal_dtype)?;
130 let x = if self.remove_mean {
131 let mean_x = (x.sum_keepdim(D::Minus1)? / hidden_size as f64)?;
132 x.broadcast_sub(&mean_x)?
133 } else {
134 x
135 };
136 let norm_x = (x.sqr()?.sum_keepdim(D::Minus1)? / hidden_size as f64)?;
137 let x_normed = x.broadcast_div(&(norm_x + self.eps)?.sqrt()?)?;
138 let x = x_normed.to_dtype(x_dtype)?.broadcast_mul(&self.weight)?;
139 match &self.bias {
140 None => Ok(x),
141 Some(bias) => x.broadcast_add(bias),
142 }
143 }
144}
145
146pub fn layer_norm<C: Into<LayerNormConfig>>(
147 size: usize,
148 config: C,
149 vb: crate::VarBuilder,
150) -> Result<LayerNorm> {
151 let config = config.into();
152
153 let weight_tensor_name = ["weight", "gamma"]
157 .iter()
158 .find(|&name| vb.contains_tensor(name))
159 .ok_or_else(|| Error::Msg("Failed to find weight tensor".into()))?;
160
161 let weight = vb.get_with_hints(size, weight_tensor_name, crate::Init::Const(1.))?;
162
163 let bias_tensor_name = ["bias", "beta"]
164 .iter()
165 .find(|&name| vb.contains_tensor(name))
166 .ok_or_else(|| Error::Msg("Failed to find weight tensor".into()))?;
167
168 let bias = if config.affine {
169 Some(vb.get_with_hints(size, bias_tensor_name, crate::Init::Const(0.))?)
170 } else {
171 None
172 };
173
174 Ok(LayerNorm {
175 weight,
176 bias,
177 remove_mean: config.remove_mean,
178 eps: config.eps,
179 })
180}
181
182pub fn layer_norm_no_bias(size: usize, eps: f64, vb: crate::VarBuilder) -> Result<LayerNorm> {
183 let config = LayerNormConfig {
184 eps,
185 remove_mean: true,
186 affine: false,
187 };
188 layer_norm(size, config, vb)
189}
190
191#[derive(Clone, Debug)]
193pub struct RmsNorm(LayerNorm);
194
195impl RmsNorm {
196 pub fn new(weight: Tensor, eps: f64) -> Self {
197 Self(LayerNorm::rms_norm(weight, eps))
198 }
199
200 pub fn into_inner(self) -> LayerNorm {
201 self.0
202 }
203
204 pub fn weight(&self) -> &Tensor {
205 self.0.weight()
206 }
207
208 pub fn eps(&self) -> f64 {
209 self.0.eps()
210 }
211
212 pub fn forward_diff(&self, xs: &Tensor) -> Result<Tensor> {
214 self.0.forward(xs)
215 }
216}
217
218impl Module for RmsNorm {
219 fn forward(&self, xs: &Tensor) -> Result<Tensor> {
220 if xs.is_contiguous() {
221 crate::ops::rms_norm(xs, &self.0.weight, self.0.eps as f32)
222 } else {
223 self.0.forward(xs)
224 }
225 }
226}
227
228pub fn rms_norm(size: usize, eps: f64, vb: crate::VarBuilder) -> Result<RmsNorm> {
229 let config = LayerNormConfig {
230 eps,
231 remove_mean: false,
232 affine: false,
233 };
234 Ok(RmsNorm(layer_norm(size, config, vb)?))
235}