diff --git a/src/typeck/intrinsics.rs b/src/typeck/intrinsics.rs index adac94c..594f7e1 100644 --- a/src/typeck/intrinsics.rs +++ b/src/typeck/intrinsics.rs @@ -374,176 +374,9 @@ impl TypeChecker { } } - /// Type-check `exp_poly_f32(v: f32xN) -> f32xN`. Restricts to f32-element - /// vectors only — scalar f32, f64xN, f16xN, integer vectors all rejected - /// at typeck. Codegen retains a defense-in-depth guard. - fn check_exp_poly_f32( - &self, - args: &[Expr], - locals: &HashMap, - span: &Span, - ) -> crate::error::Result { - if args.len() != 1 { - return Err(CompileError::type_error( - format!("exp_poly_f32 expects 1 argument, got {}", args.len()), - span.clone(), - )); - } - let t = self.check_expr(&args[0], locals)?; - match &t { - Type::Vector { elem, .. } if **elem == Type::F32 => Ok(t), - Type::Vector { elem, .. } => Err(CompileError::type_error( - format!("exp_poly_f32 expects f32 element type, got {elem}"), - span.clone(), - )), - Type::F32 | Type::FloatLiteral => Err(CompileError::type_error( - "exp_poly_f32 expects f32 vector, got scalar; use exp() for scalar libm-precision" - .to_string(), - span.clone(), - )), - _ => Err(CompileError::type_error( - format!("exp_poly_f32 expects float vector, got {t}"), - span.clone(), - )), - } - } - - /// Type-check `tanh_approx_f32(v: f32xN) -> f32xN`. Same f32-vector-only - /// shape as `exp_poly_f32`; scalar / f64 / f16 / integer all rejected. - fn check_tanh_approx_f32( - &self, - args: &[Expr], - locals: &HashMap, - span: &Span, - ) -> crate::error::Result { - if args.len() != 1 { - return Err(CompileError::type_error( - format!("tanh_approx_f32 expects 1 argument, got {}", args.len()), - span.clone(), - )); - } - let t = self.check_expr(&args[0], locals)?; - match &t { - Type::Vector { elem, .. } if **elem == Type::F32 => Ok(t), - Type::Vector { elem, .. } => Err(CompileError::type_error( - format!("tanh_approx_f32 expects f32 element type, got {elem}"), - span.clone(), - )), - Type::F32 | Type::FloatLiteral => Err(CompileError::type_error( - "tanh_approx_f32 expects f32 vector, got scalar; use tanh() for scalar libm-precision" - .to_string(), - span.clone(), - )), - _ => Err(CompileError::type_error( - format!("tanh_approx_f32 expects float vector, got {t}"), - span.clone(), - )), - } - } - - /// Type-check `log_approx_f32(v: f32xN) -> f32xN`. Same f32-vector-only - /// shape as `exp_poly_f32`; scalar / f64 / f16 / integer all rejected. - fn check_log_approx_f32( - &self, - args: &[Expr], - locals: &HashMap, - span: &Span, - ) -> crate::error::Result { - if args.len() != 1 { - return Err(CompileError::type_error( - format!("log_approx_f32 expects 1 argument, got {}", args.len()), - span.clone(), - )); - } - let t = self.check_expr(&args[0], locals)?; - match &t { - Type::Vector { elem, .. } if **elem == Type::F32 => Ok(t), - Type::Vector { elem, .. } => Err(CompileError::type_error( - format!("log_approx_f32 expects f32 element type, got {elem}"), - span.clone(), - )), - Type::F32 | Type::FloatLiteral => Err(CompileError::type_error( - "log_approx_f32 expects f32 vector, got scalar; use log() for scalar libm-precision" - .to_string(), - span.clone(), - )), - _ => Err(CompileError::type_error( - format!("log_approx_f32 expects float vector, got {t}"), - span.clone(), - )), - } - } - - /// Type-check `sin_approx_f32(v: f32xN) -> f32xN` and - /// `cos_approx_f32(v: f32xN) -> f32xN`. Same f32-vector-only shape as - /// `exp_poly_f32` — scalar / f64 / f16 / integer all rejected. Error - /// text varies by intrinsic name to point at the right libm fallback. - fn check_sin_cos_approx_f32( - &self, - name: &str, - args: &[Expr], - locals: &HashMap, - span: &Span, - ) -> crate::error::Result { - if args.len() != 1 { - return Err(CompileError::type_error( - format!("{name} expects 1 argument, got {}", args.len()), - span.clone(), - )); - } - let libm_fallback = if name == "cos_approx_f32" { - "cos" - } else { - "sin" - }; - let t = self.check_expr(&args[0], locals)?; - match &t { - Type::Vector { elem, .. } if **elem == Type::F32 => Ok(t), - Type::Vector { elem, .. } => Err(CompileError::type_error( - format!("{name} expects f32 element type, got {elem}"), - span.clone(), - )), - Type::F32 | Type::FloatLiteral => Err(CompileError::type_error( - format!( - "{name} expects f32 vector, got scalar; use {libm_fallback}() for scalar libm-precision" - ), - span.clone(), - )), - _ => Err(CompileError::type_error( - format!("{name} expects float vector, got {t}"), - span.clone(), - )), - } - } - - fn check_prefetch( - &self, - args: &[Expr], - locals: &HashMap, - span: &Span, - ) -> crate::error::Result { - if args.len() != 2 { - return Err(CompileError::type_error( - "prefetch expects 2 arguments: (ptr, offset)", - span.clone(), - )); - } - let ptr_type = self.check_expr(&args[0], locals)?; - if !matches!(ptr_type, Type::Pointer { .. }) { - return Err(CompileError::type_error( - format!("prefetch first argument must be a pointer, got {ptr_type}"), - args[0].span().clone(), - )); - } - let offset_type = self.check_expr(&args[1], locals)?; - if !offset_type.is_integer() { - return Err(CompileError::type_error( - format!("prefetch offset must be integer, got {offset_type}"), - args[1].span().clone(), - )); - } - Ok(Type::Void) - } + // Transcendental f32-vector approximation checks (exp_poly, tanh, log, + // sin/cos) live in `intrinsics_transcendental.rs`. + // Prefetch checks live in `intrinsics_prefetch.rs`. fn check_abs( &self, diff --git a/src/typeck/intrinsics_prefetch.rs b/src/typeck/intrinsics_prefetch.rs new file mode 100644 index 0000000..641cf68 --- /dev/null +++ b/src/typeck/intrinsics_prefetch.rs @@ -0,0 +1,46 @@ +//! Type check for the `prefetch` / `prefetch_write` / `prefetch_nta` family. +//! +//! All three intrinsics share the same `(ptr, integer-offset) -> void` +//! shape. The dispatcher in `intrinsics.rs` routes all three names to this +//! single `check_prefetch` because the rw / locality / cache-type fields +//! are constant-baked into the codegen, not part of the typed signature. + +use std::collections::HashMap; + +use crate::ast::Expr; +use crate::error::CompileError; +use crate::lexer::Span; + +use super::TypeChecker; +use super::types::Type; + +impl TypeChecker { + pub(super) fn check_prefetch( + &self, + args: &[Expr], + locals: &HashMap, + span: &Span, + ) -> crate::error::Result { + if args.len() != 2 { + return Err(CompileError::type_error( + "prefetch expects 2 arguments: (ptr, offset)", + span.clone(), + )); + } + let ptr_type = self.check_expr(&args[0], locals)?; + if !matches!(ptr_type, Type::Pointer { .. }) { + return Err(CompileError::type_error( + format!("prefetch first argument must be a pointer, got {ptr_type}"), + args[0].span().clone(), + )); + } + let offset_type = self.check_expr(&args[1], locals)?; + if !offset_type.is_integer() { + return Err(CompileError::type_error( + format!("prefetch offset must be integer, got {offset_type}"), + args[1].span().clone(), + )); + } + Ok(Type::Void) + } +} diff --git a/src/typeck/intrinsics_transcendental.rs b/src/typeck/intrinsics_transcendental.rs new file mode 100644 index 0000000..0b3df70 --- /dev/null +++ b/src/typeck/intrinsics_transcendental.rs @@ -0,0 +1,164 @@ +//! Type checks for the f32-vector transcendental approximation family: +//! `exp_poly_f32` (v1.11.0), `tanh_approx_f32` / `log_approx_f32` / +//! `sin_approx_f32` / `cos_approx_f32` (v1.14.0). +//! +//! All five share the same shape — accept `f32xN` for any vector width; +//! reject scalar, f64, f16, and integer vectors. Codegen retains a +//! defense-in-depth guard against malformed args. The shared shape means +//! these definitions could collapse into one helper parameterized by +//! intrinsic name and libm fallback, but the per-intrinsic error messages +//! (pointing at the right libm scalar fallback) are stable enough that +//! the deduplication isn't worth the indirection. + +use std::collections::HashMap; + +use crate::ast::Expr; +use crate::error::CompileError; +use crate::lexer::Span; + +use super::TypeChecker; +use super::types::Type; + +impl TypeChecker { + /// Type-check `exp_poly_f32(v: f32xN) -> f32xN`. Restricts to f32-element + /// vectors only — scalar f32, f64xN, f16xN, integer vectors all rejected + /// at typeck. Codegen retains a defense-in-depth guard. + pub(super) fn check_exp_poly_f32( + &self, + args: &[Expr], + locals: &HashMap, + span: &Span, + ) -> crate::error::Result { + if args.len() != 1 { + return Err(CompileError::type_error( + format!("exp_poly_f32 expects 1 argument, got {}", args.len()), + span.clone(), + )); + } + let t = self.check_expr(&args[0], locals)?; + match &t { + Type::Vector { elem, .. } if **elem == Type::F32 => Ok(t), + Type::Vector { elem, .. } => Err(CompileError::type_error( + format!("exp_poly_f32 expects f32 element type, got {elem}"), + span.clone(), + )), + Type::F32 | Type::FloatLiteral => Err(CompileError::type_error( + "exp_poly_f32 expects f32 vector, got scalar; use exp() for scalar libm-precision" + .to_string(), + span.clone(), + )), + _ => Err(CompileError::type_error( + format!("exp_poly_f32 expects float vector, got {t}"), + span.clone(), + )), + } + } + + /// Type-check `tanh_approx_f32(v: f32xN) -> f32xN`. Same f32-vector-only + /// shape as `exp_poly_f32`; scalar / f64 / f16 / integer all rejected. + pub(super) fn check_tanh_approx_f32( + &self, + args: &[Expr], + locals: &HashMap, + span: &Span, + ) -> crate::error::Result { + if args.len() != 1 { + return Err(CompileError::type_error( + format!("tanh_approx_f32 expects 1 argument, got {}", args.len()), + span.clone(), + )); + } + let t = self.check_expr(&args[0], locals)?; + match &t { + Type::Vector { elem, .. } if **elem == Type::F32 => Ok(t), + Type::Vector { elem, .. } => Err(CompileError::type_error( + format!("tanh_approx_f32 expects f32 element type, got {elem}"), + span.clone(), + )), + Type::F32 | Type::FloatLiteral => Err(CompileError::type_error( + "tanh_approx_f32 expects f32 vector, got scalar; use tanh() for scalar libm-precision" + .to_string(), + span.clone(), + )), + _ => Err(CompileError::type_error( + format!("tanh_approx_f32 expects float vector, got {t}"), + span.clone(), + )), + } + } + + /// Type-check `log_approx_f32(v: f32xN) -> f32xN`. Same f32-vector-only + /// shape as `exp_poly_f32`; scalar / f64 / f16 / integer all rejected. + pub(super) fn check_log_approx_f32( + &self, + args: &[Expr], + locals: &HashMap, + span: &Span, + ) -> crate::error::Result { + if args.len() != 1 { + return Err(CompileError::type_error( + format!("log_approx_f32 expects 1 argument, got {}", args.len()), + span.clone(), + )); + } + let t = self.check_expr(&args[0], locals)?; + match &t { + Type::Vector { elem, .. } if **elem == Type::F32 => Ok(t), + Type::Vector { elem, .. } => Err(CompileError::type_error( + format!("log_approx_f32 expects f32 element type, got {elem}"), + span.clone(), + )), + Type::F32 | Type::FloatLiteral => Err(CompileError::type_error( + "log_approx_f32 expects f32 vector, got scalar; use log() for scalar libm-precision" + .to_string(), + span.clone(), + )), + _ => Err(CompileError::type_error( + format!("log_approx_f32 expects float vector, got {t}"), + span.clone(), + )), + } + } + + /// Type-check `sin_approx_f32(v: f32xN) -> f32xN` and + /// `cos_approx_f32(v: f32xN) -> f32xN`. Same f32-vector-only shape as + /// `exp_poly_f32` — scalar / f64 / f16 / integer all rejected. Error + /// text varies by intrinsic name to point at the right libm fallback. + pub(super) fn check_sin_cos_approx_f32( + &self, + name: &str, + args: &[Expr], + locals: &HashMap, + span: &Span, + ) -> crate::error::Result { + if args.len() != 1 { + return Err(CompileError::type_error( + format!("{name} expects 1 argument, got {}", args.len()), + span.clone(), + )); + } + let libm_fallback = if name == "cos_approx_f32" { + "cos" + } else { + "sin" + }; + let t = self.check_expr(&args[0], locals)?; + match &t { + Type::Vector { elem, .. } if **elem == Type::F32 => Ok(t), + Type::Vector { elem, .. } => Err(CompileError::type_error( + format!("{name} expects f32 element type, got {elem}"), + span.clone(), + )), + Type::F32 | Type::FloatLiteral => Err(CompileError::type_error( + format!( + "{name} expects f32 vector, got scalar; use {libm_fallback}() for scalar libm-precision" + ), + span.clone(), + )), + _ => Err(CompileError::type_error( + format!("{name} expects float vector, got {t}"), + span.clone(), + )), + } + } +} diff --git a/src/typeck/mod.rs b/src/typeck/mod.rs index aa8ae95..7b1faf9 100644 --- a/src/typeck/mod.rs +++ b/src/typeck/mod.rs @@ -11,7 +11,9 @@ mod intrinsics_lane; mod intrinsics_memory; mod intrinsics_neon; mod intrinsics_pack; +mod intrinsics_prefetch; mod intrinsics_simd; +mod intrinsics_transcendental; pub mod types; use std::cell::RefCell; diff --git a/tests/common/mod.rs b/tests/common/mod.rs index 4920a66..5bf9a98 100644 --- a/tests/common/mod.rs +++ b/tests/common/mod.rs @@ -146,6 +146,116 @@ pub fn compile_to_ir(source: &str) -> String { ea_compiler::compile_to_ir(source).expect("compilation failed") } +/// Run a transcendental approximation intrinsic over `inputs` lane-by-lane, +/// compare each output to the corresponding libm reference, and assert that +/// every absolute error is within `abs_tol`. Shared by the phase14 +/// transcendental test files (`tanh_approx`, `log_approx`, `sin_cos_approx`). +/// +/// - `vector_type`: `"f32x4"` or `"f32x8"` (sets `lanes` to 4 or 8). +/// - `intrinsic`: Eä intrinsic name, e.g. `"tanh_approx_f32"`. +/// - `libm_ref`: C reference function, e.g. `"tanhf"`. +/// - `padding`: value used to pad `inputs` up to a multiple of `lanes` so the +/// kernel's `while i + lanes <= n` loop covers the original input range. +/// Must satisfy `intrinsic(padding) ≈ libm_ref(padding)` within `abs_tol` +/// (use `0.0` for tanh / sin, `1.0` for log, etc.). +#[allow(dead_code)] +pub fn assert_transcendental_accuracy( + inputs: &[f32], + vector_type: &str, + intrinsic: &str, + libm_ref: &str, + abs_tol: f32, + padding: f32, +) { + use ea_compiler::{CompileOptions, OutputMode}; + + let lanes = if vector_type == "f32x4" { 4 } else { 8 }; + + let ea = format!( + r#" + export func k(input: *f32, output: *mut f32, n: i32) {{ + let mut i: i32 = 0 + while i + {lanes} <= n {{ + let v: {vector_type} = load(input, i) + let r: {vector_type} = {intrinsic}(v) + store(output, i, r) + i = i + {lanes} + }} + }} + "# + ); + + let mut padded = inputs.to_vec(); + while !padded.len().is_multiple_of(lanes) { + padded.push(padding); + } + let n = padded.len(); + let original_n = inputs.len(); + + let inputs_str = padded + .iter() + .map(|f| format!("{f:.10e}f")) + .collect::>() + .join(", "); + + let c = format!( + r#" + #include + #include + extern void k(const float *input, float *output, int n); + int main(void) {{ + float in[{n}] = {{{inputs_str}}}; + float out[{n}] = {{0}}; + k(in, out, {n}); + for (int i = 0; i < {original_n}; ++i) {{ + float ref = {libm_ref}(in[i]); + float got = out[i]; + float abs_err = fabsf(got - ref); + if (abs_err > {abs_tol}f) {{ + printf("FAIL i=%d in=%g got=%g ref=%g abs=%g\n", i, in[i], got, ref, abs_err); + return 1; + }} + }} + printf("OK\n"); + return 0; + }} + "# + ); + + let dir = TempDir::new().unwrap(); + let obj = dir.path().join("k.o"); + let cpath = dir.path().join("h.c"); + let bin = dir.path().join("k_bin"); + let opts = CompileOptions { + opt_level: 3, + target_cpu: None, + extra_features: String::new(), + target_triple: None, + }; + ea_compiler::compile_with_options(&ea, &obj, OutputMode::ObjectFile, &opts) + .expect("compile failed"); + std::fs::write(&cpath, c).expect("write c"); + let status = Command::new("cc") + .args([ + cpath.to_str().unwrap(), + obj.to_str().unwrap(), + "-o", + bin.to_str().unwrap(), + "-lm", + ]) + .status() + .expect("link failed"); + assert!(status.success(), "linker failed"); + let out = Command::new(&bin).output().expect("run failed"); + let stdout = String::from_utf8_lossy(&out.stdout).replace("\r\n", "\n"); + assert_eq!( + stdout.trim(), + "OK", + "stderr: {}", + String::from_utf8_lossy(&out.stderr) + ); +} + /// Asserts that at least one of `expected_mnemonics` appears in the disassembly /// of the object file produced by compiling `ea_source`. /// diff --git a/tests/phase14_log_approx.rs b/tests/phase14_log_approx.rs index 9a11ecb..03a5364 100644 --- a/tests/phase14_log_approx.rs +++ b/tests/phase14_log_approx.rs @@ -1,5 +1,9 @@ +#[cfg(feature = "llvm")] +mod common; + #[cfg(feature = "llvm")] mod tests { + use super::common::assert_transcendental_accuracy; use ea_compiler::{CompileOptions, OutputMode}; use std::process::Command; use tempfile::TempDir; @@ -56,98 +60,18 @@ mod tests { assert_eq!(stdout.trim(), "0\n0\n0\n0"); } - /// Compile a kernel that runs log_approx_f32 over `inputs` lane-by-lane, - /// link with C harness that calls logf, assert absolute error ≤ 3e-6. - /// - /// Absolute (not relative) error because log(x) → 0 as x → 1, making - /// relative error blow up near 1. Across the rest of the input range - /// the magnitude of log(x) is bounded modestly, so absolute is the - /// natural metric. + /// Helper wrapping `common::assert_transcendental_accuracy` with the + /// log-specific parameters. Absolute (not relative) error tolerance + /// because log(x) → 0 as x → 1, making relative error blow up near 1. fn accuracy_test_impl(inputs: &[f32], vector_type: &str) { - let lanes = if vector_type == "f32x4" { 4 } else { 8 }; - - let ea = format!( - r#" - export func k(input: *f32, output: *mut f32, n: i32) {{ - let mut i: i32 = 0 - while i + {lanes} <= n {{ - let v: {vector_type} = load(input, i) - let r: {vector_type} = log_approx_f32(v) - store(output, i, r) - i = i + {lanes} - }} - }} - "# - ); - - let mut padded = inputs.to_vec(); - while !padded.len().is_multiple_of(lanes) { - padded.push(1.0); // pad with 1.0 so log_approx_f32 returns 0 on padding - } - let n = padded.len(); - let original_n = inputs.len(); - - let inputs_str = padded - .iter() - .map(|f| format!("{f:.10e}f")) - .collect::>() - .join(", "); - - let c = format!( - r#" - #include - #include - extern void k(const float *input, float *output, int n); - int main(void) {{ - float in[{n}] = {{{inputs_str}}}; - float out[{n}] = {{0}}; - k(in, out, {n}); - for (int i = 0; i < {original_n}; ++i) {{ - float ref = logf(in[i]); - float got = out[i]; - float abs_err = fabsf(got - ref); - if (abs_err > 3.0e-6f) {{ - printf("FAIL i=%d in=%g got=%g ref=%g abs=%g\n", i, in[i], got, ref, abs_err); - return 1; - }} - }} - printf("OK\n"); - return 0; - }} - "# - ); - - let dir = TempDir::new().unwrap(); - let obj = dir.path().join("k.o"); - let cpath = dir.path().join("h.c"); - let bin = dir.path().join("k_bin"); - let opts = CompileOptions { - opt_level: 3, - target_cpu: None, - extra_features: String::new(), - target_triple: None, - }; - ea_compiler::compile_with_options(&ea, &obj, OutputMode::ObjectFile, &opts) - .expect("compile failed"); - std::fs::write(&cpath, c).expect("write c"); - let status = Command::new("cc") - .args([ - cpath.to_str().unwrap(), - obj.to_str().unwrap(), - "-o", - bin.to_str().unwrap(), - "-lm", - ]) - .status() - .expect("link failed"); - assert!(status.success(), "linker failed"); - let out = Command::new(&bin).output().expect("run failed"); - let stdout = String::from_utf8_lossy(&out.stdout).replace("\r\n", "\n"); - assert_eq!( - stdout.trim(), - "OK", - "stderr: {}", - String::from_utf8_lossy(&out.stderr) + // padding = 1.0 because log(1) = 0 — safe to pad input arrays with. + assert_transcendental_accuracy( + inputs, + vector_type, + "log_approx_f32", + "logf", + 3.0e-6, + 1.0, ); } diff --git a/tests/phase14_sin_cos_approx.rs b/tests/phase14_sin_cos_approx.rs index 4cf11c5..492c148 100644 --- a/tests/phase14_sin_cos_approx.rs +++ b/tests/phase14_sin_cos_approx.rs @@ -1,5 +1,9 @@ +#[cfg(feature = "llvm")] +mod common; + #[cfg(feature = "llvm")] mod tests { + use super::common::assert_transcendental_accuracy; use ea_compiler::{CompileOptions, OutputMode}; use std::process::Command; use tempfile::TempDir; @@ -74,97 +78,14 @@ mod tests { assert_eq!(stdout.trim(), expected); } - /// Compile a kernel that runs `intrinsic` over `inputs` lane-by-lane, - /// link with C harness that calls `ref_fn`, assert absolute error ≤ 3e-6. - /// - /// Absolute error tolerance because sin/cos pass through zero at the - /// quadrant boundaries — relative error blows up there. + /// Helper wrapping `common::assert_transcendental_accuracy` with the + /// sin/cos shared parameters. Absolute error tolerance because sin/cos + /// pass through zero at the quadrant boundaries — relative error blows + /// up there. fn accuracy_test_impl(inputs: &[f32], vector_type: &str, intrinsic: &str, ref_fn: &str) { - let lanes = if vector_type == "f32x4" { 4 } else { 8 }; - - let ea = format!( - r#" - export func k(input: *f32, output: *mut f32, n: i32) {{ - let mut i: i32 = 0 - while i + {lanes} <= n {{ - let v: {vector_type} = load(input, i) - let r: {vector_type} = {intrinsic}(v) - store(output, i, r) - i = i + {lanes} - }} - }} - "# - ); - - let mut padded = inputs.to_vec(); - while !padded.len().is_multiple_of(lanes) { - padded.push(0.0); - } - let n = padded.len(); - let original_n = inputs.len(); - - let inputs_str = padded - .iter() - .map(|f| format!("{f:.10e}f")) - .collect::>() - .join(", "); - - let c = format!( - r#" - #include - #include - extern void k(const float *input, float *output, int n); - int main(void) {{ - float in[{n}] = {{{inputs_str}}}; - float out[{n}] = {{0}}; - k(in, out, {n}); - for (int i = 0; i < {original_n}; ++i) {{ - float ref = {ref_fn}(in[i]); - float got = out[i]; - float abs_err = fabsf(got - ref); - if (abs_err > 3.0e-6f) {{ - printf("FAIL i=%d in=%g got=%g ref=%g abs=%g\n", i, in[i], got, ref, abs_err); - return 1; - }} - }} - printf("OK\n"); - return 0; - }} - "# - ); - - let dir = TempDir::new().unwrap(); - let obj = dir.path().join("k.o"); - let cpath = dir.path().join("h.c"); - let bin = dir.path().join("k_bin"); - let opts = CompileOptions { - opt_level: 3, - target_cpu: None, - extra_features: String::new(), - target_triple: None, - }; - ea_compiler::compile_with_options(&ea, &obj, OutputMode::ObjectFile, &opts) - .expect("compile failed"); - std::fs::write(&cpath, c).expect("write c"); - let status = Command::new("cc") - .args([ - cpath.to_str().unwrap(), - obj.to_str().unwrap(), - "-o", - bin.to_str().unwrap(), - "-lm", - ]) - .status() - .expect("link failed"); - assert!(status.success(), "linker failed"); - let out = Command::new(&bin).output().expect("run failed"); - let stdout = String::from_utf8_lossy(&out.stdout).replace("\r\n", "\n"); - assert_eq!( - stdout.trim(), - "OK", - "stderr: {}", - String::from_utf8_lossy(&out.stderr) - ); + // padding = 0.0 because sin(0) = 0 and cos(0) = 1 are both within + // tolerance of the libm reference for the padding lane. + assert_transcendental_accuracy(inputs, vector_type, intrinsic, ref_fn, 3.0e-6, 0.0); } fn boundary_points() -> Vec { diff --git a/tests/phase14_tanh_approx.rs b/tests/phase14_tanh_approx.rs index 61ffc80..452b81b 100644 --- a/tests/phase14_tanh_approx.rs +++ b/tests/phase14_tanh_approx.rs @@ -1,5 +1,9 @@ +#[cfg(feature = "llvm")] +mod common; + #[cfg(feature = "llvm")] mod tests { + use super::common::assert_transcendental_accuracy; use ea_compiler::{CompileOptions, OutputMode}; use std::process::Command; use tempfile::TempDir; @@ -55,96 +59,19 @@ mod tests { assert_eq!(stdout.trim(), "0\n0\n0\n0"); } - /// Compile a kernel that runs tanh_approx_f32 over `inputs` lane-by-lane, - /// link with C harness that calls tanhf, assert absolute error ≤ 5e-6. - /// - /// Absolute (not relative) error because tanh is bounded in [-1, 1] and - /// tanh(x)→0 as x→0 makes relative error blow up near the origin. + /// Helper wrapping `common::assert_transcendental_accuracy` with the + /// tanh-specific parameters. Absolute (not relative) error tolerance + /// because tanh is bounded in [-1, 1] and tanh(x)→0 as x→0 makes + /// relative error blow up near the origin. fn accuracy_test_impl(inputs: &[f32], vector_type: &str) { - let lanes = if vector_type == "f32x4" { 4 } else { 8 }; - - let ea = format!( - r#" - export func k(input: *f32, output: *mut f32, n: i32) {{ - let mut i: i32 = 0 - while i + {lanes} <= n {{ - let v: {vector_type} = load(input, i) - let r: {vector_type} = tanh_approx_f32(v) - store(output, i, r) - i = i + {lanes} - }} - }} - "# - ); - - let mut padded = inputs.to_vec(); - while !padded.len().is_multiple_of(lanes) { - padded.push(0.0); - } - let n = padded.len(); - let original_n = inputs.len(); - - let inputs_str = padded - .iter() - .map(|f| format!("{f:.10e}f")) - .collect::>() - .join(", "); - - let c = format!( - r#" - #include - #include - extern void k(const float *input, float *output, int n); - int main(void) {{ - float in[{n}] = {{{inputs_str}}}; - float out[{n}] = {{0}}; - k(in, out, {n}); - for (int i = 0; i < {original_n}; ++i) {{ - float ref = tanhf(in[i]); - float got = out[i]; - float abs_err = fabsf(got - ref); - if (abs_err > 5.0e-6f) {{ - printf("FAIL i=%d in=%g got=%g ref=%g abs=%g\n", i, in[i], got, ref, abs_err); - return 1; - }} - }} - printf("OK\n"); - return 0; - }} - "# - ); - - let dir = TempDir::new().unwrap(); - let obj = dir.path().join("k.o"); - let cpath = dir.path().join("h.c"); - let bin = dir.path().join("k_bin"); - let opts = CompileOptions { - opt_level: 3, - target_cpu: None, - extra_features: String::new(), - target_triple: None, - }; - ea_compiler::compile_with_options(&ea, &obj, OutputMode::ObjectFile, &opts) - .expect("compile failed"); - std::fs::write(&cpath, c).expect("write c"); - let status = Command::new("cc") - .args([ - cpath.to_str().unwrap(), - obj.to_str().unwrap(), - "-o", - bin.to_str().unwrap(), - "-lm", - ]) - .status() - .expect("link failed"); - assert!(status.success(), "linker failed"); - let out = Command::new(&bin).output().expect("run failed"); - let stdout = String::from_utf8_lossy(&out.stdout).replace("\r\n", "\n"); - assert_eq!( - stdout.trim(), - "OK", - "stderr: {}", - String::from_utf8_lossy(&out.stderr) + // padding = 0.0 because tanh(0) = 0 — safe to pad input arrays with. + assert_transcendental_accuracy( + inputs, + vector_type, + "tanh_approx_f32", + "tanhf", + 5.0e-6, + 0.0, ); }