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        // Appended after the last tag; existing tags are cache-key stable.
405        CoreSemanticOp::Erf => encoder.u8(51),
406        CoreSemanticOp::ExtractDiag { axis_a, axis_b } => {
407            encoder.u8(31);
408            encoder.usize(*axis_a);
409            encoder.usize(*axis_b);
410        }
411        CoreSemanticOp::EmbedDiag { axis_a, axis_b } => {
412            encoder.u8(32);
413            encoder.usize(*axis_a);
414            encoder.usize(*axis_b);
415        }
416        CoreSemanticOp::Tril { k } => {
417            encoder.u8(33);
418            encoder.i64(*k);
419        }
420        CoreSemanticOp::Triu { k } => {
421            encoder.u8(34);
422            encoder.i64(*k);
423        }
424        CoreSemanticOp::Gather(config) => {
425            encoder.u8(35);
426            encode_gather(encoder, config);
427        }
428        CoreSemanticOp::GatherDynamicSliceSizes {
429            offset_dims,
430            collapsed_slice_dims,
431            start_index_map,
432            index_vector_dim,
433            slice_sizes,
434        } => {
435            encoder.u8(36);
436            encoder.usize_slice(offset_dims);
437            encoder.usize_slice(collapsed_slice_dims);
438            encoder.usize_slice(start_index_map);
439            encoder.usize(*index_vector_dim);
440            encode_dim_exprs(encoder, slice_sizes);
441        }
442        CoreSemanticOp::Scatter(config) => {
443            encoder.u8(37);
444            encode_scatter(encoder, config);
445        }
446        CoreSemanticOp::Slice(config) => {
447            encoder.u8(38);
448            encode_slice(encoder, config);
449        }
450        CoreSemanticOp::DynamicSlice { slice_sizes } => {
451            encoder.u8(39);
452            encoder.usize_slice(slice_sizes);
453        }
454        CoreSemanticOp::DynamicUpdateSlice => encoder.u8(40),
455        CoreSemanticOp::Pad(config) => {
456            encoder.u8(41);
457            encode_pad(encoder, config);
458        }
459        CoreSemanticOp::Concatenate { axis, input_count } => {
460            encoder.u8(42);
461            encoder.usize(*axis);
462            encoder.usize(*input_count);
463        }
464        CoreSemanticOp::Reverse { axes } => {
465            encoder.u8(43);
466            encoder.usize_slice(axes);
467        }
468        CoreSemanticOp::ShapeOf { axis } => {
469            encoder.u8(44);
470            encoder.usize(*axis);
471        }
472        CoreSemanticOp::DynamicTruncate { axis } => {
473            encoder.u8(45);
474            encoder.usize(*axis);
475        }
476        CoreSemanticOp::PadToMatch { axis } => {
477            encoder.u8(46);
478            encoder.usize(*axis);
479        }
480        CoreSemanticOp::ReduceProd { axes } => {
481            encoder.u8(47);
482            encoder.usize_slice(axes);
483        }
484        CoreSemanticOp::ReduceMax { axes } => {
485            encoder.u8(48);
486            encoder.usize_slice(axes);
487        }
488        CoreSemanticOp::ReduceMin { axes } => {
489            encoder.u8(49);
490            encoder.usize_slice(axes);
491        }
492        CoreSemanticOp::ReduceSumSquares { axes } => {
493            encoder.u8(50);
494            encoder.usize_slice(axes);
495        }
496    }
497}
498
499fn encode_metadata(encoder: &mut CanonicalEncoder, metadata: &ProgramValueMetadata) {
500    match (metadata.dtype(), metadata.scalar_identity()) {
501        // An externally defined scalar has no process-stable type identity, so its
502        // canonical name is what the program identity encodes. Two scalars that
503        // share a name would collide, which is the declaring contribution's
504        // responsibility.
505        (DType::External(_), Some(identity)) => {
506            encoder.u8(7);
507            encoder.string(identity);
508        }
509        (dtype, _) => encode_dtype(encoder, dtype),
510    }
511    encoder.usize(metadata.shape().len());
512    for extent in metadata.shape() {
513        match extent {
514            ShapeExtent::Exact(expression) => {
515                encoder.u8(0);
516                encode_dim_expr(encoder, expression);
517            }
518            ShapeExtent::UpperBound(expression) => {
519                encoder.u8(1);
520                encode_dim_expr(encoder, expression);
521            }
522            ShapeExtent::Unknown => encoder.u8(2),
523        }
524    }
525}
526
527fn encode_dtype(encoder: &mut CanonicalEncoder, dtype: DType) {
528    encoder.u8(match dtype {
529        DType::F32 => 0,
530        DType::F64 => 1,
531        DType::I32 => 2,
532        DType::I64 => 3,
533        DType::Bool => 4,
534        DType::C32 => 5,
535        DType::C64 => 6,
536        // INVARIANT: value metadata encodes an external scalar through its declared
537        // identity, and a core operation may not name one, so this arm is reached
538        // only if a value carried an external tag without an identity, which the
539        // builder rejects.
540        DType::External(_) => unreachable!("an external scalar reached a bare dtype encoding"),
541    });
542}
543
544fn encode_dim_exprs(encoder: &mut CanonicalEncoder, expressions: &[DimExpr]) {
545    encoder.usize(expressions.len());
546    for expression in expressions {
547        encode_dim_expr(encoder, expression);
548    }
549}
550
551fn encode_dim_expr(encoder: &mut CanonicalEncoder, expression: &DimExpr) {
552    match expression {
553        DimExpr::Const(value) => {
554            encoder.u8(0);
555            encoder.usize(*value);
556        }
557        DimExpr::InputDim { input_idx, axis } => {
558            encoder.u8(1);
559            encoder.usize(*input_idx);
560            encoder.usize(*axis);
561        }
562        DimExpr::Add(left, right) => encode_binary_dim_expr(encoder, 2, left, right),
563        DimExpr::Sub(left, right) => encode_binary_dim_expr(encoder, 3, left, right),
564        DimExpr::Mul(left, right) => encode_binary_dim_expr(encoder, 4, left, right),
565        DimExpr::FloorDiv(left, right) => encode_binary_dim_expr(encoder, 5, left, right),
566        DimExpr::Min(left, right) => encode_binary_dim_expr(encoder, 6, left, right),
567        DimExpr::Max(left, right) => encode_binary_dim_expr(encoder, 7, left, right),
568    }
569}
570
571fn encode_binary_dim_expr(
572    encoder: &mut CanonicalEncoder,
573    tag: u8,
574    left: &DimExpr,
575    right: &DimExpr,
576) {
577    encoder.u8(tag);
578    encode_dim_expr(encoder, left);
579    encode_dim_expr(encoder, right);
580}
581
582fn encode_effects(encoder: &mut CanonicalEncoder, effects: &[Effect]) {
583    encoder.usize(effects.len());
584    for effect in effects {
585        encoder.string(effect.resource().family());
586        encoder.u64(effect.resource().key());
587        encoder.u8(match effect.access() {
588            EffectAccess::Read => 0,
589            EffectAccess::Write => 1,
590        });
591    }
592}
593
594fn encode_aliases(encoder: &mut CanonicalEncoder, aliases: &[Alias]) {
595    encoder.usize(aliases.len());
596    for alias in aliases {
597        encoder.u8(match alias.kind() {
598            AliasKind::Fresh => 0,
599            AliasKind::ViewOf => 1,
600            AliasKind::MustAlias => 2,
601            AliasKind::ExternalAlias => 3,
602        });
603        encoder.usize(alias.output());
604        encode_option_usize(encoder, alias.input());
605        match alias.resource() {
606            Some(resource) => {
607                encoder.u8(1);
608                encoder.string(resource.family());
609                encoder.u64(resource.key());
610            }
611            None => encoder.u8(0),
612        }
613    }
614}
615
616fn encode_shape_guards(encoder: &mut CanonicalEncoder, guards: &[ShapeGuard]) {
617    encoder.usize(guards.len());
618    for guard in guards {
619        encoder.u8(match guard.relation() {
620            ProgramShapeRelation::Equal => 0,
621            ProgramShapeRelation::LessEqual => 1,
622            ProgramShapeRelation::GreaterEqual => 2,
623        });
624        encode_dim_expr(encoder, guard.lhs());
625        encode_dim_expr(encoder, guard.rhs());
626    }
627}
628
629fn encode_placement(encoder: &mut CanonicalEncoder, placement: SemanticPlacementConstraint) {
630    encoder.u8(match placement.kind() {
631        SemanticPlacementKind::Any => 0,
632        SemanticPlacementKind::SameAsInput => 1,
633    });
634    encode_option_usize(encoder, placement.input());
635}
636
637fn encode_option_usize(encoder: &mut CanonicalEncoder, value: Option<usize>) {
638    match value {
639        Some(value) => {
640            encoder.u8(1);
641            encoder.usize(value);
642        }
643        None => encoder.u8(0),
644    }
645}
646
647fn encode_compare(encoder: &mut CanonicalEncoder, direction: &CompareDir) {
648    encoder.u8(match direction {
649        CompareDir::Eq => 0,
650        CompareDir::Lt => 1,
651        CompareDir::Le => 2,
652        CompareDir::Gt => 3,
653        CompareDir::Ge => 4,
654    });
655}
656
657fn encode_dot_general(encoder: &mut CanonicalEncoder, config: &DotGeneralConfig) {
658    encoder.usize_slice(&config.lhs_contracting_dims);
659    encoder.usize_slice(&config.rhs_contracting_dims);
660    encoder.usize_slice(&config.lhs_batch_dims);
661    encoder.usize_slice(&config.rhs_batch_dims);
662}
663
664fn encode_gather(encoder: &mut CanonicalEncoder, config: &GatherConfig) {
665    encoder.usize_slice(&config.offset_dims);
666    encoder.usize_slice(&config.collapsed_slice_dims);
667    encoder.usize_slice(&config.start_index_map);
668    encoder.usize(config.index_vector_dim);
669    encoder.usize_slice(&config.slice_sizes);
670}
671
672fn encode_scatter(encoder: &mut CanonicalEncoder, config: &ScatterConfig) {
673    encoder.usize_slice(&config.update_window_dims);
674    encoder.usize_slice(&config.inserted_window_dims);
675    encoder.usize_slice(&config.scatter_dims_to_operand_dims);
676    encoder.usize(config.index_vector_dim);
677}
678
679fn encode_slice(encoder: &mut CanonicalEncoder, config: &SliceConfig) {
680    encoder.usize_slice(&config.starts);
681    encoder.usize_slice(&config.limits);
682    encoder.usize_slice(&config.strides);
683}
684
685fn encode_pad(encoder: &mut CanonicalEncoder, config: &PadConfig) {
686    encoder.i64_slice(&config.edge_padding_low);
687    encoder.i64_slice(&config.edge_padding_high);
688    encoder.i64_slice(&config.interior_padding);
689}