use std::f64::consts::{LN_2, LOG2_E};
use crate::scalar::{bf16, f16};
use crate::{
dtype::Constant,
kernel::{BOp, Kernel, Op, OpId, UOp},
};
fn constant_is_ln_2(c: &Constant) -> bool {
let val = match *c {
Constant::BF16(x) => bf16::from_le_bytes(x).to_f32() as f64,
Constant::F16(x) => f16::from_le_bytes(x).to_f32() as f64,
Constant::F32(x) => f32::from_le_bytes(x) as f64,
Constant::F64(x) => f64::from_le_bytes(x),
_ => return false,
};
(val - LN_2).abs() < 1e-6
}
fn constant_is_log2_e(c: &Constant) -> bool {
let val = match *c {
Constant::BF16(x) => bf16::from_le_bytes(x).to_f32() as f64,
Constant::F16(x) => f16::from_le_bytes(x).to_f32() as f64,
Constant::F32(x) => f32::from_le_bytes(x) as f64,
Constant::F64(x) => f64::from_le_bytes(x),
_ => return false,
};
(val - LOG2_E).abs() < 1e-6
}
impl Kernel {
pub fn exp2_to_exp(&mut self) {
let mut op_id = self.head;
while !op_id.is_null() {
let next = self.next_op(op_id);
if let &Op::Unary { x, uop: UOp::Exp2 } = self.at(op_id) {
if let &Op::Binary { x: left, y: right, bop: BOp::Mul } = self.at(x) {
let input = match (self.at(left), self.at(right)) {
(&Op::Const(c), _) if constant_is_log2_e(&c) => right,
(_, &Op::Const(c)) if constant_is_log2_e(&c) => left,
_ => OpId::NULL,
};
if input != OpId::NULL {
self.ops[op_id].op = Op::Unary { x: input, uop: UOp::Exp };
}
}
}
op_id = next;
}
}
pub fn log2_to_ln(&mut self) {
let mut op_id = self.head;
while !op_id.is_null() {
let next = self.next_op(op_id);
if let &Op::Binary { x: left, y: right, bop: BOp::Mul } = self.at(op_id) {
let ((&Op::Unary { x: log2_op, uop: UOp::Log2 }, const_op)
| (const_op, &Op::Unary { x: log2_op, uop: UOp::Log2 })) = (self.at(left), self.at(right))
else {
op_id = next;
continue;
};
if let &Op::Const(c) = const_op {
if constant_is_ln_2(&c) {
self.ops[op_id].op = Op::Unary { x: log2_op, uop: UOp::Ln };
}
}
}
op_id = next;
}
}
pub fn exp_to_exp2(&mut self) {
let mut op_id = self.head;
while !op_id.is_null() {
let next = self.next_op(op_id);
if let &Op::Unary { x, uop: UOp::Exp } = self.at(op_id) {
let dtype = self.dtype(x);
let y = self.insert_before(op_id, Op::Const(Constant::F64(LOG2_E.to_le_bytes()).cast(dtype)));
let z = self.insert_before(op_id, Op::Binary { x, y, bop: BOp::Mul });
self.ops[op_id].op = Op::Unary { x: z, uop: UOp::Exp2 };
}
op_id = next;
}
}
pub fn ln_to_log2(&mut self) {
let mut op_id = self.head;
while !op_id.is_null() {
let next = self.next_op(op_id);
if let &Op::Unary { x, uop: UOp::Ln } = self.at(op_id) {
let dtype = self.dtype(x);
let y = self.insert_before(op_id, Op::Const(Constant::F64(LN_2.to_le_bytes()).cast(dtype)));
let x = self.insert_before(op_id, Op::Unary { x, uop: UOp::Log2 });
self.ops[op_id].op = Op::Binary { x, y, bop: BOp::Mul };
}
op_id = next;
}
}
}