#![cfg(feature = "jit")]
use polydat::JitMode;
use polydat::dsl::compile::compile_polydat_to_assembler;
use polydat::kernel::PolydatKernel;
const REF_NATIVE: bool = true;
fn kernel(src: &str, mode: JitMode) -> PolydatKernel {
let mut asm = compile_polydat_to_assembler(src).unwrap_or_else(|e| panic!("{e}\n{src}"));
asm.set_jit_mode(mode);
asm.compile().unwrap_or_else(|e| panic!("{e}\n{src}"))
}
fn cones(k: &PolydatKernel) -> Vec<String> {
let p = k.program();
(0..p.node_count())
.map(|i| p.node_meta(i).name.clone())
.filter(|n| n.starts_with("jit_cone["))
.collect()
}
fn agree(src: &str, outputs: &[&str], cycles: u64, fused: &[&str]) {
let mut p1 = kernel(src, JitMode::Off);
let mut p3 = kernel(src, JitMode::Force);
let names = cones(&p3);
for member in fused {
assert!(
!REF_NATIVE || names.iter().any(|c| c.contains(member)),
"`{member}` was not fused; cones: {names:?}\n{src}"
);
}
for c in 0..cycles {
p1.set_inputs(&[c]);
p3.set_inputs(&[c]);
for out in outputs {
let a = p1.pull(out).clone();
let b = p3.pull(out).clone();
assert_eq!(a.port_type(), b.port_type(), "{out} at cycle {c}: type");
assert_eq!(
a.to_display_string(),
b.to_display_string(),
"{out} at cycle {c}\n{src}"
);
}
}
}
#[test]
fn scalar_to_string_conversions_run_in_cones() {
agree(
"input cycle: u64\nh := hash(cycle)\ns := __u64_to_string(h)\n",
&["s"],
6,
&["__u64_to_string"],
);
agree(
"input cycle: u64\nf := to_f64(hash(cycle)) / 7.0\ns := __f64_to_string(f)\n",
&["s"],
6,
&["__f64_to_string"],
);
agree(
"input cycle: u64\nb := u64_gt(hash(cycle), 1000)\ns := __bool_to_str(b)\n",
&["s"],
6,
&["__bool_to_str"],
);
}
#[test]
fn string_operations_chain_inside_one_cone() {
let src = "input cycle: u64\ns := str_upper(__u64_to_string(hash(cycle)))\nt := str_lower(s)\n";
agree(src, &["s", "t"], 6, &["str_upper", "str_lower"]);
let src = "input cycle: u64\na := __u64_to_string(hash(cycle))\nb := __u64_to_string(cycle)\nc := str_concat(a, b)\n";
agree(src, &["c"], 6, &["str_concat"]);
}
#[test]
fn a_string_literal_feeds_a_cone_as_a_boundary_handle() {
let src = "input cycle: u64\nlabel := \"row-\"\nn := __u64_to_string(hash(cycle))\nline := str_concat(label, n)\n";
agree(src, &["line"], 6, &["str_concat"]);
let mut p3 = kernel(src, JitMode::Force);
p3.set_inputs(&[2]);
assert!(p3.pull("line").as_str().starts_with("row-"));
}
#[test]
fn a_string_round_trips_through_parse_and_widening() {
let src = "input cycle: u64\ns := __u64_to_string(hash(cycle))\nn := __str_to_u64(s)\nf := __str_to_f64(s)\n";
agree(src, &["n", "f"], 6, &["__str_to_u64", "__str_to_f64"]);
let mut p3 = kernel(src, JitMode::Force);
p3.set_inputs(&[3]);
let mut p1 = kernel(src, JitMode::Off);
p1.set_inputs(&[3]);
assert_eq!(p3.pull("n").as_u64(), p1.pull("n").as_u64());
}
#[test]
fn named_string_lowerings_are_chosen_and_agree() {
use polydat::ast::{PolydatNode, PortType};
use polydat::compile::jit::{JitOp, classify_node_typed};
let u = polydat::library::convert::U64ToString::new();
assert!(matches!(
classify_node_typed(&u, &[PortType::U64]),
JitOp::U64ToStr { .. }
));
let f = polydat::library::convert::F64ToString::new();
assert!(matches!(
classify_node_typed(&f, &[PortType::F64]),
JitOp::F64ToStr { .. }
));
let c = polydat::library::string::StrConcat::new(2);
assert!(matches!(
classify_node_typed(&c, &[PortType::Str, PortType::Str]),
JitOp::StrConcat { .. }
));
assert!(
matches!(
classify_node_typed(&c, &[PortType::Str, PortType::U64]),
JitOp::SlotCall { .. }
),
"a mixed concatenation takes the kit, which reads each wire as typed"
);
let j = polydat::library::json::JsonToStr::new();
assert!(matches!(
classify_node_typed(&j, &[PortType::Json]),
JitOp::JsonToStr { .. }
));
let _ = j.meta();
let src = "input cycle: u64\n\
h := hash(cycle)\n\
neg := __i64_to_string(__u64_to_i64(mod(h, 1000)))\n\
big := __f64_to_string(f64_mul(to_f64(h), 1000000000000.0))\n\
small := __f64_to_string(f64_div(1.0, to_f64(u64_add(h, 1))))\n\
digits := __u64_to_string(h)\n\
empty := \"\"\n\
cat := str_concat(digits, empty, neg, \"|\", big)\n\
doc := json_object(json_with(\"h\", h), json_with(\"cat\", cat), json_with(\"list\", json_array(small, big, 1)))\n\
text := json_to_str(doc)\n";
agree(
src,
&["neg", "big", "small", "digits", "cat", "text"],
8,
&[],
);
let p3 = kernel(src, JitMode::Force);
let p = p3.program();
for name in [
"__i64_to_string",
"__f64_to_string",
"str_concat",
"json_to_str",
] {
assert!(
!(0..p.node_count()).any(|i| p.node_meta(i).name == name),
"`{name}` was not fused; cones: {:?}",
cones(&p3)
);
}
}
#[test]
fn the_vector_and_register_groups_have_named_lowerings() {
use polydat::ast::PortType as PT;
use polydat::compile::jit::{
JitOp, RegLaneRead, RegProducer, VecProducer, VecReducer, classify_node_typed,
};
use polydat::library::{register, vector_math};
let vec2 = [PT::VecF32, PT::VecF32];
assert_eq!(
classify_node_typed(&vector_math::VecAdd::new(), &vec2),
JitOp::VecProduce {
kind: VecProducer::Add,
scratch_base: 0
}
);
assert_eq!(
classify_node_typed(&vector_math::VecScale::new(), &[PT::VecF32, PT::F64]),
JitOp::VecProduce {
kind: VecProducer::Scale,
scratch_base: 0
}
);
assert_eq!(
classify_node_typed(&vector_math::VecNorm::new(), &[PT::VecF32]),
JitOp::VecProduce {
kind: VecProducer::Norm,
scratch_base: 0
}
);
assert_eq!(
classify_node_typed(&vector_math::HashVec::new(), &[PT::U64, PT::U64]),
JitOp::VecProduce {
kind: VecProducer::HashVec,
scratch_base: 0
}
);
assert_eq!(
classify_node_typed(&vector_math::Xxhash3Vec::new(), &[PT::U64, PT::U64]),
JitOp::VecProduce {
kind: VecProducer::XxHash3Vec,
scratch_base: 0
}
);
assert_eq!(
classify_node_typed(®ister::RegToVecF32::new(), &[PT::RegF32x4]),
JitOp::VecProduce {
kind: VecProducer::RegToVec,
scratch_base: 0
}
);
assert_eq!(
classify_node_typed(&vector_math::VecDot::new(), &vec2),
JitOp::VecReduce(VecReducer::Dot)
);
assert_eq!(
classify_node_typed(&vector_math::VecL2::new(), &vec2),
JitOp::VecReduce(VecReducer::L2)
);
assert_eq!(
classify_node_typed(&vector_math::VecCosine::new(), &vec2),
JitOp::VecReduce(VecReducer::Cosine)
);
assert_eq!(
classify_node_typed(&vector_math::LidMle::new(), &[PT::VecF32, PT::F64]),
JitOp::VecReduce(VecReducer::LidMle)
);
assert_eq!(
classify_node_typed(®ister::RegLaneF32::new(), &[PT::RegF32x4, PT::U64]),
JitOp::RegLane(RegLaneRead::F32)
);
assert_eq!(
classify_node_typed(®ister::RegLaneI16::new(), &[PT::RegI16x8, PT::U64]),
JitOp::RegLane(RegLaneRead::I16)
);
assert_eq!(
classify_node_typed(®ister::RegLaneI64::new(), &[PT::RegI64x2, PT::U64]),
JitOp::RegLane(RegLaneRead::I64)
);
assert_eq!(
classify_node_typed(
®ister::RegWithLaneF32::new(),
&[PT::RegF32x4, PT::U64, PT::F64]
),
JitOp::RegProduce(RegProducer::WithLaneF32)
);
assert_eq!(
classify_node_typed(®ister::RegGatherF32::new(), &[PT::VecF32, PT::U64]),
JitOp::RegProduce(RegProducer::GatherF32)
);
assert_eq!(
classify_node_typed(®ister::VecToRegF32::new(), &[PT::VecF32]),
JitOp::RegProduce(RegProducer::VecToRegF32)
);
assert_eq!(
classify_node_typed(®ister::RegMulI8::new(), &[PT::RegI8x16, PT::RegI8x16]),
JitOp::RegProduce(RegProducer::MulI8)
);
assert_eq!(
classify_node_typed(®ister::RegDotF32::new(), &[PT::RegF32x4, PT::RegF32x4]),
JitOp::RegDotF32
);
let mask: Vec<u64> = (0..16).rev().collect();
assert_eq!(
classify_node_typed(®ister::RegShuffleBytes::new(mask), &[PT::Reg128]),
JitOp::RegShuffleConst([15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, 0])
);
let src = "input cycle: u64\n\
a := hash_vec(cycle, 8)\n\
b := xxhash3_vec(cycle, 8)\n\
sum := vec_add(a, b)\n\
scaled := vec_scale(sum, 0.5)\n\
unit := vec_norm(scaled)\n\
dot := vec_dot(a, b)\n\
dist := vec_l2(a, b)\n\
cos := vec_cosine(a, unit)\n\
lid := lid_mle(vec_norm(a), 4.0)\n\
r := reg_gather_f32(a, 4)\n\
r2 := vec_to_reg_f32(reg_to_vec_f32(r))\n\
lane := reg_lane_f32(r2, 3)\n\
r3 := reg_with_lane_f32(r, mod(cycle, 4), lane)\n\
rd := reg_dot_f32(r, r3)\n\
i16 := reg_lane_i16(reg_splat_i16(cycle), 7)\n\
i64 := reg_lane_i64(reg_splat_i64(cycle), 1)\n\
prod := reg_mul_i8(reg_splat_i8(cycle), reg_splat_i8(u64_add(cycle, 3)))\n\
rev := reg_shuffle_bytes(prod, 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, 0)\n";
agree(
src,
&[
"sum", "scaled", "unit", "dot", "dist", "cos", "lid", "r", "r2", "lane", "r3", "rd",
"i16", "i64", "prod", "rev",
],
8,
&[],
);
let mut pure = compile_polydat_to_assembler(src)
.unwrap()
.try_compile_pure_jit()
.expect("every node of the two groups lowers");
let mut p1 = kernel(src, JitMode::Off);
for c in 0..8 {
polydat::Kernel::set_inputs(&mut pure, &[c]);
p1.set_inputs(&[c]);
for name in ["unit", "cos", "rd", "rev", "i16", "i64", "lid"] {
let want = p1.pull(name).clone();
let got = polydat::Kernel::pull(&mut pure, name).clone();
assert_eq!(got, want, "{name} at cycle {c}");
}
}
}
#[test]
fn a_string_read_is_an_owned_copy_that_outlives_the_next_write() {
let src = "input cycle: u64\nw := \"w{cycle}\"\nu := str_upper(w)\n";
let mut p3 = kernel(src, JitMode::Force);
p3.set_inputs(&[7]);
let first = p3.pull("u").clone();
assert_eq!(first.as_str(), "W7");
assert_eq!(
p3.pull("u").as_str(),
"W7",
"a second read is the same value"
);
p3.set_inputs(&[8]);
assert_eq!(p3.pull("u").as_str(), "W8");
assert_eq!(first.as_str(), "W7", "the earlier read is the reader's own");
}
#[test]
fn an_untyped_variadic_edge_lowers_with_its_wire_types() {
let src = "input cycle: u64\nh := hash(cycle)\nc := str_concat(h, h)\n";
agree(src, &["c"], 4, &["str_concat"]);
}
#[test]
fn a_non_decimal_format_u64_stays_on_p1_and_agrees() {
let src =
"input cycle: u64\nh := hash(cycle)\nd := format_u64(h, 10)\nx := format_u64(h, 16)\n";
agree(src, &["d", "x"], 4, &["format_u64"]);
let mut p3 = kernel(src, JitMode::Force);
p3.set_inputs(&[1]);
assert!(p3.pull("x").as_str().starts_with("0x"));
}
#[test]
fn a_failed_parse_is_a_diagnostic_at_every_tier() {
let src = "input cycle: u64\ns := \"nope\"\nn := __str_to_u64(s)\n";
for mode in [JitMode::Off, JitMode::Force] {
let mut k = kernel(src, mode);
k.set_inputs(&[0]);
let r = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| k.pull("n").as_u64()));
assert!(r.is_err(), "{mode:?} accepted an unparseable string");
}
}