Skip to main content

tenferro_runtime/program/
identity.rs

1use std::hash::Hasher;
2use tenferro_ops::dim_expr::DimExpr;
3use tenferro_ops::shape_extent::ShapeExtent;
4use tenferro_tensor::{
5    CompareDir, DType, DotGeneralConfig, GatherConfig, PadConfig, ScatterConfig, SliceConfig,
6};
7
8use super::op::{SemanticOp, SemanticOperation};
9use super::semantic::SemanticProgram;
10use super::{
11    Alias, AliasKind, CoreSemanticOp, Effect, EffectAccess, ProgramShapeRelation, ProgramValue,
12    ProgramValueMetadata, SemanticPlacementConstraint, SemanticPlacementKind, ShapeGuard,
13};
14
15/// Cached fixed-size identity of normalized semantic program structure.
16///
17/// Uses a fast 128-bit hash (two SipHash-1-3 passes).  Cryptographic
18/// strength is unnecessary because caches confirm matches with exact
19/// structural equality on collision.
20#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
21pub struct SemanticFingerprint([u8; 16]);
22
23impl SemanticFingerprint {
24    /// Borrow the cached hash bytes.
25    pub const fn as_bytes(&self) -> &[u8; 16] {
26        &self.0
27    }
28}
29
30pub(crate) struct SemanticIdentity {
31    pub(crate) fingerprint: SemanticFingerprint,
32    ordinals: Box<[u32]>,
33    #[cfg(test)]
34    pub(crate) fingerprint_computations: usize,
35}
36
37impl SemanticIdentity {
38    pub(crate) fn build(
39        inputs: &[ProgramValue],
40        outputs: &[ProgramValue],
41        values: &[ProgramValueMetadata],
42        operations: &[SemanticOperation],
43        shape_guards: &[ShapeGuard],
44    ) -> Self {
45        let ordinals = normalized_ordinals(inputs, values.len(), operations);
46        let mut encoder = CanonicalEncoder::new();
47        encoder.bytes(b"tenferro.semantic-program.v1");
48        encoder.usize(inputs.len());
49        for input in inputs {
50            encode_metadata(&mut encoder, &values[input.slot as usize]);
51        }
52        encoder.usize(operations.len());
53        for operation in operations {
54            encode_operation(&mut encoder, operation, values, &ordinals);
55        }
56        encoder.usize(outputs.len());
57        for output in outputs {
58            encoder.u32(ordinals[output.slot as usize]);
59        }
60        encode_shape_guards(&mut encoder, shape_guards);
61        Self {
62            fingerprint: SemanticFingerprint(encoder.finalize()),
63            ordinals,
64            #[cfg(test)]
65            fingerprint_computations: 1,
66        }
67    }
68
69    pub(crate) fn exact_eq(
70        &self,
71        left: &SemanticProgram,
72        other: &Self,
73        right: &SemanticProgram,
74    ) -> bool {
75        if self.fingerprint != other.fingerprint
76            || left.inputs.len() != right.inputs.len()
77            || left.outputs.len() != right.outputs.len()
78            || left.operations.len() != right.operations.len()
79            || left.shape_guards != right.shape_guards
80        {
81            return false;
82        }
83        if !left
84            .inputs
85            .iter()
86            .zip(&right.inputs)
87            .all(|(left_value, right_value)| {
88                left.values[left_value.slot as usize] == right.values[right_value.slot as usize]
89            })
90        {
91            return false;
92        }
93        if !left.operations.iter().zip(&right.operations).all(
94            |(left_operation, right_operation)| {
95                operation_exact_eq(
96                    left_operation,
97                    &left.values,
98                    &self.ordinals,
99                    right_operation,
100                    &right.values,
101                    &other.ordinals,
102                )
103            },
104        ) {
105            return false;
106        }
107        left.outputs
108            .iter()
109            .zip(&right.outputs)
110            .all(|(left_value, right_value)| {
111                self.ordinals[left_value.slot as usize] == other.ordinals[right_value.slot as usize]
112            })
113    }
114
115    // INVARIANT: this helper is introduced ahead of the P4-C1 cache owner
116    // wiring so retained-byte accounting can use the exact semantic identity
117    // payload size without duplicating ordinal internals.
118    #[allow(
119        dead_code,
120        reason = "P4-C1 preparation accounting consumes this helper"
121    )]
122    pub(crate) fn ordinals_retained_bytes(&self) -> Option<usize> {
123        self.ordinals.len().checked_mul(std::mem::size_of::<u32>())
124    }
125}
126
127fn normalized_ordinals(
128    inputs: &[ProgramValue],
129    value_count: usize,
130    operations: &[SemanticOperation],
131) -> Box<[u32]> {
132    let mut ordinals = vec![u32::MAX; value_count];
133    let mut next = 0_u32;
134    for input in inputs {
135        ordinals[input.slot as usize] = next;
136        next += 1;
137    }
138    for operation in operations {
139        for output in &operation.outputs {
140            ordinals[output.slot as usize] = next;
141            next += 1;
142        }
143    }
144    ordinals.into()
145}
146
147fn operation_exact_eq(
148    left: &SemanticOperation,
149    left_values: &[ProgramValueMetadata],
150    left_ordinals: &[u32],
151    right: &SemanticOperation,
152    right_values: &[ProgramValueMetadata],
153    right_ordinals: &[u32],
154) -> bool {
155    semantic_op_exact_eq(&left.op, &right.op)
156        && left.inputs.len() == right.inputs.len()
157        && left.outputs.len() == right.outputs.len()
158        && left.inputs.iter().zip(&right.inputs).all(|(left, right)| {
159            left_ordinals[left.slot as usize] == right_ordinals[right.slot as usize]
160        })
161        && left
162            .outputs
163            .iter()
164            .zip(&right.outputs)
165            .all(|(left, right)| {
166                left_values[left.slot as usize] == right_values[right.slot as usize]
167            })
168        && left.effects == right.effects
169        && left.aliases == right.aliases
170        && left.shape_guards == right.shape_guards
171        && left.placement == right.placement
172}
173
174fn semantic_op_exact_eq(left: &SemanticOp, right: &SemanticOp) -> bool {
175    match (left, right) {
176        (SemanticOp::Core(left), SemanticOp::Core(right)) => left == right,
177        (SemanticOp::Extension(left), SemanticOp::Extension(right)) => {
178            left.family_id() == right.family_id() && left.payload_eq(right.as_ref())
179        }
180        _ => false,
181    }
182}
183
184struct CanonicalEncoder {
185    buf: Vec<u8>,
186}
187
188impl CanonicalEncoder {
189    fn new() -> Self {
190        Self { buf: Vec::new() }
191    }
192
193    fn finalize(self) -> [u8; 16] {
194        let mut h1 = std::collections::hash_map::DefaultHasher::new();
195        let mut h2 = std::collections::hash_map::DefaultHasher::new();
196        h1.write(&self.buf);
197        h2.write(b"\x01");
198        h2.write(&self.buf);
199        let v1 = h1.finish().to_le_bytes();
200        let v2 = h2.finish().to_le_bytes();
201        let mut result = [0u8; 16];
202        result[..8].copy_from_slice(&v1);
203        result[8..].copy_from_slice(&v2);
204        result
205    }
206
207    fn raw(&mut self, bytes: &[u8]) {
208        self.buf.extend_from_slice(bytes);
209    }
210
211    fn bytes(&mut self, bytes: &[u8]) {
212        self.usize(bytes.len());
213        self.raw(bytes);
214    }
215
216    fn string(&mut self, value: &str) {
217        self.bytes(value.as_bytes());
218    }
219
220    fn u8(&mut self, value: u8) {
221        self.raw(&[value]);
222    }
223
224    fn u32(&mut self, value: u32) {
225        self.raw(&value.to_le_bytes());
226    }
227
228    fn u64(&mut self, value: u64) {
229        self.raw(&value.to_le_bytes());
230    }
231
232    fn i64(&mut self, value: i64) {
233        self.raw(&value.to_le_bytes());
234    }
235
236    fn usize(&mut self, value: usize) {
237        self.u64(value as u64);
238    }
239
240    fn usize_slice(&mut self, values: &[usize]) {
241        self.usize(values.len());
242        for &value in values {
243            self.usize(value);
244        }
245    }
246
247    fn i64_slice(&mut self, values: &[i64]) {
248        self.usize(values.len());
249        for &value in values {
250            self.i64(value);
251        }
252    }
253}
254
255impl Hasher for CanonicalEncoder {
256    fn finish(&self) -> u64 {
257        let mut h = std::collections::hash_map::DefaultHasher::new();
258        h.write(&self.buf);
259        h.finish()
260    }
261
262    fn write(&mut self, bytes: &[u8]) {
263        self.bytes(bytes);
264    }
265
266    fn write_u8(&mut self, value: u8) {
267        self.u8(value);
268    }
269
270    fn write_u16(&mut self, value: u16) {
271        self.raw(&value.to_le_bytes());
272    }
273
274    fn write_u32(&mut self, value: u32) {
275        self.u32(value);
276    }
277
278    fn write_u64(&mut self, value: u64) {
279        self.u64(value);
280    }
281
282    fn write_u128(&mut self, value: u128) {
283        self.raw(&value.to_le_bytes());
284    }
285
286    fn write_usize(&mut self, value: usize) {
287        self.usize(value);
288    }
289
290    fn write_i8(&mut self, value: i8) {
291        self.raw(&value.to_le_bytes());
292    }
293
294    fn write_i16(&mut self, value: i16) {
295        self.raw(&value.to_le_bytes());
296    }
297
298    fn write_i32(&mut self, value: i32) {
299        self.raw(&value.to_le_bytes());
300    }
301
302    fn write_i64(&mut self, value: i64) {
303        self.i64(value);
304    }
305
306    fn write_i128(&mut self, value: i128) {
307        self.raw(&value.to_le_bytes());
308    }
309
310    fn write_isize(&mut self, value: isize) {
311        self.i64(value as i64);
312    }
313}
314
315fn encode_operation(
316    encoder: &mut CanonicalEncoder,
317    operation: &SemanticOperation,
318    values: &[ProgramValueMetadata],
319    ordinals: &[u32],
320) {
321    match &operation.op {
322        SemanticOp::Core(op) => encode_core_op(encoder, op),
323        SemanticOp::Extension(op) => {
324            encoder.u8(1);
325            encoder.string(op.family_id());
326            op.payload_hash(encoder);
327        }
328    }
329    encoder.usize(operation.inputs.len());
330    for input in &operation.inputs {
331        encoder.u32(ordinals[input.slot as usize]);
332    }
333    encoder.usize(operation.outputs.len());
334    for output in &operation.outputs {
335        encode_metadata(encoder, &values[output.slot as usize]);
336    }
337    encode_effects(encoder, &operation.effects);
338    encode_aliases(encoder, &operation.aliases);
339    encode_shape_guards(encoder, &operation.shape_guards);
340    encode_placement(encoder, operation.placement);
341}
342
343fn encode_core_op(encoder: &mut CanonicalEncoder, op: &CoreSemanticOp) {
344    encoder.u8(0);
345    match op {
346        CoreSemanticOp::Add => encoder.u8(0),
347        CoreSemanticOp::Sub => encoder.u8(1),
348        CoreSemanticOp::Mul => encoder.u8(2),
349        CoreSemanticOp::Neg => encoder.u8(3),
350        CoreSemanticOp::Conj => encoder.u8(4),
351        CoreSemanticOp::DotGeneral { config } => {
352            encoder.u8(5);
353            encode_dot_general(encoder, config);
354        }
355        CoreSemanticOp::Transpose { perm } => {
356            encoder.u8(6);
357            encoder.usize_slice(perm);
358        }
359        CoreSemanticOp::Reshape { to_shape } => {
360            encoder.u8(7);
361            encode_dim_exprs(encoder, to_shape);
362        }
363        CoreSemanticOp::BroadcastInDim { shape, dims } => {
364            encoder.u8(8);
365            encode_dim_exprs(encoder, shape);
366            encoder.usize_slice(dims);
367        }
368        CoreSemanticOp::Convert { from, to } => {
369            encoder.u8(9);
370            encode_dtype(encoder, *from);
371            encode_dtype(encoder, *to);
372        }
373        CoreSemanticOp::Constant { dtype, bytes } => {
374            encoder.u8(10);
375            encode_dtype(encoder, *dtype);
376            encoder.bytes(bytes);
377        }
378        CoreSemanticOp::ReduceSum { axes } => {
379            encoder.u8(11);
380            encoder.usize_slice(axes);
381        }
382        CoreSemanticOp::Div => encoder.u8(12),
383        CoreSemanticOp::Rem => encoder.u8(13),
384        CoreSemanticOp::Abs => encoder.u8(14),
385        CoreSemanticOp::Sign => encoder.u8(15),
386        CoreSemanticOp::Maximum => encoder.u8(16),
387        CoreSemanticOp::Minimum => encoder.u8(17),
388        CoreSemanticOp::Compare(direction) => {
389            encoder.u8(18);
390            encode_compare(encoder, direction);
391        }
392        CoreSemanticOp::Select => encoder.u8(19),
393        CoreSemanticOp::Clamp => encoder.u8(20),
394        CoreSemanticOp::Exp => encoder.u8(21),
395        CoreSemanticOp::Log => encoder.u8(22),
396        CoreSemanticOp::Sin => encoder.u8(23),
397        CoreSemanticOp::Cos => encoder.u8(24),
398        CoreSemanticOp::Tanh => encoder.u8(25),
399        CoreSemanticOp::Sqrt => encoder.u8(26),
400        CoreSemanticOp::Rsqrt => encoder.u8(27),
401        CoreSemanticOp::Pow => encoder.u8(28),
402        CoreSemanticOp::Expm1 => encoder.u8(29),
403        CoreSemanticOp::Log1p => encoder.u8(30),
404        CoreSemanticOp::ExtractDiag { axis_a, axis_b } => {
405            encoder.u8(31);
406            encoder.usize(*axis_a);
407            encoder.usize(*axis_b);
408        }
409        CoreSemanticOp::EmbedDiag { axis_a, axis_b } => {
410            encoder.u8(32);
411            encoder.usize(*axis_a);
412            encoder.usize(*axis_b);
413        }
414        CoreSemanticOp::Tril { k } => {
415            encoder.u8(33);
416            encoder.i64(*k);
417        }
418        CoreSemanticOp::Triu { k } => {
419            encoder.u8(34);
420            encoder.i64(*k);
421        }
422        CoreSemanticOp::Gather(config) => {
423            encoder.u8(35);
424            encode_gather(encoder, config);
425        }
426        CoreSemanticOp::GatherDynamicSliceSizes {
427            offset_dims,
428            collapsed_slice_dims,
429            start_index_map,
430            index_vector_dim,
431            slice_sizes,
432        } => {
433            encoder.u8(36);
434            encoder.usize_slice(offset_dims);
435            encoder.usize_slice(collapsed_slice_dims);
436            encoder.usize_slice(start_index_map);
437            encoder.usize(*index_vector_dim);
438            encode_dim_exprs(encoder, slice_sizes);
439        }
440        CoreSemanticOp::Scatter(config) => {
441            encoder.u8(37);
442            encode_scatter(encoder, config);
443        }
444        CoreSemanticOp::Slice(config) => {
445            encoder.u8(38);
446            encode_slice(encoder, config);
447        }
448        CoreSemanticOp::DynamicSlice { slice_sizes } => {
449            encoder.u8(39);
450            encoder.usize_slice(slice_sizes);
451        }
452        CoreSemanticOp::DynamicUpdateSlice => encoder.u8(40),
453        CoreSemanticOp::Pad(config) => {
454            encoder.u8(41);
455            encode_pad(encoder, config);
456        }
457        CoreSemanticOp::Concatenate { axis, input_count } => {
458            encoder.u8(42);
459            encoder.usize(*axis);
460            encoder.usize(*input_count);
461        }
462        CoreSemanticOp::Reverse { axes } => {
463            encoder.u8(43);
464            encoder.usize_slice(axes);
465        }
466        CoreSemanticOp::ShapeOf { axis } => {
467            encoder.u8(44);
468            encoder.usize(*axis);
469        }
470        CoreSemanticOp::DynamicTruncate { axis } => {
471            encoder.u8(45);
472            encoder.usize(*axis);
473        }
474        CoreSemanticOp::PadToMatch { axis } => {
475            encoder.u8(46);
476            encoder.usize(*axis);
477        }
478        CoreSemanticOp::ReduceProd { axes } => {
479            encoder.u8(47);
480            encoder.usize_slice(axes);
481        }
482        CoreSemanticOp::ReduceMax { axes } => {
483            encoder.u8(48);
484            encoder.usize_slice(axes);
485        }
486        CoreSemanticOp::ReduceMin { axes } => {
487            encoder.u8(49);
488            encoder.usize_slice(axes);
489        }
490        CoreSemanticOp::ReduceSumSquares { axes } => {
491            encoder.u8(50);
492            encoder.usize_slice(axes);
493        }
494    }
495}
496
497fn encode_metadata(encoder: &mut CanonicalEncoder, metadata: &ProgramValueMetadata) {
498    encode_dtype(encoder, metadata.dtype());
499    encoder.usize(metadata.shape().len());
500    for extent in metadata.shape() {
501        match extent {
502            ShapeExtent::Exact(expression) => {
503                encoder.u8(0);
504                encode_dim_expr(encoder, expression);
505            }
506            ShapeExtent::UpperBound(expression) => {
507                encoder.u8(1);
508                encode_dim_expr(encoder, expression);
509            }
510            ShapeExtent::Unknown => encoder.u8(2),
511        }
512    }
513}
514
515fn encode_dtype(encoder: &mut CanonicalEncoder, dtype: DType) {
516    encoder.u8(match dtype {
517        DType::F32 => 0,
518        DType::F64 => 1,
519        DType::I32 => 2,
520        DType::I64 => 3,
521        DType::Bool => 4,
522        DType::C32 => 5,
523        DType::C64 => 6,
524    });
525}
526
527fn encode_dim_exprs(encoder: &mut CanonicalEncoder, expressions: &[DimExpr]) {
528    encoder.usize(expressions.len());
529    for expression in expressions {
530        encode_dim_expr(encoder, expression);
531    }
532}
533
534fn encode_dim_expr(encoder: &mut CanonicalEncoder, expression: &DimExpr) {
535    match expression {
536        DimExpr::Const(value) => {
537            encoder.u8(0);
538            encoder.usize(*value);
539        }
540        DimExpr::InputDim { input_idx, axis } => {
541            encoder.u8(1);
542            encoder.usize(*input_idx);
543            encoder.usize(*axis);
544        }
545        DimExpr::Add(left, right) => encode_binary_dim_expr(encoder, 2, left, right),
546        DimExpr::Sub(left, right) => encode_binary_dim_expr(encoder, 3, left, right),
547        DimExpr::Mul(left, right) => encode_binary_dim_expr(encoder, 4, left, right),
548        DimExpr::FloorDiv(left, right) => encode_binary_dim_expr(encoder, 5, left, right),
549        DimExpr::Min(left, right) => encode_binary_dim_expr(encoder, 6, left, right),
550        DimExpr::Max(left, right) => encode_binary_dim_expr(encoder, 7, left, right),
551    }
552}
553
554fn encode_binary_dim_expr(
555    encoder: &mut CanonicalEncoder,
556    tag: u8,
557    left: &DimExpr,
558    right: &DimExpr,
559) {
560    encoder.u8(tag);
561    encode_dim_expr(encoder, left);
562    encode_dim_expr(encoder, right);
563}
564
565fn encode_effects(encoder: &mut CanonicalEncoder, effects: &[Effect]) {
566    encoder.usize(effects.len());
567    for effect in effects {
568        encoder.string(effect.resource().family());
569        encoder.u64(effect.resource().key());
570        encoder.u8(match effect.access() {
571            EffectAccess::Read => 0,
572            EffectAccess::Write => 1,
573        });
574    }
575}
576
577fn encode_aliases(encoder: &mut CanonicalEncoder, aliases: &[Alias]) {
578    encoder.usize(aliases.len());
579    for alias in aliases {
580        encoder.u8(match alias.kind() {
581            AliasKind::Fresh => 0,
582            AliasKind::ViewOf => 1,
583            AliasKind::MustAlias => 2,
584            AliasKind::ExternalAlias => 3,
585        });
586        encoder.usize(alias.output());
587        encode_option_usize(encoder, alias.input());
588        match alias.resource() {
589            Some(resource) => {
590                encoder.u8(1);
591                encoder.string(resource.family());
592                encoder.u64(resource.key());
593            }
594            None => encoder.u8(0),
595        }
596    }
597}
598
599fn encode_shape_guards(encoder: &mut CanonicalEncoder, guards: &[ShapeGuard]) {
600    encoder.usize(guards.len());
601    for guard in guards {
602        encoder.u8(match guard.relation() {
603            ProgramShapeRelation::Equal => 0,
604            ProgramShapeRelation::LessEqual => 1,
605            ProgramShapeRelation::GreaterEqual => 2,
606        });
607        encode_dim_expr(encoder, guard.lhs());
608        encode_dim_expr(encoder, guard.rhs());
609    }
610}
611
612fn encode_placement(encoder: &mut CanonicalEncoder, placement: SemanticPlacementConstraint) {
613    encoder.u8(match placement.kind() {
614        SemanticPlacementKind::Any => 0,
615        SemanticPlacementKind::SameAsInput => 1,
616    });
617    encode_option_usize(encoder, placement.input());
618}
619
620fn encode_option_usize(encoder: &mut CanonicalEncoder, value: Option<usize>) {
621    match value {
622        Some(value) => {
623            encoder.u8(1);
624            encoder.usize(value);
625        }
626        None => encoder.u8(0),
627    }
628}
629
630fn encode_compare(encoder: &mut CanonicalEncoder, direction: &CompareDir) {
631    encoder.u8(match direction {
632        CompareDir::Eq => 0,
633        CompareDir::Lt => 1,
634        CompareDir::Le => 2,
635        CompareDir::Gt => 3,
636        CompareDir::Ge => 4,
637    });
638}
639
640fn encode_dot_general(encoder: &mut CanonicalEncoder, config: &DotGeneralConfig) {
641    encoder.usize_slice(&config.lhs_contracting_dims);
642    encoder.usize_slice(&config.rhs_contracting_dims);
643    encoder.usize_slice(&config.lhs_batch_dims);
644    encoder.usize_slice(&config.rhs_batch_dims);
645}
646
647fn encode_gather(encoder: &mut CanonicalEncoder, config: &GatherConfig) {
648    encoder.usize_slice(&config.offset_dims);
649    encoder.usize_slice(&config.collapsed_slice_dims);
650    encoder.usize_slice(&config.start_index_map);
651    encoder.usize(config.index_vector_dim);
652    encoder.usize_slice(&config.slice_sizes);
653}
654
655fn encode_scatter(encoder: &mut CanonicalEncoder, config: &ScatterConfig) {
656    encoder.usize_slice(&config.update_window_dims);
657    encoder.usize_slice(&config.inserted_window_dims);
658    encoder.usize_slice(&config.scatter_dims_to_operand_dims);
659    encoder.usize(config.index_vector_dim);
660}
661
662fn encode_slice(encoder: &mut CanonicalEncoder, config: &SliceConfig) {
663    encoder.usize_slice(&config.starts);
664    encoder.usize_slice(&config.limits);
665    encoder.usize_slice(&config.strides);
666}
667
668fn encode_pad(encoder: &mut CanonicalEncoder, config: &PadConfig) {
669    encoder.i64_slice(&config.edge_padding_low);
670    encoder.i64_slice(&config.edge_padding_high);
671    encoder.i64_slice(&config.interior_padding);
672}