-
Notifications
You must be signed in to change notification settings - Fork 21
Add sys.modules registration #105
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -5,9 +5,11 @@ | |
| import json | ||
| import jax | ||
| import jax.numpy as jnp | ||
| import numpy as np | ||
| import importlib | ||
| import importlib.util | ||
| import os | ||
| import sys | ||
| import traceback | ||
| import xprof_utils | ||
|
|
||
|
|
@@ -17,7 +19,16 @@ def load_module_from_path(module_name, file_path): | |
| if spec is None or spec.loader is None: | ||
| raise ImportError(f"Could not load {module_name} from {file_path}") | ||
| module = importlib.util.module_from_spec(spec) | ||
| spec.loader.exec_module(module) | ||
| # Register before exec_module: modules using `from __future__ import | ||
| # annotations` turn every annotation into a string, and dataclasses then | ||
| # resolves them via sys.modules[cls.__module__], which raises | ||
| # AttributeError on None if the module was never registered. | ||
| sys.modules[module_name] = module | ||
| try: | ||
| spec.loader.exec_module(module) | ||
| except BaseException: | ||
| sys.modules.pop(module_name, None) | ||
| raise | ||
| return module | ||
|
|
||
|
|
||
|
|
@@ -98,6 +109,44 @@ def run_xprof(): | |
| return avg_wall_time, xprof_time | ||
|
|
||
|
|
||
| def diff_metrics(b, o, chunk_elems=1 << 24): | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Is this function added because the code can run into OOM during the diff calculation? It might be beneficial if we can make this comment easier to understand.
Collaborator
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Config 3 (csa_decode_bs512) has a uint8 cache of (32769, 256, 4, 128) = 4 GiB. b and o are host numpy arrays by this point, so (b - o) / b is a numpy true division → promotes to float64 on the host. Then jnp.abs(...) ships it to the TPU as f32: a 16 GiB argument plus a 16 GiB result = 32.00 GiB against a 31.25 GiB chip, overflowing by 773 M. |
||
| \"\"\"Max absolute and max relative difference between two outputs. | ||
|
|
||
| Computed on the host in bounded-size chunks. `b` and `o` have already been | ||
| pulled off the device by `jax.device_get`, so the naive | ||
| `jnp.max(jnp.abs((b - o) / b))` ships them straight back: the true division | ||
| promotes to float, and for a large integer output (e.g. a 4 GiB uint8 paged | ||
| KV cache) that is a 16 GiB argument plus a 16 GiB result, which does not fit | ||
| in HBM. These are diagnostic metrics only -- `is_correct` comes from | ||
| `jnp.allclose` -- but raising here used to fail the whole case. | ||
|
|
||
| Chunking also fixes two latent issues with the old expression: the | ||
| subtraction no longer wraps around for unsigned dtypes, and a NaN in one | ||
| chunk no longer suppresses the maximum found in the others. | ||
| \"\"\" | ||
| fb = np.asarray(b).reshape(-1) | ||
| fo = np.asarray(o).reshape(-1) | ||
| max_abs = 0.0 | ||
| max_rel = 0.0 | ||
| for i in range(0, fb.size, chunk_elems): | ||
| x = fb[i:i + chunk_elems].astype(np.float64) | ||
| y = fo[i:i + chunk_elems].astype(np.float64) | ||
| d = np.abs(x - y) | ||
| # max() over an empty slice is undefined; size is never 0 here but guard | ||
| # anyway so a zero-sized output cannot take down the comparison. | ||
| if d.size == 0: | ||
| continue | ||
| max_abs = max(max_abs, float(np.max(d))) | ||
| with np.errstate(divide="ignore", invalid="ignore"): | ||
| r = d / np.abs(x) | ||
| # Matches the previous definition |(b - o) / b|: division by a zero | ||
| # reference stays +inf, and 0/0 stays NaN rather than being counted. | ||
| r = r[~np.isnan(r)] | ||
| if r.size: | ||
| max_rel = max(max_rel, float(np.max(r))) | ||
| return max_abs, max_rel | ||
|
|
||
|
|
||
| def main(): | ||
| try: | ||
| # Load task configuration from task.json | ||
|
|
@@ -266,8 +315,9 @@ def main(): | |
| if b.shape != o.shape: | ||
| raise ValueError(f"Shape mismatch: {b.shape} vs {o.shape}") | ||
| is_correct = is_correct and bool(jnp.allclose(b, o, atol=curr_atol, rtol=curr_rtol)) | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. If the diff can make OOM, why this line will not make OOM?
Collaborator
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. OOM happens in abs calculation. allclose survives because it's jitted and XLA fuses the promotion |
||
| max_abs_diff = max(max_abs_diff, float(jnp.max(jnp.abs(b - o)))) | ||
| max_rel_diff = max(max_rel_diff, float(jnp.max(jnp.abs((b - o) / b)))) | ||
| leaf_abs, leaf_rel = diff_metrics(b, o) | ||
| max_abs_diff = max(max_abs_diff, leaf_abs) | ||
| max_rel_diff = max(max_rel_diff, leaf_rel) | ||
| except Exception as e: | ||
| harness_logs.append(f"Correctness check failed for input {idx}: {e}") | ||
| is_correct = False | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
If
module_namewas already present insys.modulesbefore callingload_module_from_path, unconditionally popping it on failure will delete the pre-existing module instead of restoring it. It is safer to capture the previous value ofsys.modules[module_name]and restore it if execution fails.There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
If it gets to Exception, it will fail anyway.