|
4 | 4 | import itertools |
5 | 5 | import math |
6 | 6 | from contextlib import contextmanager |
| 7 | +from collections import defaultdict |
7 | 8 |
|
8 | 9 | from dataclasses import dataclass, field |
9 | 10 | from typing import Optional, Tuple, Dict, Any, Sequence, Literal |
@@ -123,6 +124,7 @@ def get_default_rounding_mode(opname: Optional[str] = None): |
123 | 124 | RoundingMode.RZ: bc.RoundingMode.ZERO, |
124 | 125 | RoundingMode.RM: bc.RoundingMode.NEGATIVE_INF, |
125 | 126 | RoundingMode.RP: bc.RoundingMode.POSITIVE_INF, |
| 127 | + RoundingMode.RA: bc.RoundingMode.NEAREST_AWAY, |
126 | 128 | RoundingMode.FULL: bc.RoundingMode.FULL, |
127 | 129 | RoundingMode.APPROX: bc.RoundingMode.APPROX, |
128 | 130 | RoundingMode.RZI: bc.RoundingMode.NEAREST_INT_TO_ZERO |
@@ -225,6 +227,79 @@ def check_shapes_eq(a: TileTy, b: TileTy, |
225 | 227 | f"got {a.shape} and {b.shape}", loc) |
226 | 228 |
|
227 | 229 |
|
| 230 | +F64 = datatype.float64 |
| 231 | +F32 = datatype.float32 |
| 232 | +TF32 = datatype.tfloat32 |
| 233 | +F16 = datatype.float16 |
| 234 | +BF16 = datatype.bfloat16 |
| 235 | +F8E5M2 = datatype.float8_e5m2 |
| 236 | +F8E8M0FNU = datatype.float8_e8m0fnu |
| 237 | +F8E4M3FN = datatype.float8_e4m3fn |
| 238 | +F4E2M1FN = datatype.float4_e2m1fn |
| 239 | +ALL = (F64, F32, TF32, F16, BF16, F8E5M2, F8E8M0FNU, F8E4M3FN, F4E2M1FN) |
| 240 | +B133 = BytecodeVersion.V_13_3 |
| 241 | +B134 = BytecodeVersion.V_13_4 |
| 242 | + |
| 243 | +_FTOF_ROUNDING_ROWS = ( |
| 244 | + {(i, F64): (RoundingMode.RN, None) for i in ALL}, |
| 245 | + {(i, F32): (RoundingMode.RN, None) for i in ALL}, |
| 246 | + {(i, TF32): (RoundingMode.RN, None) for i in ALL}, |
| 247 | + {(i, BF16): (RoundingMode.RN, None) for i in ALL}, |
| 248 | + {(i, F16): (RoundingMode.RN, None) for i in ALL}, |
| 249 | + {(i, F8E4M3FN): (RoundingMode.RN, None) for i in ALL}, |
| 250 | + {(i, F8E5M2): (RoundingMode.RN, None) for i in ALL}, |
| 251 | + {(i, F4E2M1FN): (RoundingMode.RN, None) for i in ALL}, |
| 252 | + |
| 253 | + {(i, F64): (RoundingMode.RZ, B134) for i in ALL}, |
| 254 | + {(i, F32): (RoundingMode.RZ, B134) for i in ALL}, |
| 255 | + {(i, TF32): (RoundingMode.RZ, B134) for i in ALL}, |
| 256 | + {(i, F16): (RoundingMode.RZ, B134) for i in ALL}, |
| 257 | + {(i, BF16): (RoundingMode.RZ, B134) for i in ALL}, |
| 258 | + {(i, F8E8M0FNU): (RoundingMode.RZ, B133) for i in ALL if i not in {F64, F8E5M2, F8E4M3FN}}, |
| 259 | + {(i, F8E8M0FNU): (RoundingMode.RZ, B134) for i in (F64, F8E5M2, F8E4M3FN)}, |
| 260 | + |
| 261 | + {(i, F64): (RoundingMode.RM, B134) for i in ALL}, |
| 262 | + {(i, F32): (RoundingMode.RM, B134) for i in ALL if i not in {F8E8M0FNU}}, |
| 263 | + {(i, F16): (RoundingMode.RM, B134) for i in ALL if i not in {F64, F32, TF32, F8E8M0FNU}}, |
| 264 | + {(i, F64): (RoundingMode.RP, B134) for i in ALL}, |
| 265 | + {(i, F32): (RoundingMode.RP, B134) for i in ALL if i not in {F8E8M0FNU}}, |
| 266 | + {(i, F16): (RoundingMode.RP, B134) for i in ALL if i not in {F64, F32, TF32, F8E8M0FNU}}, |
| 267 | + {(i, F8E8M0FNU): (RoundingMode.RP, B133) for i in ALL if i not in {F64, F8E5M2, F8E4M3FN}}, |
| 268 | + {(i, F8E8M0FNU): (RoundingMode.RP, B134) for i in (F64, F8E5M2, F8E4M3FN)}, |
| 269 | + |
| 270 | + {(i, F64): (RoundingMode.RA, B134) for i in ALL}, |
| 271 | + {(i, F32): (RoundingMode.RA, B134) for i in ALL if i not in {F64, F8E8M0FNU}}, |
| 272 | + {(i, TF32): (RoundingMode.RA, B134) for i in ALL if i not in {F64, F8E8M0FNU}}, |
| 273 | + {(i, F16): (RoundingMode.RA, B134) for i in ALL if i not in {F64, F32, TF32, F8E8M0FNU}} |
| 274 | +) |
| 275 | + |
| 276 | +# {(from, to): {RoundingMode_1: BC_Version, RoundingMode_2: BC_Version}} |
| 277 | +FTOF_ROUNDING_REGISTRY = defaultdict(dict) |
| 278 | +for row in _FTOF_ROUNDING_ROWS: |
| 279 | + for from_to, (mode, version) in row.items(): |
| 280 | + FTOF_ROUNDING_REGISTRY[from_to][mode] = version |
| 281 | + |
| 282 | + |
| 283 | +def get_ftof_rounding_min_version(from_dtype: datatype.DType, to_dtype: datatype.DType, |
| 284 | + rounding_mode: RoundingMode | None |
| 285 | + ) -> tuple[RoundingMode, BytecodeVersion | None]: |
| 286 | + |
| 287 | + conversion_pair = (from_dtype, to_dtype) |
| 288 | + supported = FTOF_ROUNDING_REGISTRY.get(conversion_pair, None) |
| 289 | + if supported is None: |
| 290 | + raise TileTypeError(f"float conversion from {from_dtype} to {to_dtype} " |
| 291 | + "is not supported") |
| 292 | + |
| 293 | + rounding_mode = RoundingMode.RN if rounding_mode is None else rounding_mode |
| 294 | + if rounding_mode not in supported: |
| 295 | + raise TileTypeError( |
| 296 | + f"rounding_mode={rounding_mode} is not supported " |
| 297 | + f"for conversion from {from_dtype} to {to_dtype}, " |
| 298 | + f"supported rounding modes for this conversion are {tuple(supported.keys())}") |
| 299 | + |
| 300 | + return (rounding_mode, supported[rounding_mode]) |
| 301 | + |
| 302 | + |
228 | 303 | class CompareOrdering(Enum): |
229 | 304 | ORDERED = "ordered" |
230 | 305 | UNORDERED = "unordered" |
|
0 commit comments