#![allow(unsafe_op_in_unsafe_fn, clippy::missing_safety_doc)]
use crate::common::diagnostics::NamErrorCode;
use crate::math::common::AlignedVec;
use core::arch::x86_64::*;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FiLMConfig {
pub active: bool,
pub shift: bool,
pub groups: u32,
}
impl Default for FiLMConfig {
fn default() -> Self {
Self {
active: false,
shift: true,
groups: 1,
}
}
}
#[derive(Clone)]
#[repr(align(64))]
pub struct FiLMLayer {
pub config: FiLMConfig,
pub cond_size: usize,
pub channels: usize,
pub weights: AlignedVec<f32>,
pub bias: AlignedVec<f32>,
scale_shift_buf: AlignedVec<f32>,
}
impl FiLMLayer {
pub fn load(
config: FiLMConfig,
cond_size: usize,
channels: usize,
weights: Vec<f32>,
bias: Vec<f32>,
) -> Result<Self, NamErrorCode> {
let expected_bias = if config.shift { channels * 2 } else { channels };
let mut bias_padded = bias;
if bias_padded.len() < expected_bias {
bias_padded.resize(expected_bias, 0.0);
}
Ok(Self {
config,
cond_size,
channels,
weights: AlignedVec::from_vec(weights)?,
bias: AlignedVec::from_vec(bias_padded)?,
scale_shift_buf: AlignedVec::new(channels * 2, 0.0f32)?,
})
}
#[inline(always)]
pub unsafe fn process(&mut self, input: &mut [f32], condition: &[f32]) {
debug_assert!(
condition.len() >= self.cond_size,
"FiLM process: condition slice length ({}) must be >= cond_size ({})",
condition.len(),
self.cond_size
);
self.cond_to_scale_shift(condition);
self.apply_modulation(input);
}
#[inline(always)]
unsafe fn cond_to_scale_shift(&mut self, condition: &[f32]) {
debug_assert!(
condition.len() >= self.cond_size,
"FiLM cond_to_scale_shift: condition slice length ({}) must be >= cond_size ({})",
condition.len(),
self.cond_size
);
let g = self.config.groups as usize;
let ch_per_group = self.channels / g;
let cond_per_group = self.cond_size / g;
let out_per_group = if self.config.shift {
ch_per_group * 2
} else {
ch_per_group
};
for grp in 0..g {
let cond_slice =
condition.get_unchecked(grp * cond_per_group..(grp + 1) * cond_per_group);
let w_offset = grp * out_per_group * cond_per_group;
for row in 0..out_per_group {
let global_out = if row < ch_per_group {
grp * ch_per_group + row
} else {
self.channels + grp * ch_per_group + (row - ch_per_group)
};
let mut sum = *self.bias.get_unchecked(global_out);
let w_start = w_offset + row * cond_per_group;
let w_row = self
.weights
.get_unchecked(w_start..w_start + cond_per_group);
sum += dot_product_avx2(w_row, cond_slice);
*self.scale_shift_buf.get_unchecked_mut(global_out) = sum;
}
}
if !self.config.shift {
for c in self.channels..self.channels * 2 {
*self.scale_shift_buf.get_unchecked_mut(c) = 0.0;
}
}
}
#[inline(always)]
unsafe fn apply_modulation(&mut self, input: &mut [f32]) {
let scale = &self.scale_shift_buf[..self.channels];
let shift = &self.scale_shift_buf[self.channels..self.channels * 2];
let limit = input.len().min(self.channels);
let (in_chunks, in_rem) = input[..limit].as_chunks_mut::<8>();
let (scale_chunks, _) = scale[..limit].as_chunks::<8>();
let (shift_chunks, _) = shift[..limit].as_chunks::<8>();
for i in 0..in_chunks.len() {
let v_in = _mm256_loadu_ps(in_chunks[i].as_ptr());
let v_scale = _mm256_loadu_ps(scale_chunks[i].as_ptr());
let v_shift = _mm256_loadu_ps(shift_chunks[i].as_ptr());
_mm256_storeu_ps(
in_chunks[i].as_mut_ptr(),
_mm256_fmadd_ps(v_in, v_scale, v_shift),
);
}
let off = in_chunks.len() * 8;
for c in 0..in_rem.len() {
input[off + c] = input[off + c] * scale[off + c] + shift[off + c];
}
}
}
pub struct FilmBlock<'a> {
pub conv_pre_film: Option<&'a mut FiLMLayer>,
pub conv_post_film: Option<&'a mut FiLMLayer>,
pub input_mixin_pre_film: Option<&'a mut FiLMLayer>,
pub input_mixin_post_film: Option<&'a mut FiLMLayer>,
pub activation_pre_film: Option<&'a mut FiLMLayer>,
pub activation_post_film: Option<&'a mut FiLMLayer>,
pub layer1x1_post_film: Option<&'a mut FiLMLayer>,
pub head1x1_post_film: Option<&'a mut FiLMLayer>,
}
impl<'a> FilmBlock<'a> {
pub fn empty() -> Self {
Self {
conv_pre_film: None,
conv_post_film: None,
input_mixin_pre_film: None,
input_mixin_post_film: None,
activation_pre_film: None,
activation_post_film: None,
layer1x1_post_film: None,
head1x1_post_film: None,
}
}
}
impl super::layer::A2Layer {
#[inline]
pub fn film_block(&mut self) -> FilmBlock<'_> {
FilmBlock {
conv_pre_film: self.conv_pre_film.as_mut(),
conv_post_film: self.conv_post_film.as_mut(),
input_mixin_pre_film: self.input_mixin_pre_film.as_mut(),
input_mixin_post_film: self.input_mixin_post_film.as_mut(),
activation_pre_film: self.activation_pre_film.as_mut(),
activation_post_film: self.activation_post_film.as_mut(),
layer1x1_post_film: self.layer1x1_post_film.as_mut(),
head1x1_post_film: self.head1x1_post_film.as_mut(),
}
}
}
#[inline(always)]
unsafe fn dot_product_avx2(a: &[f32], b: &[f32]) -> f32 {
let len = a.len();
let mut sum0 = _mm256_setzero_ps();
let mut sum1 = _mm256_setzero_ps();
let mut i = 0;
while i + 16 <= len {
let ha0 = _mm256_loadu_ps(a.as_ptr().add(i));
let hb0 = _mm256_loadu_ps(b.as_ptr().add(i));
sum0 = _mm256_fmadd_ps(ha0, hb0, sum0);
let ha1 = _mm256_loadu_ps(a.as_ptr().add(i + 8));
let hb1 = _mm256_loadu_ps(b.as_ptr().add(i + 8));
sum1 = _mm256_fmadd_ps(ha1, hb1, sum1);
i += 16;
}
while i + 8 <= len {
let ha = _mm256_loadu_ps(a.as_ptr().add(i));
let hb = _mm256_loadu_ps(b.as_ptr().add(i));
sum0 = _mm256_fmadd_ps(ha, hb, sum0);
i += 8;
}
let sum = _mm256_add_ps(sum0, sum1);
let hi128 = _mm256_extractf128_ps(sum, 1);
let lo128 = _mm256_castps256_ps128(sum);
let s128 = _mm_add_ps(lo128, hi128);
let shuf = _mm_movehdup_ps(s128);
let sums = _mm_add_ps(s128, shuf);
let shuf2 = _mm_movehl_ps(sums, sums);
let r = _mm_add_ss(sums, shuf2);
let mut out = _mm_cvtss_f32(r);
while i < len {
out += *a.get_unchecked(i) * *b.get_unchecked(i);
i += 1;
}
out
}
#[cfg(test)]
#[path = "film_test.rs"]
mod tests;