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#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
21pub struct SemanticFingerprint([u8; 16]);
22
23impl SemanticFingerprint {
24 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 #[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}