use rane::{mil, Program};
const BUILD_INFO: &str = concat!(
"{{\"coremlc-component-MIL\", \"3510.2.1\"}, ",
"{\"coremlc-version\", \"3505.4.1\"}, ",
"{\"coremltools-component-milinternal\", \"\"}, ",
"{\"coremltools-version\", \"9.0\"}}",
);
fn header_simple(ch: usize, sp: usize) -> String {
format!(
"program(1.3)\n[buildInfo = dict<string, string>({info})]\n{{\n func main<ios18>(tensor<fp16, [1, {ch}, 1, {sp}]> x) {{\n",
info=BUILD_INFO, ch=ch, sp=sp,
)
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
println!("MIL operation probe โ what does ANE actually allow?\n");
probe(
"P1: matmul โ cast(fp32) โ add(fp32 self,self)",
&mil_p1(),
true,
);
probe(
"P2: matmul + matmul โ add(fp16) โ cast(fp32)",
&mil::matmul_split_fp16_cast(64, 64, 64, 2).text,
false,
);
probe("P3: matmul โ cast(fp32) โ cast(fp16)", &mil_p3(), true);
probe(
"P4: matmul โ cast(fp32) โ add(fp32, fp32_const)",
&mil_p4(),
true,
);
probe("P5: 2 matmuls โ concat โ reduce_sum", &mil_p5(), true);
probe(
"P6: 2 matmuls โ cast each fp32 โ concat fp32 โ reduce_sum fp32",
&mil_p6(),
true,
);
probe(
"P7: 2 matmuls โ stack fp16 โ reduce_sum fp16 โ cast fp32",
&mil_p7(),
true,
);
probe(
"P8: stack fp16 โ reduce_sum[output_dtype=fp32]",
&mil_p8(),
true,
);
probe("P9: linear with fp16 weight", &mil_p9(), true);
probe("P10: cast x โ fp32 matmul fp32", &mil_p10(), true);
probe("P11: mul(0.5) โ matmul โ cast(fp32)", &mil_p11(), true);
probe("P12: quantize โ matmul โ dequantize", &mil_p12(), true);
probe("P13: matmul[output_dtype=fp32]", &mil_p13(), true);
probe("P14: conv2d 1x1", &mil_p14(), true);
probe(
"P15: quantize(x,scale,zero_point,output_dtype=int8) + int8 matmul",
&mil_p15(),
true,
);
probe("P16: quantize w/ axis + matmul", &mil_p16(), true);
probe(
"P17: constexpr_affine_dequantize weight + matmul",
&mil_p17(),
true,
);
probe("P18: matmul โ cast(int8)", &mil_p18(), true);
probe("P19: matmul โ cast(uint8)", &mil_p19(), true);
probe("P20: matmul โ cast(int16)", &mil_p20(), true);
probe("P21: int8 input + matmul int8", &mil_p21(), true);
probe("P22: fp_to_int_clamped", &mil_p22(), true);
probe(
"P23: constexpr_affine_dequantize per-tensor + matmul",
&mil_p23(),
true,
);
probe(
"P24: constexpr_affine_dequantize per-channel + matmul",
&mil_p24(),
true,
);
probe(
"P25: matmul int8 output declared directly",
&mil_p25(),
true,
);
probe("P26: matmul[output_dtype=int8]", &mil_p26(), true);
probe("P27: matmul โ mul(1/256) โ cast(int8)", &mil_p27(), true);
probe("P28: matmul โ mul(1/65536) โ cast(int8)", &mil_p28(), true);
Ok(())
}
fn probe(label: &str, mil_text: &str, dump: bool) {
println!("=== {label} ===");
if dump {
let path = format!("/tmp/probe_{}.mil", label.split(':').next().unwrap_or("p"));
let _ = std::fs::write(&path, mil_text);
}
let src = rane::Source {
text: mil_text.to_string(),
input_channels: 64,
input_spatial: 128,
output_channels: 64,
output_spatial: 64,
output_dtype: rane::OutputDtype::Fp32,
};
match Program::compile(&src, &[]) {
Ok(_) => println!(" โ COMPILED\n"),
Err(e) => {
let s = format!("{e}");
let safe_label = label.split(':').next().unwrap_or("p").replace(' ', "_");
let path = format!("/tmp/probe_err_{safe_label}.txt");
let _ = std::fs::write(&path, &s);
let tail: String = s
.chars()
.rev()
.take(800)
.collect::<String>()
.chars()
.rev()
.collect();
println!(" โ (tail) {tail}\n");
}
}
}
fn mil_p1() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul(&mut m, "mm", 64, 64, 64, 0, 64, "x");
m += " tensor<fp32, [1,64,1,64]> f = cast(x=mm_y, dtype=string(\"fp32\"))[name=string(\"f\")];\n";
m += " tensor<fp32, [1,64,1,64]> s = add(x=f, y=f)[name=string(\"s\")];\n";
m += &mil::mil_footer("s");
m
}
fn mil_p3() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul(&mut m, "mm", 64, 64, 64, 0, 64, "x");
m += " tensor<fp32, [1,64,1,64]> f = cast(x=mm_y, dtype=string(\"fp32\"))[name=string(\"f\")];\n";
m += " tensor<fp16, [1,64,1,64]> h = cast(x=f, dtype=string(\"fp16\"))[name=string(\"h\")];\n";
m += &mil::mil_footer("h");
m
}
fn mil_p4() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul(&mut m, "mm", 64, 64, 64, 0, 64, "x");
m += " tensor<fp32, [1,64,1,64]> f = cast(x=mm_y, dtype=string(\"fp32\"))[name=string(\"f\")];\n";
m += " tensor<fp32, [1,64,1,64]> z = const()[name=string(\"z\"), val=tensor<fp32, [1,64,1,64]>([";
for i in 0..4096 {
if i > 0 {
m += ",";
}
m += "0.0";
}
m += "])];\n";
m += " tensor<fp32, [1,64,1,64]> s = add(x=f, y=z)[name=string(\"s\")];\n";
m += &mil::mil_footer("s");
m
}
fn mil_p6() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul_chunk(&mut m, "c0", 0, 32, 64, 64, 64, "x");
mil::gen_dyn_matmul_chunk(&mut m, "c1", 32, 32, 64, 64, 64, "x");
m += " tensor<fp32, [1,64,1,64]> f0 = cast(x=c0_y, dtype=string(\"fp32\"))[name=string(\"f0\")];\n";
m += " tensor<fp32, [1,64,1,64]> f1 = cast(x=c1_y, dtype=string(\"fp32\"))[name=string(\"f1\")];\n";
m += " tensor<fp32, [1,128,1,64]> cc = concat(values=(f0, f1), axis=int32(1), interleave=bool(false))[name=string(\"cc\")];\n";
m += " tensor<int32, [4]> rsh = const()[name=string(\"rsh\"), val=tensor<int32, [4]>([1,2,64,64])];\n";
m += " tensor<fp32, [1,2,64,64]> r2 = reshape(shape=rsh, x=cc)[name=string(\"r2\")];\n";
m += " tensor<int32, [1]> axes = const()[name=string(\"axes\"), val=tensor<int32, [1]>([1])];\n";
m += " bool kd = const()[name=string(\"kd\"), val=bool(false)];\n";
m += " tensor<fp32, [1,64,64]> rs = reduce_sum(x=r2, axes=axes, keep_dims=kd)[name=string(\"rs\")];\n";
m += " tensor<int32, [4]> rsho = const()[name=string(\"rsho\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp32, [1,64,1,64]> y = reshape(shape=rsho, x=rs)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p7() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul_chunk(&mut m, "c0", 0, 32, 64, 64, 64, "x");
mil::gen_dyn_matmul_chunk(&mut m, "c1", 32, 32, 64, 64, 64, "x");
m += " tensor<fp16, [1,128,1,64]> cc = concat(values=(c0_y, c1_y), axis=int32(1), interleave=bool(false))[name=string(\"cc\")];\n";
m += " tensor<int32, [4]> rsh = const()[name=string(\"rsh\"), val=tensor<int32, [4]>([1,2,64,64])];\n";
m += " tensor<fp16, [1,2,64,64]> r2 = reshape(shape=rsh, x=cc)[name=string(\"r2\")];\n";
m += " tensor<int32, [1]> axes = const()[name=string(\"axes\"), val=tensor<int32, [1]>([1])];\n";
m += " bool kd = const()[name=string(\"kd\"), val=bool(false)];\n";
m += " tensor<fp16, [1,64,64]> rs = reduce_sum(x=r2, axes=axes, keep_dims=kd)[name=string(\"rs\")];\n";
m += " tensor<int32, [4]> rsho = const()[name=string(\"rsho\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> r = reshape(shape=rsho, x=rs)[name=string(\"r\")];\n";
m += " tensor<fp32, [1,64,1,64]> y = cast(x=r, dtype=string(\"fp32\"))[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p8() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul_chunk(&mut m, "c0", 0, 32, 64, 64, 64, "x");
mil::gen_dyn_matmul_chunk(&mut m, "c1", 32, 32, 64, 64, 64, "x");
m += " tensor<fp16, [1,128,1,64]> cc = concat(values=(c0_y, c1_y), axis=int32(1), interleave=bool(false))[name=string(\"cc\")];\n";
m += " tensor<int32, [4]> rsh = const()[name=string(\"rsh\"), val=tensor<int32, [4]>([1,2,64,64])];\n";
m += " tensor<fp16, [1,2,64,64]> r2 = reshape(shape=rsh, x=cc)[name=string(\"r2\")];\n";
m += " tensor<int32, [1]> axes = const()[name=string(\"axes\"), val=tensor<int32, [1]>([1])];\n";
m += " bool kd = const()[name=string(\"kd\"), val=bool(false)];\n";
m += " tensor<fp32, [1,64,64]> rs = reduce_sum(x=r2, axes=axes, keep_dims=kd, output_dtype=string(\"fp32\"))[name=string(\"rs\")];\n";
m += " tensor<int32, [4]> rsho = const()[name=string(\"rsho\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp32, [1,64,1,64]> y = reshape(shape=rsho, x=rs)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p9() -> String {
let mut m = header_simple(64, 128);
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
m += " tensor<int32, [4]> bw = const()[name=string(\"bw\"), val=tensor<int32, [4]>([0,0,0,64])];\n";
m += " tensor<int32, [4]> sw = const()[name=string(\"sw\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> w = slice_by_size(x=x, begin=bw, size=sw)[name=string(\"w\")];\n";
m += " tensor<int32, [2]> wsh = const()[name=string(\"wsh\"), val=tensor<int32, [2]>([64,64])];\n";
m += " tensor<fp16, [64,64]> w2 = reshape(shape=wsh, x=w)[name=string(\"w2\")];\n";
m += " tensor<int32, [2]> ash = const()[name=string(\"ash\"), val=tensor<int32, [2]>([64,64])];\n";
m += " tensor<fp16, [64,64]> a2 = reshape(shape=ash, x=a)[name=string(\"a2\")];\n";
m += " tensor<fp16, [64,64]> ly = linear(x=a2, weight=w2)[name=string(\"ly\")];\n";
m += " tensor<int32, [4]> osh = const()[name=string(\"osh\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> r = reshape(shape=osh, x=ly)[name=string(\"r\")];\n";
m += " tensor<fp32, [1,64,1,64]> y = cast(x=r, dtype=string(\"fp32\"))[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p10() -> String {
let mut m = header_simple(64, 128);
m += " tensor<fp32, [1,64,1,128]> xf = cast(x=x, dtype=string(\"fp32\"))[name=string(\"xf\")];\n";
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp32, [1,64,1,64]> a = slice_by_size(x=xf, begin=ba, size=sa)[name=string(\"a\")];\n";
m += " tensor<int32, [4]> bw = const()[name=string(\"bw\"), val=tensor<int32, [4]>([0,0,0,64])];\n";
m += " tensor<int32, [4]> sw = const()[name=string(\"sw\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp32, [1,64,1,64]> w = slice_by_size(x=xf, begin=bw, size=sw)[name=string(\"w\")];\n";
m += " tensor<int32, [4]> ra = const()[name=string(\"ra\"), val=tensor<int32, [4]>([1,1,64,64])];\n";
m += " tensor<fp32, [1,1,64,64]> a2 = reshape(shape=ra, x=a)[name=string(\"a2\")];\n";
m += " tensor<int32, [4]> pm = const()[name=string(\"pm\"), val=tensor<int32, [4]>([0,1,3,2])];\n";
m += " tensor<fp32, [1,1,64,64]> a3 = transpose(perm=pm, x=a2)[name=string(\"a3\")];\n";
m += " tensor<int32, [4]> rw = const()[name=string(\"rw\"), val=tensor<int32, [4]>([1,1,64,64])];\n";
m += " tensor<fp32, [1,1,64,64]> W = reshape(shape=rw, x=w)[name=string(\"W\")];\n";
m += " bool bF = const()[name=string(\"bF\"), val=bool(false)];\n";
m += " tensor<fp32, [1,1,64,64]> yh = matmul(transpose_x=bF, transpose_y=bF, x=a3, y=W)[name=string(\"yh\")];\n";
m += " tensor<fp32, [1,1,64,64]> yt = transpose(perm=pm, x=yh)[name=string(\"yt\")];\n";
m += " tensor<int32, [4]> ro = const()[name=string(\"ro\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp32, [1,64,1,64]> y = reshape(shape=ro, x=yt)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p11() -> String {
let mut m = header_simple(64, 128);
m += " tensor<fp16, []> half = const()[name=string(\"half\"), val=fp16(0.5)];\n";
m += " tensor<fp16, [1,64,1,128]> xh = mul(x=x, y=half)[name=string(\"xh\")];\n";
mil::gen_dyn_matmul(&mut m, "mm", 64, 64, 64, 0, 64, "xh");
m += " tensor<fp32, [1,64,1,64]> y = cast(x=mm_y, dtype=string(\"fp32\"))[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p12() -> String {
let mut m = header_simple(64, 128);
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
m += " tensor<int32, [4]> bw = const()[name=string(\"bw\"), val=tensor<int32, [4]>([0,0,0,64])];\n";
m += " tensor<int32, [4]> sw = const()[name=string(\"sw\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> w = slice_by_size(x=x, begin=bw, size=sw)[name=string(\"w\")];\n";
m += " tensor<fp16, []> scale = const()[name=string(\"scale\"), val=fp16(1.0)];\n";
m += " tensor<int8, [1,64,1,64]> aq = quantize(input=a, scale=scale, output_dtype=string(\"int8\"))[name=string(\"aq\")];\n";
m += " tensor<int8, [1,64,1,64]> wq = quantize(input=w, scale=scale, output_dtype=string(\"int8\"))[name=string(\"wq\")];\n";
m += " tensor<int32, [4]> ra = const()[name=string(\"ra\"), val=tensor<int32, [4]>([1,1,64,64])];\n";
m += " tensor<int8, [1,1,64,64]> a2 = reshape(shape=ra, x=aq)[name=string(\"a2\")];\n";
m += " tensor<int8, [1,1,64,64]> W = reshape(shape=ra, x=wq)[name=string(\"W\")];\n";
m += " bool bF = const()[name=string(\"bF\"), val=bool(false)];\n";
m += " tensor<int32, [1,1,64,64]> yh = matmul(transpose_x=bF, transpose_y=bF, x=a2, y=W)[name=string(\"yh\")];\n";
m += " tensor<int32, [4]> ro = const()[name=string(\"ro\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<int32, [1,64,1,64]> y = reshape(shape=ro, x=yh)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p13() -> String {
let mut m = header_simple(64, 128);
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
m += " tensor<int32, [4]> bw = const()[name=string(\"bw\"), val=tensor<int32, [4]>([0,0,0,64])];\n";
m += " tensor<int32, [4]> sw = const()[name=string(\"sw\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> w = slice_by_size(x=x, begin=bw, size=sw)[name=string(\"w\")];\n";
m += " tensor<int32, [4]> ra = const()[name=string(\"ra\"), val=tensor<int32, [4]>([1,1,64,64])];\n";
m += " tensor<fp16, [1,1,64,64]> a2 = reshape(shape=ra, x=a)[name=string(\"a2\")];\n";
m += " tensor<fp16, [1,1,64,64]> W = reshape(shape=ra, x=w)[name=string(\"W\")];\n";
m += " bool bF = const()[name=string(\"bF\"), val=bool(false)];\n";
m += " tensor<fp32, [1,1,64,64]> yh = matmul(transpose_x=bF, transpose_y=bF, x=a2, y=W, output_dtype=string(\"fp32\"))[name=string(\"yh\")];\n";
m += " tensor<int32, [4]> ro = const()[name=string(\"ro\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp32, [1,64,1,64]> y = reshape(shape=ro, x=yh)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p14() -> String {
let mut m = header_simple(64, 128);
m += " tensor<int32, [4]> bw = const()[name=string(\"bw\"), val=tensor<int32, [4]>([0,0,0,64])];\n";
m += " tensor<int32, [4]> sw = const()[name=string(\"sw\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> w = slice_by_size(x=x, begin=bw, size=sw)[name=string(\"w\")];\n";
m += " tensor<int32, [4]> wsh = const()[name=string(\"wsh\"), val=tensor<int32, [4]>([64,64,1,1])];\n";
m += " tensor<fp16, [64,64,1,1]> wk = reshape(shape=wsh, x=w)[name=string(\"wk\")];\n";
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
m += " tensor<int32, [2]> strides = const()[name=string(\"strides\"), val=tensor<int32, [2]>([1,1])];\n";
m += " tensor<int32, [2]> dilations = const()[name=string(\"dilations\"), val=tensor<int32, [2]>([1,1])];\n";
m += " tensor<int32, [4]> pad = const()[name=string(\"pad\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " string padt = string(\"valid\");\n";
m += " tensor<fp16, [1,64,1,64]> r = conv(x=a, weight=wk, strides=strides, pad_type=padt, pad=pad, dilations=dilations, groups=int32(1))[name=string(\"r\")];\n";
m += " tensor<fp32, [1,64,1,64]> y = cast(x=r, dtype=string(\"fp32\"))[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p15() -> String {
let mut m = header_simple(64, 128);
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
m += " tensor<int32, [4]> bw = const()[name=string(\"bw\"), val=tensor<int32, [4]>([0,0,0,64])];\n";
m += " tensor<int32, [4]> sw = const()[name=string(\"sw\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> w = slice_by_size(x=x, begin=bw, size=sw)[name=string(\"w\")];\n";
m += " tensor<fp16, []> sc = const()[name=string(\"sc\"), val=fp16(1.0)];\n";
m += " tensor<int8, []> zp = const()[name=string(\"zp\"), val=int8(0)];\n";
m += " tensor<int8, [1,64,1,64]> aq = quantize(x=a, scale=sc, zero_point=zp, output_dtype=string(\"int8\"))[name=string(\"aq\")];\n";
m += " tensor<int8, [1,64,1,64]> wq = quantize(x=w, scale=sc, zero_point=zp, output_dtype=string(\"int8\"))[name=string(\"wq\")];\n";
m += " tensor<int32, [4]> ra = const()[name=string(\"ra\"), val=tensor<int32, [4]>([1,1,64,64])];\n";
m += " tensor<int8, [1,1,64,64]> a2 = reshape(shape=ra, x=aq)[name=string(\"a2\")];\n";
m += " tensor<int8, [1,1,64,64]> W2 = reshape(shape=ra, x=wq)[name=string(\"W2\")];\n";
m += " bool bF = const()[name=string(\"bF\"), val=bool(false)];\n";
m += " tensor<int32, [1,1,64,64]> yh = matmul(transpose_x=bF, transpose_y=bF, x=a2, y=W2)[name=string(\"yh\")];\n";
m += " tensor<int32, [4]> ro = const()[name=string(\"ro\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<int32, [1,64,1,64]> y = reshape(shape=ro, x=yh)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p16() -> String {
let mut m = header_simple(64, 128);
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
m += " tensor<int32, [4]> bw = const()[name=string(\"bw\"), val=tensor<int32, [4]>([0,0,0,64])];\n";
m += " tensor<int32, [4]> sw = const()[name=string(\"sw\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> w = slice_by_size(x=x, begin=bw, size=sw)[name=string(\"w\")];\n";
m += " tensor<fp16, []> sc = const()[name=string(\"sc\"), val=fp16(1.0)];\n";
m += " tensor<int8, []> zp = const()[name=string(\"zp\"), val=int8(0)];\n";
m += " tensor<int32, []> ax = const()[name=string(\"ax\"), val=int32(-1)];\n";
m += " tensor<int8, [1,64,1,64]> aq = quantize(x=a, scale=sc, zero_point=zp, axis=ax, output_dtype=string(\"int8\"))[name=string(\"aq\")];\n";
m += " tensor<int8, [1,64,1,64]> wq = quantize(x=w, scale=sc, zero_point=zp, axis=ax, output_dtype=string(\"int8\"))[name=string(\"wq\")];\n";
m += " tensor<int32, [4]> ra = const()[name=string(\"ra\"), val=tensor<int32, [4]>([1,1,64,64])];\n";
m += " tensor<int8, [1,1,64,64]> a2 = reshape(shape=ra, x=aq)[name=string(\"a2\")];\n";
m += " tensor<int8, [1,1,64,64]> W2 = reshape(shape=ra, x=wq)[name=string(\"W2\")];\n";
m += " bool bF = const()[name=string(\"bF\"), val=bool(false)];\n";
m += " tensor<int32, [1,1,64,64]> yh = matmul(transpose_x=bF, transpose_y=bF, x=a2, y=W2)[name=string(\"yh\")];\n";
m += " tensor<int32, [4]> ro = const()[name=string(\"ro\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<int32, [1,64,1,64]> y = reshape(shape=ro, x=yh)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p17() -> String {
let mut m = header_simple(64, 128);
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
let mut weight_str = String::from("[");
for i in 0..(64 * 64) {
if i > 0 {
weight_str += ",";
}
weight_str += "1";
}
weight_str += "]";
m += &format!(" tensor<int8, [64,64]> wq = const()[name=string(\"wq\"), val=tensor<int8, [64,64]>({weight_str})];\n");
m += " tensor<fp16, []> sc = const()[name=string(\"sc\"), val=fp16(1.0)];\n";
m += " tensor<int8, []> zp = const()[name=string(\"zp\"), val=int8(0)];\n";
m += " tensor<fp16, [64,64]> wf = constexpr_affine_dequantize(quantized_data=wq, scale=sc, zero_point=zp, axis=int32(0))[name=string(\"wf\")];\n";
m += " tensor<int32, [2]> rsh = const()[name=string(\"rsh\"), val=tensor<int32, [2]>([64,64])];\n";
m += " tensor<fp16, [64,64]> a2 = reshape(shape=rsh, x=a)[name=string(\"a2\")];\n";
m += " bool bF = const()[name=string(\"bF\"), val=bool(false)];\n";
m += " tensor<fp16, [64,64]> ym = matmul(transpose_x=bF, transpose_y=bF, x=a2, y=wf)[name=string(\"ym\")];\n";
m += " tensor<int32, [4]> ro = const()[name=string(\"ro\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> yr = reshape(shape=ro, x=ym)[name=string(\"yr\")];\n";
m += " tensor<fp32, [1,64,1,64]> y = cast(x=yr, dtype=string(\"fp32\"))[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p18() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul(&mut m, "mm", 64, 64, 64, 0, 64, "x");
m += " tensor<int8, [1,64,1,64]> y = cast(x=mm_y, dtype=string(\"int8\"))[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p19() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul(&mut m, "mm", 64, 64, 64, 0, 64, "x");
m += " tensor<uint8, [1,64,1,64]> y = cast(x=mm_y, dtype=string(\"uint8\"))[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p20() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul(&mut m, "mm", 64, 64, 64, 0, 64, "x");
m += " tensor<int16, [1,64,1,64]> y = cast(x=mm_y, dtype=string(\"int16\"))[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p21() -> String {
let mut m = format!(
"program(1.3)\n[buildInfo = dict<string, string>({info})]\n{{\n func main<ios18>(tensor<int8, [1, 64, 1, 128]> x) {{\n",
info=BUILD_INFO,
);
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<int8, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
m += " tensor<int32, [4]> bw = const()[name=string(\"bw\"), val=tensor<int32, [4]>([0,0,0,64])];\n";
m += " tensor<int32, [4]> sw = const()[name=string(\"sw\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<int8, [1,64,1,64]> w = slice_by_size(x=x, begin=bw, size=sw)[name=string(\"w\")];\n";
m += " tensor<int32, [4]> ra = const()[name=string(\"ra\"), val=tensor<int32, [4]>([1,1,64,64])];\n";
m += " tensor<int8, [1,1,64,64]> a2 = reshape(shape=ra, x=a)[name=string(\"a2\")];\n";
m += " tensor<int8, [1,1,64,64]> W2 = reshape(shape=ra, x=w)[name=string(\"W2\")];\n";
m += " bool bF = const()[name=string(\"bF\"), val=bool(false)];\n";
m += " tensor<int32, [1,1,64,64]> yh = matmul(transpose_x=bF, transpose_y=bF, x=a2, y=W2)[name=string(\"yh\")];\n";
m += " tensor<int32, [4]> ro = const()[name=string(\"ro\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<int32, [1,64,1,64]> y = reshape(shape=ro, x=yh)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p22() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul(&mut m, "mm", 64, 64, 64, 0, 64, "x");
m += " tensor<int32, [1,64,1,64]> y = cast(x=mm_y, dtype=string(\"int32\"))[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p23() -> String {
let mut m = header_simple(64, 128);
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
let weight_vals = vec!["1"; 64 * 64].join(",");
m += &format!(" tensor<int8, [64,64]> wq = const()[name=string(\"wq\"), val=tensor<int8, [64,64]>([{weight_vals}])];\n");
m += " tensor<fp16, []> sc = const()[name=string(\"sc\"), val=fp16(1.0)];\n";
m += " tensor<int8, []> zp = const()[name=string(\"zp\"), val=int8(0)];\n";
m += " tensor<fp16, [64,64]> wf = constexpr_affine_dequantize(quantized_data=wq, zero_point=zp, scale=sc)[name=string(\"wf\")];\n";
m += " tensor<int32, [2]> rsh = const()[name=string(\"rsh\"), val=tensor<int32, [2]>([64,64])];\n";
m += " tensor<fp16, [64,64]> a2 = reshape(shape=rsh, x=a)[name=string(\"a2\")];\n";
m += " bool bF = const()[name=string(\"bF\"), val=bool(false)];\n";
m += " tensor<fp16, [64,64]> ym = matmul(transpose_x=bF, transpose_y=bF, x=a2, y=wf)[name=string(\"ym\")];\n";
m += " tensor<int32, [4]> ro = const()[name=string(\"ro\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> y = reshape(shape=ro, x=ym)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p24() -> String {
let mut m = header_simple(64, 128);
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
let weight_vals = vec!["1"; 64 * 64].join(",");
m += &format!(" tensor<int8, [64,64]> wq = const()[name=string(\"wq\"), val=tensor<int8, [64,64]>([{weight_vals}])];\n");
let scale_vals = vec!["1.0"; 64].join(",");
let zp_vals = vec!["0"; 64].join(",");
m += &format!(" tensor<fp16, [64]> sc = const()[name=string(\"sc\"), val=tensor<fp16, [64]>([{scale_vals}])];\n");
m += &format!(" tensor<int8, [64]> zp = const()[name=string(\"zp\"), val=tensor<int8, [64]>([{zp_vals}])];\n");
m += " tensor<fp16, [64,64]> wf = constexpr_affine_dequantize(quantized_data=wq, zero_point=zp, scale=sc, axis=int32(0))[name=string(\"wf\")];\n";
m += " tensor<int32, [2]> rsh = const()[name=string(\"rsh\"), val=tensor<int32, [2]>([64,64])];\n";
m += " tensor<fp16, [64,64]> a2 = reshape(shape=rsh, x=a)[name=string(\"a2\")];\n";
m += " bool bF = const()[name=string(\"bF\"), val=bool(false)];\n";
m += " tensor<fp16, [64,64]> ym = matmul(transpose_x=bF, transpose_y=bF, x=a2, y=wf)[name=string(\"ym\")];\n";
m += " tensor<int32, [4]> ro = const()[name=string(\"ro\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> y = reshape(shape=ro, x=ym)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p25() -> String {
let mut m = header_simple(64, 128);
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
m += " tensor<int32, [4]> bw = const()[name=string(\"bw\"), val=tensor<int32, [4]>([0,0,0,64])];\n";
m += " tensor<int32, [4]> sw = const()[name=string(\"sw\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> w = slice_by_size(x=x, begin=bw, size=sw)[name=string(\"w\")];\n";
m += " tensor<int32, [4]> ra = const()[name=string(\"ra\"), val=tensor<int32, [4]>([1,1,64,64])];\n";
m += " tensor<fp16, [1,1,64,64]> a2 = reshape(shape=ra, x=a)[name=string(\"a2\")];\n";
m += " tensor<int32, [4]> pm = const()[name=string(\"pm\"), val=tensor<int32, [4]>([0,1,3,2])];\n";
m += " tensor<fp16, [1,1,64,64]> a3 = transpose(perm=pm, x=a2)[name=string(\"a3\")];\n";
m += " tensor<int32, [4]> rw = const()[name=string(\"rw\"), val=tensor<int32, [4]>([1,1,64,64])];\n";
m += " tensor<fp16, [1,1,64,64]> W = reshape(shape=rw, x=w)[name=string(\"W\")];\n";
m += " bool bF = const()[name=string(\"bF\"), val=bool(false)];\n";
m += " tensor<int8, [1,1,64,64]> yh = matmul(transpose_x=bF, transpose_y=bF, x=a3, y=W)[name=string(\"yh\")];\n";
m += " tensor<int8, [1,1,64,64]> yt = transpose(perm=pm, x=yh)[name=string(\"yt\")];\n";
m += " tensor<int32, [4]> ro = const()[name=string(\"ro\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<int8, [1,64,1,64]> y = reshape(shape=ro, x=yt)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p26() -> String {
let mut m = header_simple(64, 128);
m += " tensor<int32, [4]> ba = const()[name=string(\"ba\"), val=tensor<int32, [4]>([0,0,0,0])];\n";
m += " tensor<int32, [4]> sa = const()[name=string(\"sa\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> a = slice_by_size(x=x, begin=ba, size=sa)[name=string(\"a\")];\n";
m += " tensor<int32, [4]> bw = const()[name=string(\"bw\"), val=tensor<int32, [4]>([0,0,0,64])];\n";
m += " tensor<int32, [4]> sw = const()[name=string(\"sw\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> w = slice_by_size(x=x, begin=bw, size=sw)[name=string(\"w\")];\n";
m += " tensor<int32, [4]> ra = const()[name=string(\"ra\"), val=tensor<int32, [4]>([1,1,64,64])];\n";
m += " tensor<fp16, [1,1,64,64]> a2 = reshape(shape=ra, x=a)[name=string(\"a2\")];\n";
m += " tensor<fp16, [1,1,64,64]> W = reshape(shape=ra, x=w)[name=string(\"W\")];\n";
m += " bool bF = const()[name=string(\"bF\"), val=bool(false)];\n";
m += " tensor<int8, [1,1,64,64]> yh = matmul(transpose_x=bF, transpose_y=bF, x=a2, y=W)[name=string(\"yh\"), output_dtype=string(\"int8\")];\n";
m += " tensor<int32, [4]> ro = const()[name=string(\"ro\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<int8, [1,64,1,64]> y = reshape(shape=ro, x=yh)[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p27() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul(&mut m, "mm", 64, 64, 64, 0, 64, "x");
m +=
" tensor<fp16, []> scale = const()[name=string(\"scale\"), val=fp16(0.00390625)];\n";
m += " tensor<fp16, [1,64,1,64]> scaled = mul(x=mm_y, y=scale)[name=string(\"scaled\")];\n";
m += " tensor<int8, [1,64,1,64]> y = cast(x=scaled, dtype=string(\"int8\"))[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p28() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul(&mut m, "mm", 64, 64, 64, 0, 64, "x");
m += " tensor<fp16, []> scale = const()[name=string(\"scale\"), val=fp16(0.0000152587890625)];\n";
m += " tensor<fp16, [1,64,1,64]> scaled = mul(x=mm_y, y=scale)[name=string(\"scaled\")];\n";
m += " tensor<int8, [1,64,1,64]> y = cast(x=scaled, dtype=string(\"int8\"))[name=string(\"y\")];\n";
m += &mil::mil_footer("y");
m
}
fn mil_p5() -> String {
let mut m = header_simple(64, 128);
mil::gen_dyn_matmul_chunk(&mut m, "c0", 0, 32, 64, 64, 64, "x");
mil::gen_dyn_matmul_chunk(&mut m, "c1", 32, 32, 64, 64, 64, "x");
m += " tensor<fp16, [1,128,1,64]> cc = concat(values=(c0_y, c1_y), axis=int32(1), interleave=bool(false))[name=string(\"cc\")];\n";
m += " tensor<int32, [4]> rsh = const()[name=string(\"rsh\"), val=tensor<int32, [4]>([1,2,64,64])];\n";
m += " tensor<fp16, [1,2,64,64]> r2 = reshape(shape=rsh, x=cc)[name=string(\"r2\")];\n";
m += " tensor<int32, [1]> axes = const()[name=string(\"axes\"), val=tensor<int32, [1]>([1])];\n";
m += " bool kd = const()[name=string(\"kd\"), val=bool(false)];\n";
m += " tensor<fp16, [1,64,64]> rs = reduce_sum(x=r2, axes=axes, keep_dims=kd)[name=string(\"rs\")];\n";
m += " tensor<int32, [4]> rsho = const()[name=string(\"rsho\"), val=tensor<int32, [4]>([1,64,1,64])];\n";
m += " tensor<fp16, [1,64,1,64]> y = reshape(shape=rsho, x=rs)[name=string(\"y\")];\n";
m += " tensor<fp32, [1,64,1,64]> yf = cast(x=y, dtype=string(\"fp32\"))[name=string(\"yf\")];\n";
m += &mil::mil_footer("yf");
m
}