| 3391 | } |
| 3392 | |
| 3393 | fn gemm_helper( |
| 3394 | t0: Type, |
| 3395 | t1: Type, |
| 3396 | array0: Vec<u64>, |
| 3397 | array1: Vec<u64>, |
| 3398 | expected: Vec<u64>, |
| 3399 | ) -> Result<()> { |
| 3400 | let trans_perm0 = transpose_permutation(t0.get_shape().len()); |
| 3401 | let trans_perm1 = transpose_permutation(t1.get_shape().len()); |
| 3402 | |
| 3403 | let c = create_context()?; |
| 3404 | let g = c.create_graph()?; |
| 3405 | let i0 = g.input(t0.clone())?; |
| 3406 | let i1 = g.input(t1.clone())?; |
| 3407 | let trans_i0 = i0.permute_axes(trans_perm0)?; |
| 3408 | let trans_i1 = i1.permute_axes(trans_perm1)?; |
| 3409 | let gemm_false_false = i0.gemm(trans_i1.clone(), false, false)?; |
| 3410 | let gemm_false_true = i0.gemm(i1.clone(), false, true)?; |
| 3411 | let gemm_true_false = trans_i0.gemm(trans_i1, true, false)?; |
| 3412 | let gemm_true_true = trans_i0.gemm(i1, true, true)?; |
| 3413 | let o = g.create_tuple(vec![ |
| 3414 | gemm_false_false.clone(), |
| 3415 | gemm_false_true, |
| 3416 | gemm_true_false, |
| 3417 | gemm_true_true, |
| 3418 | ])?; |
| 3419 | g.set_output_node(o.clone())?; |
| 3420 | g.finalize()?; |
| 3421 | c.set_main_graph(g.clone())?; |
| 3422 | c.finalize()?; |
| 3423 | |
| 3424 | let value0 = Value::from_flattened_array(&array0, t0.get_scalar_type())?; |
| 3425 | let value1 = Value::from_flattened_array(&array1, t1.get_scalar_type())?; |
| 3426 | let results = random_evaluate(g, vec![value0, value1])?.to_vector()?; |
| 3427 | |
| 3428 | let result_t = gemm_false_false.get_type()?; |
| 3429 | for result in results { |
| 3430 | assert_eq!(result.to_flattened_array_u64(result_t.clone())?, expected); |
| 3431 | } |
| 3432 | |
| 3433 | Ok(()) |
| 3434 | } |
| 3435 | |
| 3436 | fn gemm_helper_random(t0: Type, t1: Type) -> Result<()> { |
| 3437 | let trans_perm0 = transpose_permutation(t0.get_shape().len()); |