diff --git a/README.md b/README.md index 01e6809..5481c83 100644 --- a/README.md +++ b/README.md @@ -13,7 +13,7 @@ with module.spawn(100_000) as instances: The runtime is implemented in this repository: ```text -Python → wgpu-py → WGSL bytecode interpreter → wgpu-native → Metal / Vulkan / DX12 +Python → WASM validation / specialized WGSL → wgpu-py → Metal / Vulkan / DX12 ``` Python parses and validates the module, uploads initial state, dispatches work, @@ -21,6 +21,55 @@ and reads results. **All guest instructions and WASI services run on the GPU.** There is no CPU interpreter, CPU fallback, host WASI service loop, or dependency on Wasmtime in the installed package. Software GPU adapters are rejected. +`execution="auto"` (the default) and `execution="compiled"` compile every +supported function. The GPU program buffer contains metadata, jump tables and +embedded data, **no instruction bytecode**. Execution shaders do not contain an +opcode decoder or interpreter. `execution="interpreter"` explicitly selects the +original GPU interpreter for differential testing; it is never an implicit fallback. +A compilation failure is reported to the caller. + +The compiler divides the complete control-flow graph into bounded units, including +functions larger than one unit. Native pipelines are compiled lazily when their +continuations are reached. Limits of 256 lowered instructions and 64 KiB of +generated source apply per unit, not per module; source overflow splits a unit +instead of dropping functions. Large `br_table` targets are ordinary jump-table +data. The driver compiles generated WGSL into native GPU code through wgpu. +Python 3.8 / wgpu-py 0.18 uses a 16 KiB source cap per unit to reduce pressure +on its older Metal compiler. + +Basic blocks keep intermediate values and modified locals in shader variables, +spilling at continuations or traps. Calls and recursion use explicit GPU frames; +there is no WGSL recursion. A dispatch loop selects **compiled block addresses**, +never opcodes. When fuel or stack space cannot accommodate a whole block, an +unrolled, statically compiled prefix preserves the exact trap and prior effects. +It does not use the reference interpreter. Numeric implementations include software +f64 and checked word loads. WASI services use a separate shared GPU pipeline, +preserving the same embedded filesystem and operand stack. Python schedules +pipelines without executing guest instructions or WASI operations. + +This architecture removes instruction interpretation, but transitions between +compilation units require GPU scheduling and synchronization. Lazy native +compilation can dominate a first invocation; measurements must separate it from +execution. It does not establish a speedup over Wasmtime on CPython workloads. +Use the independent development watchdog below for GPU experiments. Source-size +limits alone do not guarantee driver memory usage or compilation time. + +```python +module = wasmgpu.Module("worker.wasm", execution="compiled") +with module.spawn(1000) as instances: + results = instances.call("process", inputs) + print(instances.last_call) # phase timings and compiled/interpreted counts + instances.reset() # fresh guest state, same allocated GPU buffers +``` + +`compile_functions=[...]` optionally prioritizes function indices when grouping +units; every other function remains compiled. `compile_limit` (1–256) may lower +the per-unit instruction budget. Parsed modules, generated code and GPU +pipelines have bounded in-process caches. Instances share an immutable initial +template; batched heap copies and per-instance random nonces are initialized on +the GPU. `reset()` restores memory, globals, tables, files and descriptors and +reruns the core WASM start function. It does not call a WASI command's `_start`. + This is an experimental engine with the supported profile below, **not a complete implementation of every WebAssembly proposal or a security sandbox**. Apple M4 / Metal has been tested. Vulkan and DX12 use the same shader but have @@ -78,6 +127,14 @@ transactions. A later call can reuse the instances. `proc_exit` raises a trap and records the per-instance code in `instances.exit_codes`. Explicitly close instances, preferably with a context manager, to release GPU buffers. +`call(..., cancel=callback)` checks cancellation between GPU dispatches and raises +`InterruptedError` if the callback returns true. Already completed effects persist. +It cannot preempt a GPU command already submitted to the driver. `last_call` +records lazy WGSL generation and native compilation, input preparation, upload, +dispatch/completion synchronization, result readback and decoding separately, +including counters up to cancellation. In compiled mode, +`last_call.interpreted_instructions` is always zero. + ## Embedded WASI Preview 1 Files are **byte contents embedded into each instance's GPU filesystem**. @@ -118,7 +175,7 @@ Other operating-system services have explicit virtual semantics: - Arguments/environment come from the embedded configuration. - All clocks use a per-instance virtual counter. It starts at `clock_epoch_ns` - (default 0) and advances by `clock_resolution_ns` (default 1) per interpreter + (default 0) and advances by `clock_resolution_ns` (default 1) per lowered WASM instruction. It never reads the host clock. - `poll_oneoff` reports ready virtual descriptors or advances virtual time to the earliest clock deadline, without sleeping on the host. @@ -159,8 +216,8 @@ The current memory32 addressing implementation caps memory at 65,535 pages. | `table_elements` | up to 256, at least each declared minimum | Growth capacity per table, capped by each declared maximum | | `stack_size` | 256 | 64-bit value slots per instance, including locals | | `call_depth` | 64 | Nested call frames per instance | -| `fuel` | 10,000,000 | Invocation budget; bulk work also consumes fuel | -| `quantum` | 4096 | Interpreter instructions per dispatch before resumption | +| `fuel` | 10,000,000 | Positive u64 invocation budget; bulk work also consumes fuel | +| `quantum` | 4096 | Target instruction count per dispatch; compiled blocks may exceed it by at most 31 | | `batch_size` | device-derived | Instances per dispatch/buffer group | | `max_resident_bytes` | 512 MiB | Total resident buffer allocation budget | @@ -170,6 +227,7 @@ are practical, but 100,000 workers with 1 MiB private memory require about 100 G before stacks/files. Excessive allocations fail before allocation. `resident_bytes` and `adapter_info` expose the allocation estimate and selected hardware. +Compiled dispatches yield at basic-block boundaries; fuel remains exact. A dispatch quantum is not a real-time deadline. Large individual bulk operations and WASI operations can take longer than a scalar instruction. Tune resource budgets for trusted workloads; this runtime is not suitable for hostile modules. @@ -177,20 +235,25 @@ budgets for trusted workloads; this runtime is not suitable for hostile modules. ## Verification and benchmarks ```sh -venv/bin/python -m pytest tests -q # hardware GPU required; absence fails +venv/bin/python -m tests.gpu_guard --seconds 600 -- -m pytest tests -q venv/bin/python -m pytest tests -m 'not gpu' # parser/configuration/failure tests only -WASMGPU_COVERAGE_BRANCH=true venv/bin/python -m coverage run -m pytest tests -q +venv/bin/python -m tests.gpu_guard --seconds 600 -- -m pytest tests -q --wasmgpu-execution interpreter +WASMGPU_COVERAGE_BRANCH=true venv/bin/python -m tests.gpu_guard --seconds 600 -- -m coverage run -m pytest tests -q venv/bin/python -m coverage combine venv/bin/python -m coverage report -m venv/bin/python -m tests.benchmark # run alone, without concurrent GPU jobs venv/bin/python -m build ``` -All 358 tests passed on Apple M4 / Metal with Python 3.11, including 40 official -spec suites plus API, numeric, WASI and compiled-code cases. This run covered -100% of Python statements and branches; that percentage does not measure WGSL. -Python 3.8/3.9/3.10 each passed 315 cases with their pinned GPU backend; Python 3.11 -also runs every official suite. Python 3.15rc2 passed the checks that require no GPU. +The GPU suites run on Apple M4 / Metal with Python 3.11. The interpreter was also +tested with Python 3.8/3.9/3.10 and their pinned GPU backends; native compilation is additionally checked on Python 3.8 / wgpu-py 0.18. +The full suite on older backends still requires verification. `--wasmgpu-execution interpreter` +reruns the existing suites against the reference GPU interpreter. Additional differential +tests compare compiled and interpreted continuations, precise fuel exhaustion, +WASI clocks, cancellation and reset. Large functions, cross-pipeline recursion, +exported imports and large jump tables are covered explicitly. Tests forbid use +of the interpreter and inspect shader sources and the uploaded program format. +Python coverage does not measure WGSL. Tests compare scalar operations with Wasmtime, exercise 100,003 concurrent instances, and run actual compiled C and Rust fixtures with allocation, internal calls, f64, libc/Rust formatting and embedded files. Forty unmodified official @@ -217,3 +280,43 @@ as configured for ordinary hosted runners. The Linux / Python 3.11 job uploads i coverage report to Coveralls with the `python-only` flag and saves it as a GitHub artifact. This report does not claim GPU tests have run. Full conformance and performance tests must be run locally with a hardware GPU. + +For real CPython WASI and the startup/pyflakes/mypy guest workloads used by +throng, use [benchmark_cpython.py](tests/benchmark_cpython.py). Supply a runtime +directory containing `python.wasm` and `lib/`; linter scenarios also take `--wheels` +with pure-Python wheels (pyflakes 3.3.2, mypy 1.14.1 and dependencies). The benchmark +embeds these bytes for GPU execution. Only the Wasmtime test oracle mounts host +files. The guest scans `project/` instead of `.` to exclude its embedded runtime; +throng scheduling, snapshotting and installation overhead are excluded. + +```sh +venv/bin/python -m tests.benchmark_cpython --runtime /path/to/runtime --engine wasmtime --scenario startup --output /tmp/cpu.json +venv/bin/python -m tests.gpu_guard --seconds 180 -- -m tests.benchmark_cpython --runtime /path/to/runtime --engine compiled --scenario startup --reference /tmp/cpu.json --output /tmp/gpu.json +``` + +Use `--engine interpreter` for the GPU reference, `--scenario pyflakes --wheels +/path/to/wheels --files 1` for a linter workload, and `--files 100` for its larger +corpus. `--reference` checks the runtime hash, arguments, exit code and captured +streams. Failed or interrupted runs return a nonzero status and must not be +counted as completed performance measurements. Reports are opt-in local files, +not repository artifacts. Each invocation is one observation, not a statistical +speedup claim; repeat under the same conditions before drawing conclusions. + +`--count-fuel` instruments Wasmtime to estimate a workload's size; its fuel units +are not identical to this engine's lowered-instruction accounting, and the +instrumentation changes CPU timings. Large linter workloads require explicit +`--fuel` and `--call-seconds` settings as well as a matching watchdog deadline. +Fuel and execution counters are u64 on the GPU; long invocations do not wrap at +four billion instructions. Initial compilation, cached module loading and reset timings are reported +separately. Lazy compilation during a call is reported in `metrics.codegen_seconds` +and `metrics.compile_seconds`; `call_seconds` includes these phases, while +`metrics.execute_seconds` excludes them. + +The watchdog runs separately from the worker. On macOS it accounts for physical +memory of the worker and its own Metal compiler services, checks compilation and +experiment deadlines, and terminates only those processes on a limit violation. +Defaults are 1.5 GiB and 30 seconds per pipeline compilation. The benchmark also +has a guest fuel budget and a call deadline. A driver cache hit can make a new +process's pipeline creation faster; reported compilation times do not imply a +cold Metal driver cache. Dispatch timings include completion-state readback; +the output readback phase is measured separately. diff --git a/pyproject.toml b/pyproject.toml index aa0c63e..2d3224c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,11 +4,11 @@ build-backend = "setuptools.build_meta" [project] name = "wasmgpu" -version = "0.0.1" +version = "0.0.2" authors = [ { name="Evgeniy Blinov", email="zheni-b@yandex.ru" }, ] -description = 'Batched WebAssembly execution in a GPU compute-shader interpreter' +description = 'Batched WebAssembly compilation and execution on hardware GPUs' readme = "README.md" requires-python = ">=3.8" dependencies = [ diff --git a/tests/benchmark_cpython.py b/tests/benchmark_cpython.py new file mode 100644 index 0000000..6c74fd3 --- /dev/null +++ b/tests/benchmark_cpython.py @@ -0,0 +1,220 @@ +"""Real CPython WASI and throng guest workloads; no scheduler overhead included. + +Supply a runtime directory containing python.wasm and lib/, and (for linters) +a directory of pure-Python wheels. Downloads and input preparation are untimed. +Run GPU variants through tests.gpu_guard. Results are printed, never checked in. +""" +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import platform +import tempfile +import time +from dataclasses import asdict +from importlib.metadata import version +from pathlib import Path +from zipfile import ZipFile + +import wasmtime + +import wasmgpu +from wasmgpu.binary import Reader + +SOURCE = b'from typing import Iterable\n\ndef total(values: Iterable[int]) -> int:\n return sum(values)\n\nanswer: int = total([1, 2, 3])\n' + + +def function_names(data): + reader = Reader(data) + reader.take(8) + names = {} + while reader.pos < len(reader.data): + section_id = reader.byte() + section = Reader(reader.take(reader.leb())) + if section_id != 0 or section.name() != 'name': + continue + while section.pos < len(section.data): + subsection_id = section.byte() + subsection = Reader(section.take(section.leb())) + if subsection_id == 1: + for _ in range(subsection.leb()): + index = subsection.leb() + names[index] = subsection.name() + return names + + +def workload(runtime, scenario, wheels, count): + files = {'python/' + path.relative_to(runtime).as_posix(): path.read_bytes() + for path in sorted((runtime / 'lib').rglob('*')) if path.is_file()} + env = {'PYTHONHOME': '/python', 'PYTHONDONTWRITEBYTECODE': '1', 'PYTHONHASHSEED': '0'} + if scenario == 'smoke': + args = ['-S', '-c', 'print(6 * 7)'] + elif scenario == 'startup': + args = ['-c', 'pass'] + else: + if not wheels or not list(wheels.glob('*.whl')): + raise ValueError('linter scenarios require --wheels containing pure-Python wheels') + for wheel in sorted(wheels.glob('*.whl')): + with ZipFile(wheel) as archive: + for entry in archive.infolist(): + if not entry.is_dir(): + files['packages/' + entry.filename] = archive.read(entry) + env['PYTHONPATH'] = '/packages' + for index in range(count): + files[f'project/module_{index:03}.py'] = SOURCE + # Only the project is scanned; stdlib and wheels share the embedded root. + args = ['-m', scenario] + if scenario == 'mypy': + args += ['--no-incremental', '--cache-dir=/dev/null', '--no-site-packages', '--python-version=3.13', '--platform=linux'] + args += ['project'] + return files, ['python', *args], env + + +def cpu(data, files, args, env, *, count_fuel=False): + with tempfile.TemporaryDirectory(prefix='wasmgpu-oracle-') as directory: + root = Path(directory) + for name, content in files.items(): + path = root / name + path.parent.mkdir(parents=True, exist_ok=True) + path.write_bytes(content) + started = time.perf_counter() + engine_config = wasmtime.Config() + engine_config.consume_fuel = count_fuel + engine = wasmtime.Engine(engine_config) + module = wasmtime.Module(engine, data) + compilation = time.perf_counter() - started + started = time.perf_counter() + store = wasmtime.Store(engine) + if count_fuel: + store.set_fuel(1 << 60) + config = wasmtime.WasiConfig() + config.argv, config.env = args, list(env.items()) + config.preopen_dir(directory, '.') + config.stdout_file, config.stderr_file = str(root / 'stdout'), str(root / 'stderr') + store.set_wasi(config) + linker = wasmtime.Linker(engine) + linker.define_wasi() + instance = linker.instantiate(store, module) + preparation = time.perf_counter() - started + started = time.perf_counter() + code = 0 + try: + instance.exports(store)['_start'](store) + except wasmtime.ExitTrap as error: + code = error.code + execution = time.perf_counter() - started + return {'compile_seconds': compilation, 'prepare_seconds': preparation, 'execute_seconds': execution, + 'wasmtime_fuel': (1 << 60) - store.get_fuel() if count_fuel else None, + 'exit_code': code, 'stdout': (root / 'stdout').read_text(), 'stderr': (root / 'stderr').read_text()} + + +def gpu(data, files, args, env, *, options, names): # noqa: PLR0913 - Workload and measurement configuration. + if not os.environ.get('WASMGPU_GUARD_STATUS'): + raise RuntimeError('run this development benchmark through python -m tests.gpu_guard') + selected = [index for index, name in names.items() if name in options.functions] + if options.engine == 'compiled' and len(selected) != len(options.functions): + raise ValueError('selected function name missing or ambiguous in this runtime') + module = wasmgpu.Module(data, execution=options.engine, compile_functions=selected) + plan = module.compiled + print(json.dumps({'stage': 'planned', 'functions': len(plan.functions) if plan else 0, 'instructions': plan.instructions if plan else 0, # noqa: T201 - Benchmark CLI. + 'compilation_units': len(plan.regions) if plan else 0}), flush=True) + wasi = wasmgpu.Wasi(args=args, env=env, files=files, storage_size=48 * 1024 * 1024, max_files=4096) + started = time.perf_counter() + with module.spawn(1, memory_pages=1024, stack_size=8192, call_depth=1024, fuel=options.fuel, + quantum=262144, wasi=wasi) as instances: + preparation = time.perf_counter() - started - instances.pipeline_seconds - instances.codegen_seconds + print(json.dumps({'stage': 'pipeline', 'seconds': instances.pipeline_seconds}), flush=True) # noqa: T201 + started = time.perf_counter() + error = None + reported = started + + def cancelled(): + nonlocal reported + now = time.perf_counter() + if now - reported >= 15: + state = instances._batches[0].state + print(json.dumps({'stage': 'execute', 'seconds': round(now - started, 2), # noqa: T201 - Bounded-run progress. + 'fuel_used': options.fuel - state[9] - (state[12] << 32), + 'function': names.get(state[3], str(state[3])), + 'pipelines_created': instances.last_call.pipelines_created, + 'compile_seconds': round(instances.last_call.compile_seconds, 2), + 'execute_seconds': round(instances.last_call.execute_seconds, 2)}), flush=True) + reported = now + return now - started > options.call_seconds + try: + instances.call('_start', cancel=cancelled) + except wasmgpu.Trap as trap: + if any(value != 'WASI proc_exit' for value in trap.traps.values()): + error = str(trap) + except InterruptedError as interrupted: + error = str(interrupted) + wall = time.perf_counter() - started + metrics = asdict(instances.last_call) + metrics['function_samples'] = {names.get(index, str(index)): count for index, count in + sorted(instances.last_call.function_samples.items(), key=lambda item: -item[1])[:20]} + started = time.perf_counter() + stdout, stderr = instances.stdout[0].decode(), instances.stderr[0].decode() + stream_readback = time.perf_counter() - started + exit_code = (instances.exit_codes[0] or 0) if error is None else None + cached = wasmgpu.Module(data, execution=options.engine, compile_functions=selected) + started = time.perf_counter() + instances.reset() + reset_seconds = time.perf_counter() - started + return {'load_seconds': module.load_seconds, 'codegen_seconds': module.codegen_seconds + instances.codegen_seconds, + 'cached_load_seconds': cached.load_seconds, 'cached_codegen_seconds': cached.codegen_seconds, + 'reset_seconds': reset_seconds, + 'compiled_functions': [names.get(index, str(index)) for index in plan.functions[:32]] if plan else [], + 'compiled_function_count': len(plan.functions) if plan else 0, + 'compilation_units': len(plan.regions) if plan else 0, + 'compiled_static_instructions': plan.instructions if plan else 0, + 'compile_seconds': instances.pipeline_seconds, 'pipeline_cache_hit': instances.pipeline_cache_hit, + 'prepare_seconds': preparation, 'call_seconds': wall, 'stream_readback_seconds': stream_readback, + 'metrics': metrics, 'exit_code': exit_code, 'error': error, + 'stdout': stdout, 'stderr': stderr, 'adapter': instances.adapter_info} + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--runtime', required=True, type=Path) + parser.add_argument('--engine', required=True, choices=['wasmtime', 'interpreter', 'compiled']) + parser.add_argument('--scenario', default='smoke', choices=['smoke', 'startup', 'pyflakes', 'mypy']) + parser.add_argument('--wheels', type=Path) + parser.add_argument('--files', type=int, default=1, choices=[1, 100]) + parser.add_argument('--functions', nargs='+', default=['_PyCode_Quicken', '_PyPegen_is_memoized', 'strlen', 'visit_decref']) + parser.add_argument('--fuel', type=int, default=300_000_000) + parser.add_argument('--call-seconds', type=float, default=120) + parser.add_argument('--output', type=Path) + parser.add_argument('--reference', type=Path, help='assert observable results against a previous engine report') + parser.add_argument('--count-fuel', action='store_true', help='instrument Wasmtime for a separate work-volume estimate; changes CPU timings') + options = parser.parse_args() + root = Path(wasmgpu.__file__).parent + source_digest = hashlib.sha256() + for path in sorted([*root.glob('*.py'), *root.glob('*.wgsl')]): + source_digest.update(path.name.encode() + b'\0' + path.read_bytes()) + data = (options.runtime / 'python.wasm').read_bytes() + files, args, env = workload(options.runtime, options.scenario, options.wheels, options.files) + result = (cpu(data, files, args, env, count_fuel=options.count_fuel) if options.engine == 'wasmtime' else + gpu(data, files, args, env, options=options, names=function_names(data))) + result.update(engine=options.engine, scenario=options.scenario, files=options.files, + runtime_sha256=hashlib.sha256(data).hexdigest(), args=args, env=env, + engine_source_sha256=source_digest.hexdigest(), python=platform.python_version(), + system=platform.platform(), wgpu=version('wgpu'), wasmtime=version('wasmtime')) + digest = hashlib.sha256() + for name, content in sorted(files.items()): + digest.update(name.encode() + b'\0' + len(content).to_bytes(8, 'little') + content) + result['workload_sha256'] = digest.hexdigest() + if options.reference: + reference = json.loads(options.reference.read_text()) + keys = ('runtime_sha256', 'workload_sha256', 'scenario', 'files', 'args', 'env', 'exit_code', 'stdout', 'stderr') + result['matches_reference'] = all(result[key] == reference.get(key) for key in keys) and result.get('error') is None + output = json.dumps(result, indent=2) + print(output, flush=True) # noqa: T201 + if options.output: + options.output.write_text(output + '\n') + return int(result.get('error') is not None or result.get('matches_reference') is False) + + +if __name__ == '__main__': + raise SystemExit(main()) diff --git a/tests/conftest.py b/tests/conftest.py index 0bbf627..c2e8ab8 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -8,6 +8,24 @@ import wasmgpu +def pytest_addoption(parser): + parser.addoption('--wasmgpu-execution', choices=['auto', 'compiled', 'interpreter'], default='auto', + help='Default backend for the existing GPU conformance suite') + + +@pytest.fixture(scope='session', autouse=True) +def execution_mode(request): + mode = request.config.getoption('--wasmgpu-execution') + original = wasmgpu.Module.__init__ + if mode != 'auto': + def configured(self, *args, **kwargs): + kwargs.setdefault('execution', mode) + original(self, *args, **kwargs) + wasmgpu.Module.__init__ = configured + yield + wasmgpu.Module.__init__ = original + + @pytest.fixture(scope='session') def engine(): return wasmtime.Engine() diff --git a/tests/gpu_guard.py b/tests/gpu_guard.py new file mode 100644 index 0000000..8728b4c --- /dev/null +++ b/tests/gpu_guard.py @@ -0,0 +1,149 @@ +"""Run a Python experiment with an independent watchdog (no GPU in this process). + +Example: venv/bin/python -m tests.gpu_guard -- -m pytest tests/test_codegen.py +On macOS, account for the worker's Metal XPC compiler, including compressed +memory. Only services belonging to the worker's launchd domain may be killed. +This bounds development experiments; it cannot prevent a GPU driver panic. +""" +from __future__ import annotations + +import argparse +import ctypes +import json +import os +import re +import signal +import subprocess +import sys +import tempfile +import time +from pathlib import Path +from typing import ClassVar + + +class Usage(ctypes.Structure): + _fields_: ClassVar[list[tuple[str, type]]] = [('uuid', ctypes.c_uint8 * 16), *[ + (name, ctypes.c_uint64) for name in ( + 'user', 'system', 'idle', 'interrupts', 'pageins', 'wired', + 'resident', 'footprint', 'start', 'exit', + ) + ]] + + +def memory(pid): + """Physical footprint and process identity, or None after process exit.""" + if sys.platform == 'darwin': + usage = Usage() + library = ctypes.CDLL('/usr/lib/libproc.dylib') + if library.proc_pid_rusage(pid, 0, ctypes.byref(usage)) == 0: + return max(usage.footprint, usage.resident), usage.start + return None + try: + status = Path('/proc/{}/status'.format(pid)).read_text() + resident = re.search(r'^VmRSS:\s+(\d+)', status, re.MULTILINE) + swapped = re.search(r'^VmSwap:\s+(\d+)', status, re.MULTILINE) + identity = Path('/proc/{}/stat'.format(pid)).read_text().rsplit(')', 1)[1].split()[19] + return (int(resident[1]) + (int(swapped[1]) if swapped else 0)) * 1024, identity + except (OSError, TypeError): + return None + + +def compiler_pids(domain): + services = re.search(r'^\s*services = \{(.*?)^\s*\}', domain, re.MULTILINE | re.DOTALL) + if services is None: + return set() + return {int(match[1]) for match in re.finditer( + r'^\s*(\d+)\s+[^\n]*\scom\.apple\.MTLCompilerService(?:\.[\w-]+)?\s*$', + services[1], re.MULTILINE, + ) if int(match[1]) > 0} + + +def owned_compilers(pid): + if sys.platform != 'darwin': + return set() + result = subprocess.run(['launchctl', 'print', 'pid/{}'.format(pid)], capture_output=True, text=True, timeout=1, check=False) + return compiler_pids(result.stdout) + + +def terminate(worker, compilers): + try: + os.killpg(worker.pid, signal.SIGKILL) + except ProcessLookupError: + pass + for pid, identity in compilers.items(): + current = memory(pid) + if current is not None and current[1] == identity: + try: + os.kill(pid, signal.SIGKILL) + except ProcessLookupError: + pass + worker.wait(timeout=5) + + +def run(command, *, seconds=180, memory_mib=1536, compile_seconds=30, interval=0.1): + if sys.platform not in ('darwin', 'linux'): + raise RuntimeError('development watchdog currently supports macOS and Linux') + started = time.monotonic() + peak = 0 + worker_peak = compiler_peak = 0 + compilers = {} + reason = None + with tempfile.TemporaryDirectory(prefix='wasmgpu-guard-') as directory: + status = Path(directory) / 'phase.json' + env = dict(os.environ, WASMGPU_GUARD_STATUS=str(status), PYTHONUNBUFFERED='1') + worker = subprocess.Popen([sys.executable, *command], env=env, start_new_session=True) + try: + while worker.poll() is None: + for pid in owned_compilers(worker.pid): + usage = memory(pid) + if usage is not None: + compilers[pid] = usage[1] + usage = memory(worker.pid) + total = usage[0] if usage else 0 + worker_peak = max(worker_peak, total) + compiler_total = 0 + for pid, identity in compilers.items(): + usage = memory(pid) + if usage is not None and usage[1] == identity: + total += usage[0] + compiler_total += usage[0] + compiler_peak = max(compiler_peak, compiler_total) + peak = max(peak, total) + now = time.monotonic() + if total > memory_mib * 1024 * 1024: + reason = 'memory limit' + elif now - started > seconds: + reason = 'experiment timeout' + if status.exists(): + phase = json.loads(status.read_text()) + if phase['compiling'] and now - phase['started'] > compile_seconds: + reason = 'shader compilation timeout' + if reason: + terminate(worker, compilers) + break + time.sleep(interval) + finally: + if worker.poll() is None: + terminate(worker, compilers) + report = {'guard': reason or 'completed', 'peak_mib': round(peak / 1048576, 1), + 'worker_peak_mib': round(worker_peak / 1048576, 1), 'compiler_peak_mib': round(compiler_peak / 1048576, 1), + 'seconds': round(time.monotonic() - started, 3), 'metal_compilers': len(compilers)} + print(json.dumps(report), file=sys.stderr) # noqa: T201 - CLI diagnostics. + return 124 if reason else worker.returncode + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--seconds', type=float, default=180) + parser.add_argument('--memory-mib', type=int, default=1536) + parser.add_argument('--compile-seconds', type=float, default=30) + parser.add_argument('command', nargs=argparse.REMAINDER) + args = parser.parse_args() + command = args.command[1:] if args.command[:1] == ['--'] else args.command + if not command or min(args.seconds, args.memory_mib, args.compile_seconds) <= 0: + parser.error('a Python command and positive limits are required') + return run(command, seconds=args.seconds, memory_mib=args.memory_mib, compile_seconds=args.compile_seconds) + + +if __name__ == '__main__': + raise SystemExit(main()) diff --git a/tests/test_codegen.py b/tests/test_codegen.py new file mode 100644 index 0000000..58f8f5d --- /dev/null +++ b/tests/test_codegen.py @@ -0,0 +1,457 @@ +from __future__ import annotations + +import gc +import struct +import weakref +from array import array + +import pytest + +import wasmgpu +from tests.conftest import binary, oracle +from wasmgpu import compiler, runtime +from wasmgpu.compiler import MAX_COMPILED_INSTRUCTIONS, MAX_GENERATED_BYTES + +LOOP = '''(module (memory 1 2) (global $g (mut i32) (i32.const 0)) + (func (export "run") (param $n i32) (result i32) + block $done loop $next + local.get $n i32.eqz br_if $done + global.get $g i32.const 1 i32.add global.set $g + i32.const 0 global.get $g i32.store + local.get $n i32.const 1 i32.sub local.set $n br $next + end end global.get $g))''' + + +def test_codegen_limits_and_cache(): + wasm = binary(LOOP) + first = wasmgpu.Module(wasm, execution='compiled') + second = wasmgpu.Module(wasm, execution='compiled') + assert first.compiled is second.compiled + assert first._binary is second._binary + assert first.compiled.instructions > 0 + assert 'let v' in first.compiled.source + assert 'invoke(' not in first.compiled.source + assert wasmgpu.Module(wasm, execution='auto').compiled is first.compiled + assert wasmgpu.Module(wasm, execution='interpreter').compiled is None + with pytest.raises(wasmgpu.ResourceLimitError, match='limited'): + wasmgpu.Module(wasm, execution='compiled', compile_limit=0) + assert wasmgpu.Module(wasm, execution='compiled', compile_functions=[0, 0]).compiled.functions == (0,) + with pytest.raises(wasmgpu.ResourceLimitError, match='limited'): + wasmgpu.Module(wasm, execution='compiled', compile_limit=MAX_COMPILED_INSTRUCTIONS + 1) + with pytest.raises(ValueError, match='function index'): + wasmgpu.Module(wasm, execution='compiled', compile_functions=[1]) + with pytest.raises(ValueError, match='execution'): + wasmgpu.Module(wasm, execution='cpu') + + +def test_source_budget_splits_units_without_losing_functions(monkeypatch): + functions = ' '.join(f'(func (export "f{i}") (param f64) (result f64) local.get 0 f64.sqrt)' for i in range(60)) + module = wasmgpu.Module(binary('(module ' + functions + ')'), execution='interpreter') + monkeypatch.setattr(compiler, 'MAX_GENERATED_BYTES', 32768) + plan = compiler.compile_module(module._binary) + assert len(plan.source.encode()) <= min(32768, MAX_GENERATED_BYTES) + assert plan.complete + assert len(plan.functions) == 60 + index = 0 + while index < len(plan.regions): + assert len(plan.source_for(index).encode()) <= 32768 + index += 1 + assert len(plan.regions) > 1 + + +@pytest.mark.parametrize('index', [True, 1.5, '0', -1]) +def test_invalid_function_selection(index): + with pytest.raises((ValueError, TypeError), match='function index'): + wasmgpu.Module(binary(LOOP), execution='compiled', compile_functions=[index]) + + +def test_evicted_module_is_not_retained_by_codegen_cache(): + def make(value): + return wasmgpu.Module(binary(f'(module (func (result i32) i32.const {value}))'), execution='compiled') + module = make(701) + reference = weakref.ref(module._binary) + del module + assert reference() is not None + for value in range(702, 707): + make(value) + gc.collect() + assert reference() is None + + +def outcome(instances, values, fuel): + try: + result, traps = instances.call('run', values, fuel=fuel), {} + except wasmgpu.Trap as error: + result, traps = error.results, error.traps + state = instances._batches[0].state + # Compare continuation, memory size, trap, fuel and proc_exit, excluding counters. + return result, traps, list(state[:6]) + list(state[7:11]), instances.read_memory(0, 4) + + +@pytest.mark.gpu +@pytest.mark.parametrize('quantum', [1, 3, 17, 4096]) +def test_exact_fuel_prefix(quantum): + modules = [wasmgpu.Module(binary(LOOP), execution=mode) for mode in ('compiled', 'interpreter')] + for fuel in [1, 2, 5, 13, 37, 100]: + with modules[0].spawn(1, quantum=quantum) as compiled, modules[1].spawn(1, quantum=quantum) as interpreted: + assert outcome(compiled, [3], fuel) == outcome(interpreted, [3], fuel) + assert compiled.last_call.compiled_instructions + compiled.last_call.interpreted_instructions == interpreted.last_call.interpreted_instructions + if quantum == 4096 and fuel == 100: + assert compiled.last_call.compiled_instructions > 0 + + +@pytest.mark.gpu +def test_recursion_indirect_and_pipeline_cache(engine): + wat = '''(module (type $f (func (param i32) (result i64))) + (table 1 funcref) (elem (i32.const 0) $f) + (func $f (export "run") (type $f) + local.get 0 i32.eqz if (result i64) i64.const 1 else + local.get 0 i64.extend_i32_u local.get 0 i32.const 1 i32.sub + i32.const 0 call_indirect (type $f) i64.mul end))''' + values = [(n,) for n in range(15)] + module = wasmgpu.Module(binary(wat), execution='compiled') + with module.spawn(len(values)) as instances: + assert instances.call('run', values) == oracle(engine, wat, 'run', values) + assert instances.last_call.interpreted_instructions == 0 + assert instances.last_call.compiled_instructions > 100 + pipeline = instances._pipeline + with module.spawn(1) as instances: + assert instances._pipeline is pipeline + assert instances.pipeline_cache_hit + + +@pytest.mark.gpu +@pytest.mark.parametrize('indirect', [False, True]) +def test_wasi_clock_and_import_continuation(indirect): + call = 'i32.const 0 call_indirect (type $clock)' if indirect else 'call $clock' + wat = f'''(module (type $clock (func (param i32 i64 i32) (result i32))) + (import "wasi_snapshot_preview1" "clock_time_get" (func $clock (type $clock))) + (memory 1) (table 1 funcref) (elem (i32.const 0) $clock) + (func (export "run") (result i64) + i32.const 0 i64.const 0 i32.const 0 {call} drop i32.const 0 i64.load))''' + outputs = [] + for mode in ('compiled', 'interpreter'): + with wasmgpu.Module(binary(wat), execution=mode).spawn(1, wasi=wasmgpu.Wasi(clock_resolution_ns=7)) as instances: + outputs.append(instances.call('run')) + assert instances.read_memory(0, 8) == struct.pack(' 0 + assert instances.last_call.interpreted_instructions == 0 + assert outputs[0] == outputs[1] + + +@pytest.mark.gpu +def test_cancel_retains_metrics_and_effects(): + module = wasmgpu.Module(binary(LOOP), execution='compiled') + with module.spawn(1, quantum=17) as instances: + with pytest.raises(InterruptedError): + instances.call('run', [1000], cancel=lambda: instances.last_call.dispatches == 2) + assert instances.last_call.dispatches == 2 + assert instances.last_call.compiled_instructions > 0 + assert instances.read_memory(0, 4) != b'\0' * 4 + assert instances.call('run', [0])[0] > 0 + with pytest.raises(TypeError, match='cancel'): + instances.call('run', [0], cancel=True) + + +@pytest.mark.gpu +@pytest.mark.parametrize('count', [1, 5]) +def test_reset_reuses_buffers_and_restores_state(count, monkeypatch): + wat = '''(module (memory 1 2) (data (i32.const 1) "init") + (global $g (mut i32) (i32.const 3)) + (table 1 2 funcref) (elem (i32.const 0) $start) + (func $start global.get $g i32.const 4 i32.add global.set $g) (start $start) + (func (export "run") (result i32) global.get $g) + (func (export "mutate") + i32.const 1 memory.grow drop i32.const 1 i32.const 999 i32.store + i32.const 0 ref.null func table.set i32.const 40 global.set $g))''' + with wasmgpu.Module(binary(wat), execution='compiled', files={'a': b'initial'}).spawn(count, batch_size=2) as instances: + buffers = [id(buffer) for batch in instances._batches for buffer in batch.buffers] + assert instances.call('run') == [7] * count + instances.call('mutate') + assert instances.call('run') == [40] * count + + def forbidden(*_args, **_kwargs): + raise AssertionError('reset allocated a new GPU buffer') + monkeypatch.setattr(instances._context, 'buffer', forbidden) + instances.reset() + assert buffers == [id(buffer) for batch in instances._batches for buffer in batch.buffers] + assert instances.call('run') == [7] * count + for index in range(count): + assert instances.read_memory(1, 4, instance=index) == b'init' + assert instances.read_file('a', instance=index) == b'initial' + with pytest.raises(IndexError): + instances.read_memory(65536, 1, instance=index) + instances.call('mutate') + + +@pytest.mark.gpu +def test_all_functions_and_indirect_traps(engine): + wat = '''(module (type $t (func (param i32) (result i32))) (table 3 funcref) + (func $f (type $t) local.get 0 i32.const 7 i32.mul) + (func $wrong (result i64) i64.const 1) + (elem (i32.const 0) $f $wrong) + (func (export "run") (param i32) (result i32) + i32.const 6 local.get 0 call_indirect (type $t)))''' + module = wasmgpu.Module(binary(wat), execution='compiled', compile_functions=[0]) + assert module.compiled.complete + assert oracle(engine, wat, 'run', [(0,)]) == [42] + with module.spawn(4) as instances: + with pytest.raises(wasmgpu.Trap) as error: + instances.call('run', [0, 1, 2, 3]) + assert error.value.results == [42, None, None, None] + assert error.value.traps == {1: 'indirect call type mismatch', 2: 'uninitialized element', 3: 'out of bounds table access'} + assert instances.last_call.compiled_instructions > 4 + assert instances.last_call.interpreted_instructions == 0 + + +@pytest.mark.gpu +@pytest.mark.parametrize(('ty', 'store', 'load', 'value'), [ + ('i32', 'store', 'load', -2147483648), ('i32', 'store8', 'load8_s', -128), + ('i32', 'store16', 'load16_s', -32768), ('i64', 'store', 'load', -(1 << 63)), + ('i64', 'store8', 'load8_s', -128), ('i64', 'store16', 'load16_s', -32768), + ('i64', 'store32', 'load32_s', -2147483648), ('f32', 'store', 'load', 1.25), + ('f64', 'store', 'load', 1e-310), +]) +def test_compiled_word_loads(engine, ty, store, load, value): + wat = f'''(module (memory 1) (func $compiled) + (func (export "run") (param i32 {ty}) (result {ty}) + local.get 0 local.get 1 {ty}.{store} align=1 local.get 0 {ty}.{load} align=1))''' + rows = [(offset, value) for offset in [0, 1, 2, 3, 7, 65528]] + expected = oracle(engine, wat, 'run', rows) + with wasmgpu.Module(binary(wat), execution='compiled', compile_functions=[0]).spawn(len(rows)) as instances: + assert instances.call('run', rows) == expected + assert instances.last_call.compiled_instructions > 0 + assert instances.last_call.interpreted_instructions == 0 + with pytest.raises(wasmgpu.Trap, match='out of bounds memory'): + instances.call('run', [(-1, value)] * len(rows)) + + +@pytest.mark.gpu +@pytest.mark.parametrize('mode', ['compiled', 'interpreter']) +@pytest.mark.parametrize('fuel', [2**32 - 1, 2**32, 2**32 + 3, 2**64 - 1]) +def test_u64_fuel_borrow_and_bulk_charge(mode, fuel): + wat = '''(module (memory 1) + (func (export "run") (result i32) + i32.const 0 i32.const 7 i32.const 16 memory.fill i32.const 0 i32.load))''' + with wasmgpu.Module(binary(wat), execution=mode).spawn(2, fuel=fuel) as instances: + assert instances.call('run') == [0x07070707] * 2 + state = instances._batches[0].state + assert [state[base + 9] + (state[base + 12] << 32) for base in (0, 16)] == [fuel - 23] * 2 + assert instances.last_call.compiled_instructions + instances.last_call.interpreted_instructions == 14 + with pytest.raises(wasmgpu.Trap, match='fuel exhausted'): + instances.call('run', fuel=4) + with pytest.raises(ValueError, match='fuel'): + instances.call('run', fuel=2**64) + + +def test_u64_fuel_overflow_rejected_before_gpu(monkeypatch): + def forbidden(): + raise AssertionError('invalid fuel reached the GPU') + monkeypatch.setattr(runtime, '_context', forbidden) + with pytest.raises(ValueError, match='fuel'): + wasmgpu.Module(binary(LOOP), execution='interpreter').spawn(1, fuel=2**64) + + +@pytest.mark.gpu +def test_hardlink_readback_and_filesystem_reset(): + wat = r'''(module + (import "wasi_snapshot_preview1" "path_link" (func $link (param i32 i32 i32 i32 i32 i32 i32) (result i32))) + (import "wasi_snapshot_preview1" "fd_write" (func $write (param i32 i32 i32 i32) (result i32))) + (memory 1) (data (i32.const 0) "aalias") + (data (i32.const 8) "\10\00\00\00\01\00\00\00X") + (func (export "link") (result i32) + i32.const 3 i32.const 0 i32.const 0 i32.const 1 i32.const 3 i32.const 1 i32.const 5 call $link) + (func (export "write") (result i32) + i32.const 1 i32.const 8 i32.const 1 i32.const 24 call $write))''' + with wasmgpu.Module(binary(wat), execution='compiled', files={'a': b'content'}).spawn(2) as instances: + assert instances.call('link') == [0, 0] + assert instances.call('write') == [0, 0] + assert instances.read_file('alias', instance=1) == b'content' + assert instances.stdout == [b'X', b'X'] + instances.reset() + assert instances.stdout == [b'', b''] + assert instances.read_file('a') == b'content' + with pytest.raises(FileNotFoundError): + instances.read_file('alias') + assert instances.call('link') == [0, 0] + + +@pytest.mark.gpu +@pytest.mark.parametrize('mode', ['compiled', 'interpreter']) +def test_instruction_counter_carry(mode): + module = wasmgpu.Module(binary('(module (func (export "run") (result i32) i32.const 42))'), execution=mode) + with module.spawn(1) as instances: + with pytest.raises(InterruptedError): + instances.call('run', cancel=lambda: True) + batch = instances._batches[0] + low, high = (11, 14) if mode == 'compiled' else (6, 13) + # Seed a valid continuation near overflow to exercise the GPU carry + # without running four billion instructions in a unit test. + batch.state[low], batch.state[high] = 0xfffffffe, 7 + device = instances._context.device + device.queue.write_buffer(batch.buffers[5], 0, batch.state) + encoder = device.create_command_encoder(label='counter carry test') + compute = encoder.begin_compute_pass() + compute.set_pipeline(instances._pipeline) + compute.set_bind_group(0, batch.group) + compute.dispatch_workgroups(1) + compute.end() + device.queue.submit([encoder.finish()]) + state = array('I') + state.frombytes(device.queue.read_buffer(batch.buffers[5])) + assert state[low] == 0 + assert state[high] == 8 + assert state[7] == 1 + + +@pytest.mark.gpu +def test_large_function_has_no_bytecode_or_interpreter(engine, monkeypatch): + operations = 'i32.const 3 i32.add ' * 400 + wat = f'(module (func (export "run") (param i32) (result i32) local.get 0 {operations}))' + module = wasmgpu.Module(binary(wat), compile_limit=64) + reference = wasmgpu.Module(binary(wat), execution='interpreter') + assert module.compiled.complete + assert module.compiled.instructions > MAX_COMPILED_INSTRUCTIONS + assert len(module.compiled.regions) > 1 + # No instruction stream is uploaded: the segment directory immediately + # follows function metadata and the bytecode pointer is absent. + assert module._program[0] == 0 + assert module._program[1] == 16 + len(module._binary.functions) * 8 + assert len(reference._program) - len(module._program) == module.compiled.total_instructions * 4 + + def forbidden(_self): + raise AssertionError('compiled execution requested the interpreter') + monkeypatch.setattr(runtime._Context, 'pipeline', property(forbidden)) + sources = [] + context = runtime._context() + create = context.device.create_shader_module + + def capture(**kwargs): + sources.append(kwargs['code']) + return create(**kwargs) + monkeypatch.setattr(context.device, 'create_shader_module', capture) + rows = [(0,), (7,), (-100,), (2147483640,)] + with module.spawn(len(rows)) as instances: + assert instances.call('run', rows) == oracle(engine, wat, 'run', rows) + assert instances.last_call.interpreted_instructions == 0 + assert instances.last_call.compiled_instructions == module.compiled.instructions * len(rows) + assert sources + for source in sources: + assert 'interpret_step' not in source + assert 'switch op' not in source + assert 'program[0]' not in source + + +@pytest.mark.gpu +def test_mutual_recursion_across_pipelines(engine): + wat = '''(module (type $t (func (param i32) (result i32))) + (table 2 funcref) (elem (i32.const 0) $even $odd) + (func $even (export "run") (type $t) + local.get 0 i32.eqz if (result i32) i32.const 1 else + local.get 0 i32.const 1 i32.sub i32.const 1 call_indirect (type $t) end) + (func $odd (type $t) + local.get 0 i32.eqz if (result i32) i32.const 0 else + local.get 0 i32.const 1 i32.sub call $even end))''' + rows = [(n,) for n in range(12)] + module = wasmgpu.Module(binary(wat), compile_limit=5) + with module.spawn(len(rows), call_depth=32) as instances: + assert instances.call('run', rows) == oracle(engine, wat, 'run', rows) + assert instances.last_call.interpreted_instructions == 0 + assert instances.last_call.dispatches > 10 + + +@pytest.mark.gpu +def test_exact_stack_exhaustion_in_compiled_prefix(): + wat = '''(module (memory 1) + (func $run (export "run") (param i32) (result i32) + i32.const 0 local.get 0 i32.store + local.get 0 i32.const 1 i32.add call $run))''' + results = [] + for mode in ('compiled', 'interpreter'): + with wasmgpu.Module(binary(wat), execution=mode, compile_limit=3).spawn(1, stack_size=8, call_depth=64) as instances: + results.append(outcome(instances, [0], 1000)) + if mode == 'compiled': + assert instances.last_call.interpreted_instructions == 0 + assert results[0] == results[1] + assert results[0][1] == {0: 'stack exhausted'} + + +@pytest.mark.gpu +def test_exported_import_is_a_compiled_continuation(): + wat = '''(module (import "wasi_snapshot_preview1" "args_sizes_get" + (func $sizes (param i32 i32) (result i32))) + (memory 1) (export "run" (func $sizes)))''' + results = [] + for mode in ('compiled', 'interpreter'): + with wasmgpu.Module(binary(wat), execution=mode).spawn(1, wasi=wasmgpu.Wasi(args=['a', 'bb'])) as instances: + results.append((instances.call('run', [(0, 4)]), instances.read_memory(0, 8))) + if mode == 'compiled': + assert instances.last_call.interpreted_instructions == 0 + assert results == [([0], struct.pack(' 10000 + assert module.compiled.instructions < 10 + assert len(module.compiled.source) < 10000 + assert module._program[0] == 0 + rows = [(0,), (1,), (9999,), (10000,), (-1,)] + with module.spawn(len(rows)) as instances: + assert instances.call('run', rows) == oracle(engine, wat, 'run', rows) + assert instances.last_call.interpreted_instructions == 0 + + +@pytest.mark.gpu +def test_small_loop_stays_in_one_native_unit(): + prefix = 'i32.const 0 drop ' * 27 + work = 'i32.const 1 drop ' * 15 + wat = f'''(module (func (export "run") (param i32) (result i32) + {prefix} loop $next {work} + local.get 0 i32.const 1 i32.sub local.tee 0 br_if $next + end local.get 0))''' + with wasmgpu.Module(binary(wat), compile_limit=64).spawn(1) as instances: + assert instances.call('run', [1000]) == [0] + assert instances.last_call.interpreted_instructions == 0 + # The backedge must not require a pipeline transition on every iteration. + assert instances.last_call.dispatches < 30 + + +@pytest.mark.gpu +def test_native_pipeline_eviction_and_compile_failure(monkeypatch): + monkeypatch.setattr(runtime, '_PIPELINE_CACHE_ENTRIES', 2) + modules = [wasmgpu.Module(binary(f'(module (func (export "run") (result i32) i32.const {n}))')) for n in (70123, 70124, 70125)] + with modules[0].spawn(1) as first: + for module in modules[1:]: + with module.spawn(1) as instances: + instances.call('run') + assert len(first._context._compiled) == 2 + assert first._context._compiled_bytes == sum(first._context._compiled_sizes.values()) + assert first.call('run') == [70123] + assert first.last_call.pipelines_created == 1 + assert first.last_call.interpreted_instructions == 0 + + def failed(*_args): + raise RuntimeError('native compiler rejected shader') + def forbidden(_self): + raise AssertionError('native failure fell back to the interpreter') + monkeypatch.setattr(runtime._Context, 'compiled_pipeline', failed) + monkeypatch.setattr(runtime._Context, 'pipeline', property(forbidden)) + with pytest.raises(RuntimeError, match='native compiler rejected'): + modules[0].spawn(1) + + +@pytest.mark.gpu +def test_source_limit_can_split_one_basic_block(monkeypatch): + monkeypatch.setattr(compiler, 'MAX_GENERATED_BYTES', 2500) + wat = '(module (func (export "run") (result i32) i32.const 0 ' + 'i32.const 1 i32.add ' * 25 + '))' + module = wasmgpu.Module(binary(wat)) + with module.spawn(1) as instances: + assert instances.call('run') == [25] + assert instances.last_call.interpreted_instructions == 0 + assert len(module.compiled.regions) > 2 diff --git a/tests/test_gpu_guard.py b/tests/test_gpu_guard.py new file mode 100644 index 0000000..a4ac6de --- /dev/null +++ b/tests/test_gpu_guard.py @@ -0,0 +1,61 @@ +from __future__ import annotations + +import os +import sys +from types import SimpleNamespace + +import pytest + +from tests import gpu_guard +from tests.gpu_guard import compiler_pids, memory, run + +pytestmark = pytest.mark.skipif(sys.platform not in ('darwin', 'linux'), reason='development watchdog supports macOS/Linux') + + +def test_compiler_ownership(): + domain = '''services = { + 0 - com.apple.MTLCompilerService + 123 (pe) com.apple.MTLCompilerService.0000-ABCD + 987 - com.apple.AnotherService + } + service stubs = { + 456 - com.apple.MTLCompilerService + }''' + assert compiler_pids(domain) == {123} + assert compiler_pids('services = {\n}') == set() + + +def test_physical_memory(): + size, identity = memory(os.getpid()) + assert size > 0 + assert identity + assert memory(2147483647) is None + + +def test_normal_exit(): + assert run(['-c', 'raise SystemExit(7)']) == 7 + + +def test_timeout(): + assert run(['-c', 'import time; time.sleep(20)'], seconds=0.3) == 124 + + +def test_memory_limit(): + assert run(['-c', 'import time; data = bytearray(80 * 1024 * 1024); time.sleep(20)'], memory_mib=64) == 124 + + +def test_compile_timeout(): + code = 'from wasmgpu.runtime import _compile_phase; import time; _compile_phase(True); time.sleep(20)' + assert run(['-c', code], seconds=10, compile_seconds=0.3) == 124 + + +def test_termination_only_signals_confirmed_owned_compilers(monkeypatch): + groups, killed = [], [] + worker = SimpleNamespace(pid=123, wait=lambda **_options: 0) + monkeypatch.setattr(gpu_guard.os, 'killpg', lambda pid, _signal: groups.append(pid)) + monkeypatch.setattr(gpu_guard.os, 'kill', lambda pid, _signal: killed.append(pid)) + # A recycled PID must never receive a signal intended for an old compiler. + monkeypatch.setattr(gpu_guard, 'memory', lambda pid: (100, 9) if pid == 200 else (100, 4)) + gpu_guard.terminate(worker, {200: 8, 300: 4}) + assert groups == [123] + assert killed == [300] diff --git a/wasmgpu/compiled.wgsl b/wasmgpu/compiled.wgsl new file mode 100644 index 0000000..fe8d803 --- /dev/null +++ b/wasmgpu/compiled.wgsl @@ -0,0 +1,64 @@ +// Boundaries preserve the interpreter's instruction accounting and virtual time. +fn compiled_tick(count: u32) { + consume_fuel(count); + let before = vm.compiled; + vm.compiled += count; + if vm.compiled < before { vm.compiled_high += 1u; } + if program[12] != 0u { + let delta = mul64(vec2u(count, 0u), vec2u(read_heap(config.fs_offset + 4u), 0u)); + let clock = add64(vec2u(read_heap(config.fs_offset + 1u), read_heap(config.fs_offset + 2u)), delta); + write_heap(config.fs_offset + 1u, clock.x); write_heap(config.fs_offset + 2u, clock.y); + } +} +fn compiled_indirect(type_id: u32, table: u32, index: u32) { + if index >= table_size(table) { fail(3u); return; } + let reference = read_heap(table_offset(table) + index); + if reference == 0u { fail(9u); return; } + if program[16u + (reference - 1u) * 8u + 4u] != type_id { fail(10u); return; } + invoke(reference - 1u); +} +fn compiled_branch_table(offset: u32, count: u32, index: u32) { + let location = program[5] + (offset + min(index, count - 1u)) * 3u; + branch(program[location], program[location + 1u], program[location + 2u]); +} +fn compiled_load(address: u32, width: u32) -> vec2u { + let word = config.memory_offset + address / 4u; + let shift = (address & 3u) * 8u; + var value = vec2u(read_heap(word) >> shift, 0u); + if width == 1u { return value & vec2u(255u, 0u); } + if width == 2u { + if shift == 24u { value.x |= read_heap(word + 1u) << 8u; } + return value & vec2u(65535u, 0u); + } + if shift != 0u { value.x |= read_heap(word + 1u) << (32u - shift); } + if width == 8u { + value.y = read_heap(word + 1u) >> shift; + if shift != 0u { value.y |= read_heap(word + 2u) << (32u - shift); } + } + return value; +} +fn compiled_store(address: u32, width: u32, value: vec2u) { + if (address & 3u) == 0u && width >= 4u { + write_heap(config.memory_offset + address / 4u, value.x); + if width == 8u { write_heap(config.memory_offset + address / 4u + 1u, value.y); } + } else { + for (var i = 0u; i < width; i += 1u) { write_byte(address + i, shr64(value, i * 8u).x); } + } +} +// Dispatch only statically generated basic blocks. An unknown continuation +// belongs to a different compilation unit; the host schedules its native pipeline. +@compute @workgroup_size(64) +fn run(@builtin(global_invocation_id) id: vec3u) { + lane = id.x; + if lane >= config.count { return; } + vm = states[lane]; + var remaining = config.quantum; + loop { + if vm.status != 0u || remaining == 0u { break; } + let consumed = compiled_dispatch(); + if consumed == 0u { break; } + // Quantum boundaries are basic-block boundaries (at most 31 extra ops). + remaining -= min(remaining, consumed); + } + states[lane] = vm; +} diff --git a/wasmgpu/compiler.py b/wasmgpu/compiler.py new file mode 100644 index 0000000..e69d988 --- /dev/null +++ b/wasmgpu/compiler.py @@ -0,0 +1,475 @@ +"""Specialize validated WASM basic blocks into WGSL with local SSA values. + +Calls use explicit frames and compiled continuations. Bounded compilation units +cover the whole module; no unit contains a bytecode interpreter. Native pipelines +are compiled on demand and cached without imposing a size limit on the module. +""" +from __future__ import annotations + +import re +import threading +from array import array +from collections import OrderedDict, deque +from dataclasses import dataclass +from pathlib import Path +from typing import Iterable, cast + +from .binary import ( + BRANCH, + BRANCH_IF, + BRANCH_TABLE, + IF_ZERO, + JUMP, + RETURN, + BinaryModule, + Function, + numeric_signature, +) +from .errors import ResourceLimitError + +TERMINATORS = {0, RETURN, 0x10, 0x11, JUMP, IF_ZERO, BRANCH, BRANCH_IF, BRANCH_TABLE} +MAX_COMPILED_INSTRUCTIONS = 256 +MAX_GENERATED_BYTES = 65536 +TRAPPING_NUMERIC: set[int] = {0x6D, 0x6E, 0x6F, 0x70, 0x7F, 0x80, 0x81, 0x82} +TRAPPING_NUMERIC.update(range(0xA8, 0xAC)) +TRAPPING_NUMERIC.update(range(0xAE, 0xB2)) + + +@dataclass(frozen=True) +class Block: + start: int + end: int + height: int + peak: int + + +def _effect(module: BinaryModule, instruction: list[int]) -> int: + op, a, _, _ = instruction + if op in (0x10, 0x11): + args, results = module.signature(a) if op == 0x10 else module.types[a] + return len(results) - len(args) - int(op == 0x11) + if op in (0x20, 0x23, 0x3F, 0x41, 0x42, 0x43, 0x44, 0xFC10): + return 1 + if op in (0x1A, 0x21, 0x24, IF_ZERO, BRANCH_IF, BRANCH_TABLE, 0xFC0F): + return -1 + if op in (0x1B, 0x26) or 0x36 <= op <= 0x3E: + return -2 + if op in (0xFC08, 0xFC0A, 0xFC0B, 0xFC0C, 0xFC0E, 0xFC11): + return -3 + signature = numeric_signature(op) + return 1 - len(signature[0]) if signature else 0 + + +def blocks(module: BinaryModule, function: Function, maximum: int = 32) -> list[Block]: + """Infer reachable stack heights on the lowered CFG, excluding br_table data.""" + code = function.instructions + if not code: + return [] + heights: dict[int, int] = {} + leaders = {0} + pending = deque([(0, 0)]) + while pending: + pc, height = pending.popleft() + if pc in heights: + if heights[pc] != height: + raise ValueError('inconsistent validated operand stack at a branch') + continue + if not 0 <= pc < len(code): + raise ValueError('validated branch outside function') + heights[pc] = height + op, a, b, c = code[pc] + after = height + _effect(module, code[pc]) + edges = [] + if op in (JUMP, BRANCH, BRANCH_IF, IF_ZERO): + edges.append((a, b + c if op in (BRANCH, BRANCH_IF) else after)) + leaders.add(a) + if op == BRANCH_TABLE: + for _, target, target_height, arity in code[a:a + b]: + edges.append((target, target_height + arity)) + leaders.add(target) + if op not in (0, RETURN, JUMP, BRANCH, BRANCH_TABLE): + edges.append((pc + 1, after)) + if op in TERMINATORS or op >= 0xFC08: + leaders.add(pc + 1) + pending.extend(edges) + result = [] + pc = 0 + while pc < len(code): + if pc not in heights: + pc += 1 + continue + start, peak = pc, heights[pc] + while True: + op = code[pc][0] + peak = max(peak, heights[pc], heights[pc] + _effect(module, code[pc])) + pc += 1 + if op in TERMINATORS or op >= 0xFC08 or pc in leaders or pc not in heights or pc - start >= maximum: + break + result.append(Block(start, pc, heights[start], peak)) + return result + + +def _interval_width(interval: tuple[int, int]) -> int: + return interval[1] - interval[0] + + +def clusters(module: BinaryModule, function: Function, limit: int) -> list[list[Block]]: + """Keep bounded loops together so a backedge need not launch another kernel.""" + body = blocks(module, function, min(32, limit)) + positions = {block.start: index for index, block in enumerate(body)} + totals = [0] + intervals: set[tuple[int, int]] = set() + for index, block in enumerate(body): + totals.append(totals[-1] + block.end - block.start) + op, a, b, _ = function.instructions[block.end - 1] + targets = [a] if op in (JUMP, IF_ZERO, BRANCH, BRANCH_IF) else [] + if op == BRANCH_TABLE: + targets = [target for _, target, _, _ in function.instructions[a:a + b]] + for target in targets: + if target < block.start: + intervals.add((positions[target], index + 1)) + joined: set[int] = set() + for interval in sorted(intervals, key=_interval_width): + start, end = interval + if totals[end] - totals[start] > limit: + continue + while start in joined: + start -= 1 + while end in joined: + end += 1 + if totals[end] - totals[start] <= limit: + joined.update(range(start + 1, end)) + result: list[list[Block]] = [] + for index, block in enumerate(body): + if index not in joined: + result.append([]) + result[-1].append(block) + return result + + +class BlockEmitter: + def __init__(self, module: BinaryModule, function: Function, block: Block, canonical: list[int], tables: dict[int, int], *, checked: bool = False) -> None: # noqa: PLR0913 - Static lowering context and checked-prefix mode. + self.module = module + self.function, self.block, self.canonical = function, block, canonical + self.stack = [f's{i}' for i in range(block.height)] + self.locals: dict[int, str] = {} + self.dirty: set[int] = set() + self.lines: list[str] = [] + self.numeric: set[int] = set() + self.pc = block.start + self.checked = checked + self.tables = tables + + def emit(self, text: str) -> None: + self.lines.append(text) + + def push(self, expression: str) -> None: + if self.checked and len(self.stack) >= self.block.height: + self.trap_if(f'vm.base + {len(self.function.locals) + len(self.stack)}u >= config.stack_cap', '7u') + name = f'v{self.pc}' + self.emit(f'let {name} = {expression};') + self.stack.append(name) + + def local(self, index: int) -> str: + if index not in self.locals: + self.locals[index] = f'l{index}' + self.emit(f'let l{index} = get_value(vm.base + {index}u);') + return self.locals[index] + + def finish(self, action: str = '', *, consumed: int | None = None) -> str: + """Spill only at a continuation or trap, with the exact executed prefix.""" + count = self.pc - self.block.start + 1 if consumed is None else consumed + lines = [f'set_value(vm.base + {index}u, {self.locals[index]});' for index in sorted(self.dirty)] + offset = len(self.function.locals) + lines.extend(f'set_value(vm.base + {offset + index}u, {value});' for index, value in enumerate(self.stack) if value != f's{index}') + lines.extend([f'vm.sp = vm.base + {offset + len(self.stack)}u;', f'vm.pc = {self.function.offset + self.block.start + count}u;', f'compiled_tick({count}u);']) + if action: + lines.append(action) + lines.append(f'return {count}u;') + return '\n'.join(lines) + + def trap_if(self, condition: str, code: str) -> None: + self.emit(f'if {condition} {{\n{self.finish(f"fail({code});")}\n}}') + + def instruction(self, instruction: list[int]) -> bool: # noqa: PLR0915 - Explicit instruction lowering. + op, a, b, c = instruction + pop = self.stack.pop + if op == 0: + self.emit(self.finish('fail(1u);')) + elif op == RETURN: + self.emit(self.finish('return_function();')) + elif op == JUMP: + self.emit(self.finish(f'vm.pc = {self.function.offset + a}u;')) + elif op in (IF_ZERO, BRANCH, BRANCH_IF): + condition = f'{pop()}.x {"==" if op == IF_ZERO else "!="} 0u' if op != BRANCH else 'true' + action = f'vm.pc = {self.function.offset + a}u;' if op == IF_ZERO else f'branch({self.function.offset + a}u, {b}u, {c}u);' + self.emit(self.finish(f'if {condition} {{ {action} }}')) + elif op == BRANCH_TABLE: + selector = pop() + offset = self.tables[self.function.offset + self.pc] + self.emit(self.finish(f'compiled_branch_table({offset}u, {b}u, {selector}.x);')) + elif op == 0x10: + self.emit(self.finish(f'invoke({a}u);')) + elif op == 0x11: + index = pop() + self.emit(self.finish(f'compiled_indirect({self.canonical[a]}u, {b}u, {index}.x);')) + elif op >= 0xFC08: + self.emit(self.finish(f'bulk_operation({op - 0xFC00}u, {a}u, {b}u);')) + elif op == 0x1A: + pop() + elif op == 0x1B: + condition, rhs, lhs = pop(), pop(), pop() + self.push(f'select({rhs}, {lhs}, {condition}.x != 0u)') + elif op == 0x20: + self.push(self.local(a)) + elif op in (0x21, 0x22): + self.locals[a] = pop() if op == 0x21 else self.stack[-1] + self.dirty.add(a) + elif op == 0x23: + self.push(f'vec2u(read_heap({a * 2}u), read_heap({a * 2 + 1}u))') + elif op == 0x24: + value = pop() + self.emit(f'write_heap({a * 2}u, {value}.x); write_heap({a * 2 + 1}u, {value}.y);') + elif op in (0x25, 0x26): + value = pop() if op == 0x26 else '' + index = pop() + self.trap_if(f'{index}.x >= table_size({a}u)', '3u') + if op == 0x25: + self.push(f'vec2u(read_heap(table_offset({a}u) + {index}.x), 0u)') + else: + self.emit(f'write_heap(table_offset({a}u) + {index}.x, {value}.x);') + elif 0x28 <= op <= 0x3E: + value = pop() if op >= 0x36 else '' + pointer = pop() + address = f'address{self.pc}' + self.emit(f'let {address} = {pointer}.x + {a}u;') + self.trap_if(f'{address} < {pointer}.x || !bounds({address}, {b}u)', '2u') + if op >= 0x36: + self.emit(f'compiled_store({address}, {b}u, {value});') + else: + expression = f'compiled_load({address}, {b}u)' + if op in (0x2C, 0x2E, 0x30, 0x32, 0x34): + expression = f'sar64(shl64({expression}, {64 - b * 8}u), {64 - b * 8}u)' + if op in (0x28, 0x2A, 0x2C, 0x2D, 0x2E, 0x2F): + expression = f'vec2u(({expression}).x, 0u)' + self.push(expression) + elif op == 0x3F: + self.push('vec2u(vm.pages, 0u)') + elif op == 0x40: + delta = pop() + self.emit(f'let old{self.pc} = vm.pages;') + self.emit(f'let fits{self.pc} = {delta}.x <= config.memory_cap - vm.pages;') + self.emit(f'if fits{self.pc} {{ vm.pages += {delta}.x; }}') + self.push(f'vec2u(select(0xffffffffu, old{self.pc}, fits{self.pc}), 0u)') + elif 0x41 <= op <= 0x44: + self.push(f'vec2u({a}u, {b}u)') + else: + signature = numeric_signature(op) + if signature is None: + raise ValueError(f'cannot compile opcode {op}') + rhs = pop() if len(signature[0]) == 2 else 'vec2u(0u)' + lhs = pop() + self.numeric.add(op) + self.emit(f'let answer{self.pc} = numeric_{op:x}({lhs}, {rhs});') + if op in TRAPPING_NUMERIC: + self.trap_if(f'answer{self.pc}.trap != 0u', f'answer{self.pc}.trap') + self.push(f'answer{self.pc}.value') + return op in TERMINATORS or op >= 0xFC08 + + def render(self) -> str: + block = self.block + count = block.end - block.start + if self.checked: + self.emit('if !fuel_available(1u) { fail(8u); return 0u; }') + else: + self.emit(f'if !fuel_available({count}u) || vm.base + {len(self.function.locals) + block.peak}u > config.stack_cap {{ return prefix_{self.function.offset + block.start}(); }}') + terminal = False + for self.pc in range(block.start, block.end): + terminal = self.instruction(self.function.instructions[self.pc]) + if not terminal: + self.emit(self.finish()) + body = '\n'.join(self.lines[1:]) + # Untouched incoming stack slots already reside in storage and need no + # load/spill. Omitting dead loads also keeps native compiler input small. + used = sorted({int(match[1]) for match in re.finditer(r'\bs(\d+)\b', body)}) + loads = [f'let s{index} = get_value(vm.base + {len(self.function.locals) + index}u);' for index in used] + return '\n'.join([self.lines[0], *loads, body]) + + +def _numeric(opcodes: set[int]) -> str: + """Reuse exact numeric implementations, removing the runtime opcode switch.""" + source = Path(__file__).with_name('operations.wgsl').read_text() + body = source[source.index(' var result'):source.index(' switch op')] + cases = list(re.finditer(r'^ (case [^\n]+|default): \{', source, re.MULTILINE)) + implementations: dict[int, str] = {} + default = '' + for index, match in enumerate(cases): + end = cases[index + 1].start() if index + 1 < len(cases) else source.rindex(' }') + code = source[match.end():end].rstrip() + code = code[:code.rfind('}')] + label = cast(str, match.group(1)) + if label == 'default': + default = code + else: + for value in re.finditer(r'0x([0-9a-f]+)u', label): + implementations[int(cast(str, value.group(1)), 16)] = code + return '\n'.join(f'fn numeric_{op:x}(a: vec2u, b: vec2u) -> NumericResult {{\nlet op = {op}u;\n{body}{implementations.get(op, default)}\nreturn NumericResult(result, trap);\n}}' for op in sorted(opcodes)) + + + +def vm_support() -> str: + """Shared memory/frames/services ABI, deliberately excluding bytecode execution.""" + source = Path(__file__).with_name('vm.wgsl').read_text().split('fn interpret_step()', 1)[0] + start, end = source.index('fn memory_operation('), source.index('fn table_size(') + return source[:start] + source[end:] + + +class CompiledModule: + """Complete CFG with bounded, lazily generated native compilation units. + + The pc map is host scheduling metadata. GPU buffers contain function metadata + and data/element segments only, never the WASM instruction stream. + """ + + def __init__(self, module: BinaryModule, order: list[int], instruction_limit: int, source_limit: int) -> None: + self.module = module + self.program: array[int] = array('I') + self.source_limit = source_limit + self.functions = tuple(range(len(module.functions))) + self.wasi = any(fn.imported is not None for fn in module.functions) + self.complete = True + self.canonical = [module.types.index(signature) for signature in module.types] + self.bodies = list(module.functions) + for index, function in enumerate(self.bodies): + if function.imported is not None: + params, _ = module.signature(index) + instructions = [[0x20, i, 0, 0] for i in range(len(params))] + [[0x10, index, 0, 0], [RETURN, 0, 0, 0]] + self.bodies[index] = Function(function.type_index, list(params), instructions, function.imported, function.offset, len(params)) + self.total_instructions = sum(len(fn.instructions) for fn in self.bodies) + self.branch_tables: dict[int, int] = {} + offset = 0 + for function in self.bodies: + for pc, (op, _, count, _) in enumerate(function.instructions): + if op == BRANCH_TABLE: + self.branch_tables[function.offset + pc] = offset + offset += count + self.opcodes = tuple(sorted({op for fn in self.bodies for op, _, _, _ in fn.instructions})) + self.regions: list[list[tuple[int, Block]]] = [[]] + size = 0 + self.blocks = self.instructions = 0 + self.locations = array('I', [0]) * self.total_instructions + for index in order: + function = self.bodies[index] + groups = clusters(module, function, instruction_limit) + function_count = sum(block.end - block.start for group in groups for block in group) + if size and size + function_count > instruction_limit: + self.regions.append([]) + size = 0 + for group in groups: + count = sum(block.end - block.start for block in group) + if size + count > instruction_limit: + self.regions.append([]) + size = 0 + for block in group: + self.regions[-1].append((index, block)) + self.locations[function.offset + block.start] = len(self.regions) + self.blocks += len(group) + self.instructions += count + size += count + self._sources: OrderedDict[int, str] = OrderedDict() + self._lock = threading.RLock() + + @property + def source(self) -> str: + """First compilation unit, retained for diagnostics of small modules.""" + return self.source_for(0) + + def region_for(self, pc: int) -> int: + if not 0 <= pc < len(self.locations) or not self.locations[pc]: + raise RuntimeError(f'invalid compiled continuation {pc}') + return self.locations[pc] - 1 + + def _render(self, region: int) -> str: + functions, dispatch = [], [] + numeric: set[int] = set() + for index, block in self.regions[region]: + function = self.bodies[index] + address = function.offset + block.start + emitter = BlockEmitter(self.module, function, block, self.canonical, self.branch_tables) + functions.append(f'fn block_{address}() -> u32 {{\n{emitter.render()}\n}}') + numeric.update(emitter.numeric) + prefix = ['var consumed = 0u;'] + height = block.height + for pc in range(block.start, block.end): + after = height + _effect(self.module, function.instructions[pc]) + step = BlockEmitter(self.module, function, Block(pc, pc + 1, height, max(height, after)), self.canonical, self.branch_tables, checked=True) + step_address = function.offset + pc + functions.append(f'fn step_{step_address}() -> u32 {{\n{step.render()}\n}}') + numeric.update(step.numeric) + prefix.append(f'consumed += step_{step_address}(); if vm.status != 0u {{ return consumed; }}') + height = after + functions.append(f'fn prefix_{address}() -> u32 {{\n' + '\n'.join(prefix) + '\nreturn consumed;\n}') + dispatch.append(f'case {address}u: {{ return block_{address}(); }}') + return (_numeric(numeric) + '\n' + '\n'.join(functions) + + '\nfn compiled_dispatch() -> u32 {\nswitch vm.pc {\n' + '\n'.join(dispatch) + + '\ndefault: { return 0u; }\n}\n}\n') + + def _split(self, region: int) -> None: + units = self.regions[region] + if not units: + raise ResourceLimitError('empty compilation unit exceeds the generated source budget') + if len(units) == 1: + index, block = units[0] + if block.end - block.start == 1: + raise ResourceLimitError('one compiled instruction exceeds the generated source budget') + middle = (block.start + block.end) // 2 + height = block.height + peak = height + function = self.bodies[index] + for pc in range(block.start, middle): + height += _effect(self.module, function.instructions[pc]) + peak = max(peak, height) + units = [(index, Block(block.start, middle, block.height, peak)), + (index, Block(middle, block.end, height, block.peak))] + self.blocks += 1 + middle = len(units) // 2 + self.regions[region] = units[:middle] + self.regions.append(units[middle:]) + for index, block in self.regions[-1]: + self.locations[self.bodies[index].offset + block.start] = len(self.regions) + + def source_for(self, region: int) -> str: + with self._lock: + if region not in self._sources: + source = self._render(region) + while len(source.encode()) > min(MAX_GENERATED_BYTES, self.source_limit): + self._split(region) + source = self._render(region) + self._sources[region] = source + self._sources.move_to_end(region) + while len(self._sources) > 16: + self._sources.popitem(last=False) + return self._sources[region] + + +def compile_module(module: BinaryModule, functions: Iterable[int] | None = None, *, instruction_limit: int = MAX_COMPILED_INSTRUCTIONS, + source_limit: int = MAX_GENERATED_BYTES) -> CompiledModule: + """Compile every function; the limit bounds each unit, never module coverage. + + Optional function indices influence grouping order only. All other functions + follow them and are compiled when reached, with no interpreter fallback. + """ + if not 1 <= instruction_limit <= MAX_COMPILED_INSTRUCTIONS: + raise ResourceLimitError(f'compilation units are limited to 1..{MAX_COMPILED_INSTRUCTIONS} instructions') + if source_limit < 256: + raise ResourceLimitError('generated source limit is too small') + order: list[int] = [] + selected: set[int] = set() + for index in functions or (): + if index not in selected: + selected.add(index) + order.append(index) + if any(not 0 <= index < len(module.functions) for index in order): + raise ValueError('compiled function index out of range') + order.extend(index for index in range(len(module.functions)) if index not in selected) + return CompiledModule(module, order, instruction_limit, source_limit) diff --git a/wasmgpu/filesystem.wgsl b/wasmgpu/filesystem.wgsl index 6571b4f..5eb3108 100644 --- a/wasmgpu/filesystem.wgsl +++ b/wasmgpu/filesystem.wgsl @@ -26,8 +26,8 @@ fn fs_file(fd: u32) -> u32 { fn fs_inode(index: u32) -> u32 { return fs_entry(read_heap(fs_entry(index) + 5u)); } fn fs_right(fd: u32, right: u32) -> bool { return (read_heap(fs_descriptor(fd) + 4u) & right) == right; } fn fs_charge(amount: u32) -> bool { - if amount > vm.fuel { fail(8u); return false; } - vm.fuel -= amount; return true; + if !fuel_available(amount) { fail(8u); return false; } + consume_fuel(amount); return true; } fn fs_open_entry(index: u32) -> bool { for (var fd = 0u; fd < config.fs_fds; fd += 1u) { if fs_file(fd) == index { return true; } } diff --git a/wasmgpu/initialize.wgsl b/wasmgpu/initialize.wgsl new file mode 100644 index 0000000..bfa8498 --- /dev/null +++ b/wasmgpu/initialize.wgsl @@ -0,0 +1,13 @@ +struct InitialConfig { count: u32, words: u32, nonce: u32, first: u32 } +@group(0) @binding(0) var initial: array; +@group(0) @binding(1) var heap: array; +@group(0) @binding(2) var config: InitialConfig; +@compute @workgroup_size(256) +fn initialize(@builtin(global_invocation_id) id: vec3u) { + for (var i = id.x; i < config.words * config.count; i += 16776960u) { + let word = i / config.count; + var value = initial[word]; + if word == config.nonce { value = config.first + i % config.count; } + heap[i] = value; + } +} diff --git a/wasmgpu/runtime.py b/wasmgpu/runtime.py index b62a980..97b32bd 100644 --- a/wasmgpu/runtime.py +++ b/wasmgpu/runtime.py @@ -2,21 +2,34 @@ from __future__ import annotations import copy +import hashlib import importlib +import json import os import struct +import sys import threading +import time from array import array +from collections import OrderedDict +from dataclasses import dataclass, field from pathlib import Path -from typing import Iterable, Mapping, Sequence, Tuple, cast - -from .binary import F32, F64, I32, I64, BinaryModule +from typing import Callable, Iterable, Mapping, Sequence, Tuple, cast + +from .binary import BRANCH_TABLE, F32, F64, I32, I64, BinaryModule +from .compiler import ( + MAX_COMPILED_INSTRUCTIONS, + CompiledModule, + compile_module, + vm_support, +) from .errors import GPUUnavailableError, ResourceLimitError, Trap from .types import ( AdapterInfo, Backend, Blob, Buffer, + Pipeline, RequestAdapter, RequestDevice, Result, @@ -25,7 +38,7 @@ ) from .wasi import WASI_IDS, Wasi, normalize_path -_TRAPS = { +_TRAPS: dict[int, str] = { 1: 'unreachable', 2: 'out of bounds memory access', 3: 'out of bounds table access', 4: 'integer divide by zero', 5: 'integer overflow', 6: 'invalid conversion to integer', 7: 'stack exhausted', 8: 'fuel exhausted', 9: 'uninitialized element', @@ -33,6 +46,43 @@ } _CONTEXT: _Context | None = None _CONTEXT_LOCK = threading.Lock() +_MODULE_CACHE: OrderedDict[bytes, tuple[BinaryModule, array[int]]] = OrderedDict() +_CODE_CACHE: OrderedDict[tuple[BinaryModule, tuple[int, ...] | None, int, int], CompiledModule] = OrderedDict() +_MODULE_LOCK = threading.RLock() +_STATE_WORDS = 16 +_PIPELINE_CACHE_BYTES = 16 * 1024 * 1024 +_PIPELINE_CACHE_ENTRIES = 128 + + +def _compile_phase(active: bool) -> None: + """Notify the external development watchdog without executing guest code.""" + destination = os.environ.get('WASMGPU_GUARD_STATUS') + if destination: + path = Path(destination) + temporary = path.with_suffix('.tmp') + status: dict[str, bool | float] = {'compiling': active, 'started': time.monotonic()} + temporary.write_text(json.dumps(status)) + temporary.replace(path) + + +@dataclass +class CallMetrics: + """Wall-clock phases; execution includes dispatch and completion-state synchronization.""" + + codegen_seconds: float = 0.0 + compile_seconds: float = 0.0 + pipelines_created: int = 0 + pipeline_cache_hits: int = 0 + prepare_seconds: float = 0.0 + upload_seconds: float = 0.0 + execute_seconds: float = 0.0 + max_dispatch_seconds: float = 0.0 + readback_seconds: float = 0.0 + decode_seconds: float = 0.0 + dispatches: int = 0 + compiled_instructions: int = 0 + interpreted_instructions: int = 0 + function_samples: dict[int, int] = field(default_factory=dict) def _words(blob: Blob) -> array[int]: @@ -41,10 +91,10 @@ def _words(blob: Blob) -> array[int]: return words -def _positive(value: int, name: str, zero: bool = False) -> int: +def _positive(value: int, name: str, zero: bool = False, *, bits: int = 32) -> int: if isinstance(value, bool) or not isinstance(value, int): raise TypeError(f'{name} must be an integer') - if value < (0 if zero else 1) or value > 0xFFFFFFFF: + if value < (0 if zero else 1) or value >= 1 << bits: raise ValueError(f'{name} out of range') return value @@ -100,8 +150,93 @@ def __init__(self) -> None: raise GPUUnavailableError(f'could not initialize a hardware GPU: {error}') from error constants = '\n'.join(f'const WASI_{name.upper()}: u32 = {index}u;' for name, index in WASI_IDS.items()) source = constants + '\n' + '\n'.join((Path(__file__).parent / name).read_text() for name in ('numeric.wgsl', 'operations.wgsl', 'vm.wgsl', 'filesystem.wgsl')) - shader = self.device.create_shader_module(label='wasmgpu interpreter', code=source) - self.pipeline = self.device.create_compute_pipeline(layout='auto', compute={'module': shader, 'entry_point': 'run'}) + self._interpreter_source = source + self._interpreter: Pipeline | None = None + self._initializer: Pipeline | None = None + self._services: Pipeline | None = None + self._compiled: OrderedDict[str, Pipeline] = OrderedDict() + self._compiled_sizes: dict[str, int] = {} + self._compiled_bytes = 0 + self._compile_lock = threading.Lock() + self._constants = constants + self._binding_layout = self.device.create_bind_group_layout(entries=[ + {'binding': index, 'visibility': 4, 'buffer': {'type': 'uniform' if index == 1 else 'read-only-storage' if index == 0 else 'storage'}} + for index in range(7) + ]) + self._pipeline_layout = self.device.create_pipeline_layout(bind_group_layouts=[self._binding_layout]) + + @property + def pipeline(self) -> Pipeline: + with self._compile_lock: + if self._interpreter is None: + _compile_phase(True) + try: + shader = self.device.create_shader_module(label='wasmgpu interpreter', code=self._interpreter_source) + self._interpreter = self.device.create_compute_pipeline(layout=self._pipeline_layout, compute={'module': shader, 'entry_point': 'run'}) + finally: + _compile_phase(False) + return self._interpreter + + def compiled_pipeline(self, generated: str) -> tuple[Pipeline, bool]: + key = hashlib.sha256(generated.encode()).hexdigest() + with self._compile_lock: + if key in self._compiled: + self._compiled.move_to_end(key) + return self._compiled[key], True + root = Path(__file__).parent + vm = vm_support() + filesystem = ''' +fn wasi_dispatch(syscall: u32, base: u32) -> u32 { + output[lane] = vec2u(syscall, base); vm.status = 5u; return 0u; +} +fn fs_charge(amount: u32) -> bool { + if !fuel_available(amount) { fail(8u); return false; } + consume_fuel(amount); return true; +}''' + source = '\n'.join([self._constants, (root / 'numeric.wgsl').read_text(), 'struct NumericResult { value: vec2u, trap: u32 }', vm, + filesystem, generated, (root / 'compiled.wgsl').read_text()]) + _compile_phase(True) + try: + shader = self.device.create_shader_module(label='wasmgpu compiled module', code=source) + pipeline = self.device.create_compute_pipeline(layout=self._pipeline_layout, compute={'module': shader, 'entry_point': 'run'}) + finally: + _compile_phase(False) + self._compiled[key] = pipeline + cost = len(generated.encode()) + self._compiled_sizes[key] = cost + self._compiled_bytes += cost + while len(self._compiled) > 1 and (len(self._compiled) > _PIPELINE_CACHE_ENTRIES or self._compiled_bytes > _PIPELINE_CACHE_BYTES): + evicted, _ = self._compiled.popitem(last=False) + self._compiled_bytes -= self._compiled_sizes.pop(evicted) + return pipeline, False + + @property + def services(self) -> Pipeline: + with self._compile_lock: + if self._services is None: + root = Path(__file__).parent + vm = vm_support() + source = '\n'.join([self._constants, (root / 'numeric.wgsl').read_text(), vm, + (root / 'filesystem.wgsl').read_text(), (root / 'services.wgsl').read_text()]) + _compile_phase(True) + try: + shader = self.device.create_shader_module(label='wasmgpu GPU services', code=source) + self._services = self.device.create_compute_pipeline(layout=self._pipeline_layout, compute={'module': shader, 'entry_point': 'service'}) + finally: + _compile_phase(False) + return self._services + + @property + def initializer(self) -> Pipeline: + with self._compile_lock: + if self._initializer is None: + _compile_phase(True) + try: + shader = self.device.create_shader_module(label='wasmgpu initialization', code=Path(__file__).with_name('initialize.wgsl').read_text()) + self._initializer = self.device.create_compute_pipeline(layout='auto', compute={'module': shader, 'entry_point': 'initialize'}) + finally: + _compile_phase(False) + return self._initializer def buffer(self, size: int = 0, data: Blob | array[int] | None = None, uniform: bool = False) -> Buffer: flags = self.wgpu.BufferUsage @@ -119,26 +254,28 @@ def _context() -> _Context: return _CONTEXT -def _program(module: BinaryModule) -> array[int]: +def _program(module: BinaryModule, *, bytecode: bool = True) -> array[int]: canonical = [module.types.index(signature) for signature in module.types] words = [0] * 16 code: list[int] = [] + offset = 0 for fn in module.functions: params, results = module.types[fn.type_index] instructions = fn.instructions if fn.imported is not None: fn.locals = list(params) instructions = [[0x20, index, 0, 0] for index in range(len(params))] + [[0x10, module.functions.index(fn), 0, 0], [0x0F, 0, 0, 0]] - fn.offset = len(code) // 4 + fn.offset = offset + offset += len(instructions) words.extend([fn.offset, len(params), len(results), len(fn.locals), canonical[fn.type_index], WASI_IDS[fn.imported[1]] if fn.imported else 0, 0, 0]) - for op, operand, b, c in instructions: + for op, operand, b, c in instructions if bytecode else (): a = operand if 0x100 <= op <= 0x104: a += fn.offset elif op == 0x11: a = canonical[a] code.extend([op, a & 0xFFFFFFFF, b & 0xFFFFFFFF, c & 0xFFFFFFFF]) - words[0] = len(words) + words[0] = len(words) if bytecode else 0 words.extend(code) words[1] = len(words) words.extend([0] * (len(module.data) * 2)) @@ -151,9 +288,49 @@ def _program(module: BinaryModule) -> array[int]: for i, (_, elements, _, _, _) in enumerate(module.elements): words[words[2] + i * 2:words[2] + i * 2 + 2] = [len(words), len(elements)] words.extend(elements) + if not bytecode: + # Native jump tables contain continuation addresses and stack moves, + # never opcodes. Keeping them in data avoids enormous shader switches. + words[5] = len(words) + for function in module.functions: + for op, start, count, _ in function.instructions: + if op == BRANCH_TABLE: + for _, target, height, arity in function.instructions[start:start + count]: + words.extend([function.offset + target, height, arity]) return array('I', words) +def _load_module(data: bytes) -> tuple[BinaryModule, array[int]]: + with _MODULE_LOCK: + if data not in _MODULE_CACHE: + binary = BinaryModule(data) + Wasi.validate_imports(binary) + _MODULE_CACHE[data] = binary, _program(binary) + _MODULE_CACHE.move_to_end(data) + while len(_MODULE_CACHE) > 1 and (len(_MODULE_CACHE) > 4 or sum(map(len, _MODULE_CACHE)) > 64 * 1024 * 1024): + _, (evicted, _) = _MODULE_CACHE.popitem(last=False) + # Code-cache keys must not keep large evicted syntax trees alive. + for key in list(_CODE_CACHE): + if key[0] is evicted: + del _CODE_CACHE[key] + return _MODULE_CACHE[data] + + +def _compile_module(binary: BinaryModule, functions: tuple[int, ...] | None, limit: int) -> CompiledModule: + with _MODULE_LOCK: + # The older Metal backend used by Python 3.8 exhausted compiler + # memory on a larger single shader. Keep each of its units smaller. + source_limit = 16384 if sys.version_info < (3, 9) else 65536 + key = binary, functions, limit, source_limit + if key not in _CODE_CACHE: + _CODE_CACHE[key] = compile_module(binary, functions, instruction_limit=limit, source_limit=source_limit) + _CODE_CACHE[key].program = _program(binary, bytecode=False) + _CODE_CACHE.move_to_end(key) + while len(_CODE_CACHE) > 16: + _CODE_CACHE.popitem(last=False) + return _CODE_CACHE[key] + + class Module: """Validate a binary WASM file. Function bodies execute only on the GPU. @@ -161,17 +338,28 @@ class Module: runtime dependency; tests use Wasmtime's assembler. """ - def __init__(self, source: str | os.PathLike[str] | Blob, *, files: Mapping[str, Blob] | None = None) -> None: + def __init__(self, source: str | os.PathLike[str] | Blob, *, files: Mapping[str, Blob] | None = None, + execution: str = 'auto', compile_functions: Iterable[int] | None = None, compile_limit: int = MAX_COMPILED_INSTRUCTIONS) -> None: + started = time.perf_counter() + if execution not in ('auto', 'compiled', 'interpreter'): + raise ValueError('execution must be auto, compiled or interpreter') + _positive(compile_limit, 'compile_limit', zero=True) if isinstance(source, (str, os.PathLike)): data = Path(source).read_bytes() elif isinstance(source, (bytes, bytearray, memoryview)): data = bytes(source) else: raise TypeError('Module requires a filesystem path or WASM bytes') - self._binary = BinaryModule(data) - Wasi.validate_imports(self._binary) - self._program = _program(self._binary) + self._binary, self._program = _load_module(data) self._files = Wasi(files=files).files + self.load_seconds = time.perf_counter() - started + started = time.perf_counter() + self.execution = execution + selection = tuple(_positive(index, 'compiled function index', zero=True) for index in compile_functions) if compile_functions is not None else None + self.compiled = _compile_module(self._binary, selection, compile_limit) if execution != 'interpreter' else None + if self.compiled is not None: + self._program = self.compiled.program + self.codegen_seconds = time.perf_counter() - started @property def exports(self) -> dict[str, str]: @@ -199,42 +387,8 @@ def __init__(self, owner: Instances, count: int, first: int) -> None: self.context = owner._context device = self.context.device binary = owner.module._binary - words = array('I', [0]) * (owner._heap_words * count) - - def repeat(offset: int, value: int) -> None: - words[offset * count:(offset + 1) * count] = array('I', [value]) * count - - for index, (_, _, value) in enumerate(binary.globals): - repeat(index * 2, value & 0xFFFFFFFF) - repeat(index * 2 + 1, value >> 32) - initial_memory = bytearray((binary.memory[0] if binary.memory else 0) * 65536) - for index, (offset, blob) in enumerate(binary.data): - if offset is not None: - if offset > len(initial_memory) or len(blob) > len(initial_memory) - offset: - raise Trap({first: 'out of bounds memory access during instantiation'}, []) - initial_memory[offset:offset + len(blob)] = blob - repeat(owner._data_flags + index, 1) - for index, value in enumerate(_words(initial_memory)): - if value: - repeat(owner._memory_offset + index, value) - for index, (_, (initial_size, _)) in enumerate(binary.tables): - repeat(owner._table_lengths + index, initial_size) - for index, (offset, elements, declarative, _, table) in enumerate(binary.elements): - if offset is not None: - initial_size = binary.tables[table][1][0] - if offset > initial_size or len(elements) > initial_size - offset: - raise Trap({first: 'out of bounds table access during instantiation'}, []) - for element_index, value in enumerate(elements): - repeat(owner._table_offsets[table] + offset + element_index, value) - if offset is not None or declarative: - repeat(owner._element_flags + index, 1) - for offset, value in enumerate(owner._filesystem): - if value: - repeat(owner._fs_offset + offset, value) - if owner._filesystem: - for lane in range(count): - words[(owner._fs_offset + 18) * count + lane] = first + lane self.buffers: list[Buffer] = [] + self.initialization_config: Buffer | None = None try: assert owner._program_buffer is not None self.buffers = [owner._program_buffer] @@ -244,19 +398,24 @@ def repeat(offset: int, value: int) -> None: owner._fs_offset, owner._wasi.max_files, owner._wasi.max_fds, owner._wasi.storage_size]) self.buffers.append(self.context.buffer(data=config, uniform=True)) self.buffers.append(self.context.buffer(size=count * owner.stack_size * 8)) - self.buffers.append(self.context.buffer(data=words)) + self.buffers.append(self.context.buffer(size=owner._heap_words * count * 4)) self.buffers.append(self.context.buffer(size=count * owner.call_depth * 16)) - state = array('I', [0]) * (count * 12) - for lane in range(count): - state[lane * 12 + 5] = binary.memory[0] if binary.memory else 0 - state[lane * 12 + 6] = 0 + state = array('I', [0, 0, 0, 0, 0, binary.memory[0] if binary.memory else 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]) * count self.buffers.append(self.context.buffer(data=state)) self.buffers.append(self.context.buffer(size=count * owner._output_slots * 8)) self.state = state - self.group = device.create_bind_group(layout=self.context.pipeline.get_bind_group_layout(0), entries=[ + self.group = device.create_bind_group(layout=owner._pipeline.get_bind_group_layout(0), entries=[ {'binding': index, 'resource': {'buffer': buffer, 'offset': 0, 'size': buffer.size}} for index, buffer in enumerate(self.buffers) ]) + if owner._template_buffer is not None: + self.initialization_config = self.context.buffer(data=array('I', [count, owner._heap_words, + owner._fs_offset + 18 if owner._uses_wasi else 0xffffffff, first]), uniform=True) + self.initialization_group = device.create_bind_group(layout=self.context.initializer.get_bind_group_layout(0), entries=[ + {'binding': index, 'resource': {'buffer': buffer, 'offset': 0, 'size': buffer.size}} + for index, buffer in enumerate((owner._template_buffer, self.buffers[3], self.initialization_config)) + ]) + self.initialize_heap() except Exception: self.close() raise @@ -265,48 +424,101 @@ def close(self) -> None: for buffer in self.buffers[1:]: buffer.destroy() self.buffers = [] + if self.initialization_config is not None: + self.initialization_config.destroy() + self.initialization_config = None - def execute(self, function_index: int, inputs: Sequence[tuple[int, ...]], fuel: int, raw: bool = False) -> tuple[list[Result], dict[int, str]]: + def initialize_heap(self) -> None: + if self.owner._template_buffer is None: + self.context.device.queue.write_buffer(self.buffers[3], 0, self.owner._initial_words) + else: + encoder = self.context.device.create_command_encoder(label='wasmgpu initialize batch') + compute = encoder.begin_compute_pass() + compute.set_pipeline(self.context.initializer) + compute.set_bind_group(0, self.initialization_group) + compute.dispatch_workgroups(min(65535, (self.count * self.owner._heap_words + 255) // 256)) + compute.end() + self.context.device.queue.submit([encoder.finish()]) + + def execute(self, function_index: int, inputs: Sequence[tuple[int, ...]], fuel: int, raw: bool = False, # noqa: PLR0915 - Dispatch, phase timing and decoding. + cancel: Callable[[], bool] | None = None) -> tuple[list[Result], dict[int, str]]: owner = self.owner device = self.context.device fn = owner.module._binary.functions[function_index] _, returns = owner.module._binary.signature(function_index) + metrics = owner.last_call + started = time.perf_counter() arguments = array('Q', [0]) * (self.count * max(1, len(fn.locals))) for lane, row in enumerate(inputs): for index, value in enumerate(row): arguments[index * self.count + lane] = value - base = lane * 12 - pages, table_len = self.state[base + 5], self.state[base + 6] - self.state[base:base + 12] = array('I', [fn.offset, len(fn.locals), 0, function_index, 0, pages, table_len, 0, 0, fuel, 0, 0]) + base = lane * _STATE_WORDS + pages = self.state[base + 5] + self.state[base:base + _STATE_WORDS] = array('I', [fn.offset, len(fn.locals), 0, function_index, 0, pages, 0, 0, 0, fuel & 0xffffffff, 0, 0, fuel >> 32, 0, 0, 0]) + metrics.prepare_seconds += time.perf_counter() - started + started = time.perf_counter() device.queue.write_buffer(self.buffers[2], 0, arguments) device.queue.write_buffer(self.buffers[5], 0, self.state) + metrics.upload_seconds += time.perf_counter() - started + previous_compiled = previous_interpreted = 0 + last_region = -1 while True: + if cancel is not None and cancel(): + raise InterruptedError('GPU invocation cancelled between dispatches') + if any(status == 5 for status in self.state[7::_STATE_WORDS]): + pipeline = self.context.services + elif owner.module.compiled is not None: + plan = owner.module.compiled + regions = sorted({plan.region_for(self.state[base]) for base in range(0, len(self.state), _STATE_WORDS) if self.state[base + 7] == 0}) + region = next((index for index in regions if index > last_region), regions[0]) + pipeline = owner._native_pipeline(region, metrics) + last_region = region + else: + pipeline = owner._pipeline + if cancel is not None and cancel(): + raise InterruptedError('GPU invocation cancelled between dispatches') + started = time.perf_counter() encoder = device.create_command_encoder(label='wasmgpu dispatch') compute = encoder.begin_compute_pass() - compute.set_pipeline(self.context.pipeline) + compute.set_pipeline(pipeline) compute.set_bind_group(0, self.group) compute.dispatch_workgroups((self.count + 63) // 64) compute.end() device.queue.submit([encoder.finish()]) self.state = array('I') self.state.frombytes(device.queue.read_buffer(self.buffers[5])) - statuses = self.state[7::12] + elapsed = time.perf_counter() - started + metrics.execute_seconds += elapsed + metrics.max_dispatch_seconds = max(metrics.max_dispatch_seconds, elapsed) + metrics.dispatches += 1 + compiled = sum(self.state[11::_STATE_WORDS]) + (sum(self.state[14::_STATE_WORDS]) << 32) + interpreted = sum(self.state[6::_STATE_WORDS]) + (sum(self.state[13::_STATE_WORDS]) << 32) + metrics.compiled_instructions += compiled - previous_compiled + metrics.interpreted_instructions += interpreted - previous_interpreted + previous_compiled, previous_interpreted = compiled, interpreted + for function in self.state[3::_STATE_WORDS]: + metrics.function_samples[function] = metrics.function_samples.get(function, 0) + 1 + statuses = self.state[7::_STATE_WORDS] if all(status in (1, 2) for status in statuses): break + started = time.perf_counter() output = array('Q') output.frombytes(device.queue.read_buffer(self.buffers[6])) + metrics.readback_seconds += time.perf_counter() - started + started = time.perf_counter() results: list[Result] = [] traps: dict[int, str] = {} for lane in range(self.count): - if self.state[lane * 12 + 7] == 2: - code = self.state[lane * 12 + 8] + if self.state[lane * _STATE_WORDS + 7] == 2: + code = self.state[lane * _STATE_WORDS + 8] if code == 12: - owner.exit_codes[self.first + lane] = self.state[lane * 12 + 10] + owner.exit_codes[self.first + lane] = self.state[lane * _STATE_WORDS + 10] traps[self.first + lane] = _TRAPS[code] results.append(None) else: output_row = tuple(output[index * self.count + lane] if raw else _decode(output[index * self.count + lane], ty) for index, ty in enumerate(returns)) results.append(output_row[0] if len(output_row) == 1 else output_row if output_row else None) + metrics.decode_seconds += time.perf_counter() - started return results, traps @@ -318,15 +530,17 @@ def __init__(self, module: Module, count: int, *, memory_pages: int | None, tabl stack_size: int, call_depth: int, fuel: int, quantum: int, batch_size: int | None, max_resident_bytes: int, wasi: Wasi | None) -> None: self.module = module + self.last_call = CallMetrics() self.count = _positive(count, 'count', zero=True) self.stack_size = _positive(stack_size, 'stack_size') self.call_depth = _positive(call_depth, 'call_depth') - self.fuel = _positive(fuel, 'fuel') + self.fuel = _positive(fuel, 'fuel', bits=64) self.quantum = _positive(quantum, 'quantum') self._closed = False self._lock = threading.RLock() self._batches: list[_Batch] = [] self._program_buffer: Buffer | None = None + self._template_buffer: Buffer | None = None binary = module._binary initial_pages, declared_pages = binary.memory or (0, 0) self.memory_pages = min(declared_pages if declared_pages is not None else 65535, max(initial_pages, 16)) if memory_pages is None else _positive(memory_pages, 'memory_pages', zero=True) @@ -355,13 +569,14 @@ def __init__(self, module: Module, count: int, *, memory_pages: int | None, tabl if not isinstance(self._wasi, Wasi): raise TypeError('wasi must be a Wasi configuration') uses_wasi = any(fn.imported is not None for fn in binary.functions) or bool(module._files) or bool(self._wasi.files) + self._uses_wasi = uses_wasi filesystem_words = 48 + self._wasi.max_files * 76 + self._wasi.max_fds * 8 + (self._wasi.storage_size + 3) // 4 if uses_wasi else 0 self._heap_words = max(4, self._fs_offset + filesystem_words) self._output_slots = max([8] + [max(len(args), len(results)) for args, results in binary.types]) for fn in binary.functions: if len(fn.locals) + fn.max_stack > self.stack_size: raise ResourceLimitError('stack_size is smaller than a function requires') - per_instance = self.stack_size * 8 + self.call_depth * 16 + self._heap_words * 4 + 48 + self._output_slots * 8 + per_instance = self.stack_size * 8 + self.call_depth * 16 + self._heap_words * 4 + _STATE_WORDS * 4 + self._output_slots * 8 if isinstance(max_resident_bytes, bool) or not isinstance(max_resident_bytes, int): raise TypeError('max_resident_bytes must be an integer') if max_resident_bytes <= 0: @@ -369,7 +584,7 @@ def __init__(self, module: Module, count: int, *, memory_pages: int | None, tabl budget = max_resident_bytes environment_strings = [*self._wasi.args, *(f'{key}={value}' for key, value in self._wasi.env.items())] program_bytes = 4 * (len(module._program) + len(binary.tables) * 4 + sum(2 + (len(value.encode()) + 4) // 4 for value in environment_strings)) - self.resident_bytes = per_instance * self.count + program_bytes + self.resident_bytes = per_instance * self.count + program_bytes + (self._heap_words * 4 if count > 1 else 0) if self.resident_bytes > budget: raise ResourceLimitError(f'{self.count} instances require approximately {self.resident_bytes} bytes; budget is {budget}') if uses_wasi: @@ -381,7 +596,7 @@ def __init__(self, module: Module, count: int, *, memory_pages: int | None, tabl limits = self._context.limits if program_bytes > min(limits['max-storage-buffer-binding-size'], limits['max-buffer-size']): raise ResourceLimitError('program and guest environment exceed the device buffer limit') - largest = max(self.stack_size * 8, self.call_depth * 16, self._heap_words * 4, 48, self._output_slots * 8) + largest = max(self.stack_size * 8, self.call_depth * 16, self._heap_words * 4, _STATE_WORDS * 4, self._output_slots * 8) max_batch = min(limits['max-storage-buffer-binding-size'] // largest, limits['max-compute-workgroups-per-dimension'] * 64, (128 * 1024 * 1024) // per_instance) @@ -390,13 +605,28 @@ def __init__(self, module: Module, count: int, *, memory_pages: int | None, tabl self.batch_size = max_batch if batch_size is None else _positive(batch_size, 'batch_size') if self.batch_size > max_batch: raise ResourceLimitError(f'batch_size exceeds device/runtime limit {max_batch}') - self.resident_bytes += 64 * ((count + self.batch_size - 1) // self.batch_size) + self.resident_bytes += (80 if count > 1 else 64) * ((count + self.batch_size - 1) // self.batch_size) if self.resident_bytes > budget: raise ResourceLimitError('batch configuration buffers exceed the resident allocation budget') - self._filesystem = self._build_filesystem() if uses_wasi else array('I') + self.codegen_seconds = 0.0 + started = time.perf_counter() + if count > 1: + _ = self._context.initializer + if module.compiled is not None: + if module.compiled.wasi: + _ = self._context.services + generated_at = time.perf_counter() + generated = module.compiled.source + self.codegen_seconds = time.perf_counter() - generated_at + self._pipeline, self.pipeline_cache_hit = self._context.compiled_pipeline(generated) + else: + self.pipeline_cache_hit = self._context._interpreter is not None + self._pipeline = self._context.pipeline + self.pipeline_seconds = time.perf_counter() - started - self.codegen_seconds + self._initial_words = self._build_initial_heap() if count else array('I') try: program = array('I', module._program) - program[12] = int(bool(self._filesystem)) + program[12] = int(uses_wasi) program[4] = len(program) for index, (offset, capacity) in enumerate(zip(self._table_offsets, self._table_capacities)): program.extend([self._table_lengths + index, offset, capacity, 0]) @@ -409,19 +639,83 @@ def __init__(self, module: Module, count: int, *, memory_pages: int | None, tabl program[info_offset + index * 2:info_offset + index * 2 + 2] = array('I', [len(program), len(blob)]) program.extend(_words(blob + b'\0' * (-len(blob) % 4))) self._program_buffer = self._context.buffer(data=program) + if count > 1: + self._template_buffer = self._context.buffer(data=self._initial_words) + self._initial_words = array('I') for first in range(0, count, self.batch_size): self._batches.append(_Batch(self, min(self.batch_size, count - first), first)) + self._synchronize_initialization() if binary.start is not None: self._call(binary.start, [()] * count, self.fuel) except Exception: self.close() raise + def _native_pipeline(self, region: int, metrics: CallMetrics) -> Pipeline: + assert self.module.compiled is not None + started = time.perf_counter() + source = self.module.compiled.source_for(region) + metrics.codegen_seconds += time.perf_counter() - started + started = time.perf_counter() + pipeline, hit = self._context.compiled_pipeline(source) + metrics.compile_seconds += time.perf_counter() - started + metrics.pipelines_created += int(not hit) + metrics.pipeline_cache_hits += int(hit) + return pipeline + @property def adapter_info(self) -> AdapterInfo: assert self._context.adapter is not None return dict(self._context.adapter.info) + def _build_initial_heap(self) -> array[int]: + binary = self.module._binary + words = array('I', [0]) * self._heap_words + for index, (_, _, value) in enumerate(binary.globals): + words[index * 2:index * 2 + 2] = array('I', [value & 0xffffffff, value >> 32]) + memory = bytearray((binary.memory[0] if binary.memory else 0) * 65536) + for index, (offset, blob) in enumerate(binary.data): + if offset is not None: + if offset > len(memory) or len(blob) > len(memory) - offset: + raise Trap({0: 'out of bounds memory access during instantiation'}, []) + memory[offset:offset + len(blob)] = blob + words[self._data_flags + index] = 1 + words[self._memory_offset:self._memory_offset + len(memory) // 4] = _words(memory) + for index, (_, (size, _)) in enumerate(binary.tables): + words[self._table_lengths + index] = size + for index, (offset, elements, declarative, _, table) in enumerate(binary.elements): + if offset is not None: + size = binary.tables[table][1][0] + if offset > size or len(elements) > size - offset: + raise Trap({0: 'out of bounds table access during instantiation'}, []) + start = self._table_offsets[table] + offset + words[start:start + len(elements)] = array('I', elements) + if offset is not None or declarative: + words[self._element_flags + index] = 1 + if self._uses_wasi: + filesystem = self._build_filesystem() + words[self._fs_offset:self._fs_offset + len(filesystem)] = filesystem + return words + + def _synchronize_initialization(self) -> None: + if self._batches: + self._context.device.queue.read_buffer(self._batches[-1].buffers[5], 0, 4) + + def reset(self) -> None: + """Restore fresh guest state and rerun its start function, reusing GPU buffers.""" + with self._lock: + self._check_open() + binary = self.module._binary + for batch in self._batches: + batch.initialize_heap() + batch.state = array('I', [0, 0, 0, 0, 0, binary.memory[0] if binary.memory else 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]) * batch.count + self._context.device.queue.write_buffer(batch.buffers[5], 0, batch.state) + self._synchronize_initialization() + self.exit_codes = [None] * self.count + self.last_call = CallMetrics() + if binary.start is not None: + self._call(binary.start, [()] * self.count, self.fuel) + @property def stdout(self) -> list[bytes]: return [self._read_file_index(1, instance) for instance in range(self.count)] @@ -462,31 +756,36 @@ def _build_filesystem(self) -> array[int]: def _read_file_index(self, index: int, instance: int) -> bytes: with self._lock: self._check_open() - if not self._filesystem: + if not self._uses_wasi: return b'' batch, lane = self._locate(instance) - raw = self._context.device.queue.read_buffer(batch.buffers[3]) - words = _words(raw)[lane::batch.count] entry = self._fs_offset + 48 + index * 12 - inode = words[entry + 5] - entry = self._fs_offset + 48 + inode * 12 - length, start = words[entry + 2], words[entry + 3] + metadata = self._read_heap(batch, lane, entry, 12) + inode = metadata[5] + if inode != index: + metadata = self._read_heap(batch, lane, self._fs_offset + 48 + inode * 12, 12) + length, start = metadata[2], metadata[3] + if length == 0: + return b'' data_offset = self._fs_offset + 48 + self._wasi.max_files * 76 + self._wasi.max_fds * 8 - content = array('I', words[data_offset + start // 4:data_offset + (start + length + 3) // 4]).tobytes() + content = self._read_heap(batch, lane, data_offset + start // 4, (start % 4 + length + 3) // 4).tobytes() return content[start % 4:start % 4 + length] + def _read_heap(self, batch: _Batch, lane: int, start: int, count: int) -> array[int]: + raw = self._context.device.queue.read_buffer(batch.buffers[3], start * batch.count * 4, count * batch.count * 4) + return _words(raw)[lane::batch.count] + def read_file(self, path: str, *, instance: int = 0) -> bytes: """Read a file from one instance's GPU filesystem.""" with self._lock: self._check_open() encoded_path = normalize_path(path).encode() batch, lane = self._locate(instance) - if self._filesystem: - raw = self._context.device.queue.read_buffer(batch.buffers[3]) - words = _words(raw)[lane::batch.count] - names = self._fs_offset + 48 + self._wasi.max_files * 12 + if self._uses_wasi: + words = self._read_heap(batch, lane, self._fs_offset, 48 + self._wasi.max_files * 76) + names = 48 + self._wasi.max_files * 12 for index in range(4, self._wasi.max_files): - entry = self._fs_offset + 48 + index * 12 + entry = 48 + index * 12 name = array('I', words[names + index * 64:names + (index + 1) * 64]).tobytes()[:words[entry + 1]] if words[entry] == 1 and name == encoded_path: return self._read_file_index(index, instance) @@ -499,7 +798,8 @@ def _check_open(self) -> None: if self._closed: raise RuntimeError('instances are closed') - def call(self, name: str, inputs: Iterable[Row] | None = None, *, fuel: int | None = None) -> list[Result]: + def call(self, name: str, inputs: Iterable[Row] | None = None, *, fuel: int | None = None, + cancel: Callable[[], bool] | None = None) -> list[Result]: """Invoke one export per instance, in input order. A scalar per instance is accepted for a single argument. For multiple @@ -508,6 +808,9 @@ def call(self, name: str, inputs: Iterable[Row] | None = None, *, fuel: int | No """ with self._lock: self._check_open() + started = time.perf_counter() + if cancel is not None and not callable(cancel): + raise TypeError('cancel must be a callable') binary = self.module._binary if name not in binary.exports or binary.exports[name][0] != 0: raise KeyError(f'no exported function {name!r}') @@ -527,13 +830,20 @@ def call(self, name: str, inputs: Iterable[Row] | None = None, *, fuel: int | No if not isinstance(normalized, (tuple, list)) or len(normalized) != len(params): raise TypeError(f'each input row must contain {len(params)} arguments') encoded.append(tuple(_encode(value, ty) for value, ty in zip(cast(Sequence[Scalar], normalized), params))) - return self._call(index, encoded, self.fuel if fuel is None else _positive(fuel, 'fuel')) - - def _call(self, index: int, inputs: Sequence[tuple[int, ...]], fuel: int, raw: bool = False) -> list[Result]: + budget = self.fuel if fuel is None else _positive(fuel, 'fuel', bits=64) + preparation = time.perf_counter() - started + try: + return self._call(index, encoded, budget, cancel=cancel) + finally: + self.last_call.prepare_seconds += preparation + + def _call(self, index: int, inputs: Sequence[tuple[int, ...]], fuel: int, raw: bool = False, + cancel: Callable[[], bool] | None = None) -> list[Result]: + self.last_call = CallMetrics() results: list[Result] = [] traps: dict[int, str] = {} for batch in self._batches: - output, errors = batch.execute(index, inputs[batch.first:batch.first + batch.count], fuel, raw=raw) + output, errors = batch.execute(index, inputs[batch.first:batch.first + batch.count], fuel, raw=raw, cancel=cancel) results.extend(output) traps.update(errors) if traps: @@ -554,7 +864,7 @@ def read_memory(self, offset: int, size: int, *, instance: int = 0) -> bytes: _positive(offset, 'offset', zero=True) _positive(size, 'size', zero=True) batch, lane = self._locate(instance) - length = batch.state[lane * 12 + 5] * 65536 + length = batch.state[lane * _STATE_WORDS + 5] * 65536 if offset > length or size > length - offset: raise IndexError('memory range out of bounds') if size == 0: @@ -590,6 +900,9 @@ def close(self) -> None: batch.close() if self._program_buffer is not None: self._program_buffer.destroy() + if self._template_buffer is not None: + self._template_buffer.destroy() + self._initial_words = array('I') self._closed = True def __enter__(self) -> Instances: diff --git a/wasmgpu/services.wgsl b/wasmgpu/services.wgsl new file mode 100644 index 0000000..841ed67 --- /dev/null +++ b/wasmgpu/services.wgsl @@ -0,0 +1,16 @@ +// The compiled tier outlines WASI into a shared GPU kernel. Pending arguments +// remain on the operand stack; output[0] holds a private continuation record +// until return_function writes final results. No syscall is handled by Python. +@compute @workgroup_size(64) +fn service(@builtin(global_invocation_id) id: vec3u) { + lane = id.x; + if lane >= config.count { return; } + vm = states[lane]; + if vm.status != 5u { return; } + let pending = output[lane]; + vm.status = 0u; + let answer = wasi_dispatch(pending.x, pending.y); + vm.sp = pending.y; + if pending.x != WASI_PROC_EXIT { push(vec2u(answer, 0u)); } + states[lane] = vm; +} diff --git a/wasmgpu/types.py b/wasmgpu/types.py index 0633a77..67cba70 100644 --- a/wasmgpu/types.py +++ b/wasmgpu/types.py @@ -49,7 +49,9 @@ class Device(Protocol): def create_buffer(self, *, size: int, usage: int) -> Buffer: ... def create_buffer_with_data(self, *, data: Blob | array[int], usage: int) -> Buffer: ... def create_shader_module(self, *, label: str, code: str) -> object: ... - def create_compute_pipeline(self, *, layout: str, compute: Mapping[str, object]) -> Pipeline: ... + def create_compute_pipeline(self, *, layout: object, compute: Mapping[str, object]) -> Pipeline: ... + def create_bind_group_layout(self, *, entries: Sequence[Mapping[str, object]]) -> object: ... + def create_pipeline_layout(self, *, bind_group_layouts: Sequence[object]) -> object: ... def create_bind_group(self, *, layout: object, entries: Sequence[Mapping[str, object]]) -> object: ... def create_command_encoder(self, *, label: str) -> Encoder: ... diff --git a/wasmgpu/vm.wgsl b/wasmgpu/vm.wgsl index 7ae7ba3..ed6cb9e 100644 --- a/wasmgpu/vm.wgsl +++ b/wasmgpu/vm.wgsl @@ -6,8 +6,9 @@ struct Config { } struct State { pc: u32, sp: u32, base: u32, function_id: u32, - depth: u32, pages: u32, table_len: u32, status: u32, - trap: u32, fuel: u32, exit_code: u32, reserved: u32, + depth: u32, pages: u32, interpreted: u32, status: u32, + trap: u32, fuel: u32, exit_code: u32, compiled: u32, + fuel_high: u32, interpreted_high: u32, compiled_high: u32, padding: u32, } @group(0) @binding(0) var program: array; @group(0) @binding(1) var config: Config; @@ -19,6 +20,11 @@ struct State { var lane: u32; var vm: State; fn fail(code: u32) { vm.trap = code; vm.status = 2u; } +fn fuel_available(amount: u32) -> bool { return vm.fuel_high != 0u || vm.fuel >= amount; } +fn consume_fuel(amount: u32) { + if vm.fuel < amount { vm.fuel_high -= 1u; } + vm.fuel -= amount; +} fn get_value(index: u32) -> vec2u { return values[index * config.count + lane]; } fn set_value(index: u32, value: vec2u) { values[index * config.count + lane] = value; } fn push(value: vec2u) { @@ -57,10 +63,17 @@ fn invoke(function_id: u32) { if program[metadata + 5u] != 0u { let syscall = program[metadata + 5u]; let answer = wasi_dispatch(syscall, base); + if vm.status == 5u { return; } vm.sp = base; if program[metadata + 2u] != 0u { push(vec2u(answer, 0u)); } return; } + invoke_native(function_id); +} +fn invoke_native(function_id: u32) { + let metadata = 16u + function_id * 8u; + let params = program[metadata + 1u]; let locals = program[metadata + 3u]; + let base = vm.sp - params; if vm.depth >= config.frame_cap || base + locals > config.stack_cap { fail(7u); return; } frames[vm.depth * config.count + lane] = vec4u(vm.pc, vm.base, vm.function_id, base); vm.depth += 1u; vm.base = base; vm.function_id = function_id; @@ -140,15 +153,11 @@ fn bulk_operation(op: u32, index: u32, other: u32) { } } } -@compute @workgroup_size(64) -fn run(@builtin(global_invocation_id) id: vec3u) { - lane = id.x; - if lane >= config.count { return; } - vm = states[lane]; - if vm.status != 0u { return; } - for (var step = 0u; step < config.quantum; step += 1u) { - if vm.fuel == 0u { fail(8u); break; } - vm.fuel -= 1u; +fn interpret_step() { + if !fuel_available(1u) { fail(8u); return; } + consume_fuel(1u); + vm.interpreted += 1u; + if vm.interpreted == 0u { vm.interpreted_high += 1u; } if program[12] != 0u { let clock = add64(vec2u(read_heap(config.fs_offset + 1u), read_heap(config.fs_offset + 2u)), vec2u(read_heap(config.fs_offset + 4u), 0u)); write_heap(config.fs_offset + 1u, clock.x); write_heap(config.fs_offset + 2u, clock.y); @@ -216,6 +225,15 @@ fn run(@builtin(global_invocation_id) id: vec3u) { } } } +} +@compute @workgroup_size(64) +fn run(@builtin(global_invocation_id) id: vec3u) { + lane = id.x; + if lane >= config.count { return; } + vm = states[lane]; + if vm.status != 0u { return; } + for (var step = 0u; step < config.quantum; step += 1u) { + interpret_step(); if vm.status != 0u { break; } } states[lane] = vm;