MCPcopy Create free account
hub / github.com/ciphermodelabs/ciphercore / broadcast_shapes

Function broadcast_shapes

ciphercore-base/src/broadcast.rs:5–27  ·  view source on GitHub ↗
(s1: ArrayShape, s2: ArrayShape)

Source from the content-addressed store, hash-verified

3use std::cmp::max;
4
5pub(super) fn broadcast_shapes(s1: ArrayShape, s2: ArrayShape) -> Result<ArrayShape> {
6 let result_length = max(s1.len(), s2.len());
7 let offset1 = result_length - s1.len();
8 let offset2 = result_length - s2.len();
9 let mut result = vec![];
10 for i in 0..result_length {
11 let mut value1 = 1;
12 if i >= offset1 {
13 value1 = s1[i - offset1];
14 }
15 let mut value2 = 1;
16 if i >= offset2 {
17 value2 = s2[i - offset2];
18 }
19 if value1 > 1 && value2 > 1 && value1 != value2 {
20 return Err(runtime_error!(
21 "Invalid broadcast: shapes {s1:?} and {s2:?}"
22 ));
23 }
24 result.push(max(value1, value2));
25 }
26 Ok(result)
27}
28
29fn broadcast_pair(t1: Type, t2: Type) -> Result<Type> {
30 if t1.get_scalar_type() != t2.get_scalar_type() {

Callers 7

broadcast_pairFunction · 0.85
mixed_multiply_inferenceFunction · 0.85
matmul_type_inferenceFunction · 0.85
gemm_type_inferenceFunction · 0.85
newMethod · 0.85

Calls 1

pushMethod · 0.45