module Component.RotaryFrequency import Float32Exact import FloatLiteralSpec import Std.List import Std.Natural -- The Llama half-split rotation uses theta_i = base^(-i / columns). -- Keep the exact root and its f32 residual together: the SM86 table writer -- uses both words to keep 4096-position angles accurate after reduction. def rotaryFrequencyResolution : Nat = (specPow2 96) def rotaryFrequencyRoot = (lambda unrestricted base : Nat . (lambda unrestricted columns : Nat . (lambda unrestricted index : Nat . (float32ExactRootKLow columns 1 (specPower base index) rotaryFrequencyResolution)))) def rotaryFrequencyHigh = (lambda unrestricted base : Nat . (lambda unrestricted columns : Nat . (lambda unrestricted index : Nat . (let unrestricted root = (rotaryFrequencyRoot base columns index) in (float32ExactBetween root rotaryFrequencyResolution (succ root) rotaryFrequencyResolution))))) def rotaryFrequencyLow = (lambda unrestricted base : Nat . (lambda unrestricted columns : Nat . (lambda unrestricted index : Nat . (let unrestricted root = (rotaryFrequencyRoot base columns index) in (float32ExactRemainder root (succ root) rotaryFrequencyResolution (rotaryFrequencyHigh base columns index)))))) def rotaryFrequencyRoundingMagic : Nat = (compile-time (float32ExactRational 12582912 1)) -- Adjacent high/low words let a launch schedule index a column without -- reevaluating the large exact root for every scalar patch. def rotaryFrequencyWords = (lambda unrestricted base : Nat . (lambda unrestricted columns : Nat . (nat-eliminate (lambda unrestricted current : Nat . (family StdList Nat)) (constructor StdList StdListEmpty Nat) (lambda unrestricted predecessor : Nat . (lambda unrestricted rest : (family StdList Nat) . (let unrestricted index = (naturalSaturatingSubtract (naturalSaturatingSubtract columns 1) predecessor) in (constructor StdList StdListCons Nat (rotaryFrequencyHigh base columns index) (constructor StdList StdListCons Nat (rotaryFrequencyLow base columns index) rest))))) columns)))