-
Notifications
You must be signed in to change notification settings - Fork 53
Expand file tree
/
Copy pathreference.py
More file actions
76 lines (59 loc) · 2.25 KB
/
Copy pathreference.py
File metadata and controls
76 lines (59 loc) · 2.25 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
# SPDX-FileCopyrightText: Copyright (C) 2026 Advanced Micro Devices, Inc. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
import torch
import numpy as np
from ml_dtypes import bfloat16
def reference(A, B):
"""CPU reference: matrix-vector product ``C = A @ B`` (ground truth)."""
return A @ B
def generate_golden_reference(
M=128, K=128, seed=42
): # Defaults are tile-aligned minimums; tests always pass explicit values
"""
Generate golden reference data for GEMV (General Matrix-Vector Multiplication).
Parameters:
M: Number of rows of matrix A
K: Number of columns of matrix A (equals vector B length)
seed: Random seed
Returns:
dict: Contains 'A' (matrix), 'B' (vector), 'C' (output vector)
"""
torch.manual_seed(seed)
# Generate golden inputs
val_range = 4
A = torch.randn(M, K, dtype=torch.bfloat16) * val_range
B = torch.randn(K, dtype=torch.bfloat16) * val_range
# Generate golden outputs
C = reference(A, B)
return {
"A": A,
"B": B,
"C": C,
}
def generate_golden_reference_batched(M=128, K=128, num_batches=2, seed=42):
"""
Generate golden reference data for a batched GEMV (num_batches independent
matrix-vector products stacked contiguously, matching the GEMV op layout).
Parameters:
M: Number of rows of each matrix A
K: Number of columns of each matrix A (equals vector B length)
num_batches: Number of independent GEMVs
seed: Random seed
Returns:
dict: Contains 'A' (matrices), 'B' (vectors), 'C' (output vectors)
"""
torch.manual_seed(seed)
val_range = 4
A = torch.randn(num_batches, M, K, dtype=torch.bfloat16) * val_range
B = torch.randn(num_batches, K, dtype=torch.bfloat16) * val_range
C = torch.empty(num_batches, M, dtype=torch.bfloat16)
for b in range(num_batches):
C[b] = A[b] @ B[b]
return {"A": A, "B": B, "C": C}
def gelu_tanh_approx(x):
"""Tanh-approximation GELU, matching aie_kernels/aie2p/gelu.cc.
0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3))). Computed in float32.
"""
xf = np.asarray(x, dtype=np.float32)
inner = 0.79788456 * (xf + 0.044715 * xf**3)
return 0.5 * xf * (1.0 + np.tanh(inner))