MCPcopy Create free account
hub / github.com/evilsocket/cake / make_pos_emb

Method make_pos_emb

cake-core/src/models/luxtts/text_encoder.rs:97–121  ·  view source on GitHub ↗

Create CompactRelPositionalEncoding [1, 2*seq-1, pos_dim].

(&self, seq_len: usize, device: &Device, dtype: DType)

Source from the content-addressed store, hash-verified

95
96 /// Create CompactRelPositionalEncoding [1, 2*seq-1, pos_dim].
97 fn make_pos_emb(&self, seq_len: usize, device: &Device, dtype: DType) -> Result<Tensor> {
98 let pos_len = 2 * seq_len - 1;
99 let half_dim = self.pos_dim / 2;
100 let compression_length = (self.pos_dim as f32).sqrt();
101 let length_scale = 1.0 * self.pos_dim as f32 / (2.0 * std::f32::consts::PI);
102
103 let mut pos_data = vec![0.0f32; pos_len * self.pos_dim];
104 for pos in 0..pos_len {
105 let t = pos as f32 - (seq_len as f32 - 1.0);
106 let x_compressed = compression_length
107 * t.signum()
108 * ((t.abs() + compression_length).ln() - compression_length.ln());
109 let x_atan = (x_compressed / length_scale).atan();
110
111 for i in 0..half_dim {
112 let freq = (i + 1) as f32;
113 pos_data[pos * self.pos_dim + 2 * i] = (x_atan * freq).cos();
114 pos_data[pos * self.pos_dim + 2 * i + 1] = (x_atan * freq).sin();
115 }
116 pos_data[pos * self.pos_dim + self.pos_dim - 1] = 1.0;
117 }
118
119 let pos_emb = Tensor::from_vec(pos_data, (1, pos_len, self.pos_dim), device)?;
120 Ok(pos_emb.to_dtype(dtype)?)
121 }
122}

Callers 1

forwardMethod · 0.45

Calls

no outgoing calls

Tested by

no test coverage detected