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::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 (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 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}