pub unsafe fn accumulate_head_fallback(dest: &mut [f32], src: &[f32]) {
let len = core::cmp::min(dest.len(), src.len());
for i in 0..len {
unsafe {
let acc = *dest.get_unchecked_mut(i) as f64 + *src.get_unchecked(i) as f64;
*dest.get_unchecked_mut(i) = acc as f32;
}
}
}
pub unsafe fn tanh_and_accumulate_block_fallback(head_input: &mut [f32], block: &mut [f32]) {
let len = head_input.len();
for i in 0..len {
let v = block[i];
let activated = v.tanh(); block[i] = activated; let acc = head_input[i] as f64 + activated as f64;
head_input[i] = acc as f32; }
}
pub unsafe fn gated_activation_and_accumulate_block_fallback(
head_input: &mut [f32],
block: &mut [f32],
ch: usize, ) {
let num_frames = head_input.len() / ch;
for f in 0..num_frames {
let block_offset = f * 2 * ch;
let head_offset = f * ch;
for c in 0..ch {
let z1 = block[block_offset + c];
let z2 = block[block_offset + ch + c];
let activated = z1.tanh() * (1.0 / (1.0 + (-z2).exp()));
block[block_offset + c] = activated;
let acc = head_input[head_offset + c] as f64 + activated as f64;
head_input[head_offset + c] = acc as f32;
}
}
}
pub unsafe fn tanh_and_overwrite_block_fallback(head_input: &mut [f32], block: &mut [f32]) {
let len = head_input.len();
for i in 0..len {
let v = block[i];
let activated = v.tanh();
block[i] = activated;
head_input[i] = activated;
}
}
pub unsafe fn tanh_and_accumulate_with_seed_fallback(
head_input: &mut [f32],
block: &mut [f32],
seed: &[f32],
) {
let len = head_input.len();
for i in 0..len {
let v = block[i];
let activated = v.tanh();
block[i] = activated;
let acc = seed[i] as f64 + activated as f64;
head_input[i] = acc as f32;
}
}
pub unsafe fn gated_activation_and_overwrite_block_fallback(
head_input: &mut [f32],
block: &mut [f32],
ch: usize,
) {
let num_frames = head_input.len() / ch;
for f in 0..num_frames {
let block_offset = f * 2 * ch;
let head_offset = f * ch;
for c in 0..ch {
let z1 = block[block_offset + c];
let z2 = block[block_offset + ch + c];
let activated = z1.tanh() * (1.0 / (1.0 + (-z2).exp()));
block[block_offset + c] = activated;
head_input[head_offset + c] = activated;
}
}
}