diff --git a/.gitignore b/.gitignore deleted file mode 100644 index 70b7b24..0000000 --- a/.gitignore +++ /dev/null @@ -1,46 +0,0 @@ -# Python cache files -__pycache__/ -*.py[cod] -*$py.class - -# Native/Cython build artifacts -*.so -*.pyd -*.dylib -*.dll -charmnumeric/_native_region.cpp - -# Packaging outputs -build/ -dist/ -.eggs/ -*.egg -*.egg-info/ -pip-wheel-metadata/ - -# Test and coverage outputs -.pytest_cache/ -.mypy_cache/ -.ruff_cache/ -.coverage -.coverage.* -htmlcov/ - -# Virtual environments -.venv/ -venv/ -env/ -ENV/ -.python-version - -# Local CMake build directories -src/build/ -tests/build/ -cmake-build-*/ - -# Editor/system files -.idea/ -.vscode/ -*.swp -*.swo -.DS_Store diff --git a/CODEMAP.md b/CODEMAP.md deleted file mode 100644 index 738bc45..0000000 --- a/CODEMAP.md +++ /dev/null @@ -1,163 +0,0 @@ -# Codemap: Charmnumeric (`example/charmnumeric/`) - -Distributed N-dimensional array DSL built on charmtyles. Supports up to 3D arrays with tile decomposition. Python `Array`/`ArrayView` classes extend `FrontendObject`; C++ backend uses Charm++ chare arrays (`Partition`) with optional MLIR JIT and GPU (Kokkos) support. - -## Python (`charmnumeric/`) - -### `charmnumeric/charmnumeric.py` — Core array types & operations (1063 lines) -| Symbol | Line | Description | -|--------|------|-------------| -| `DType` | 34 | Type encoding — FLOAT32=0, FLOAT64=1, INT32=2, INT64=3; `from_numpy()`, `to_numpy()`, `promote()` | -| `ArrayOperation` | 139 | Opcodes: add=0, sub=1, mul=2, div=3, matmul=4, tanh=5, exp=6, tile=7, reduce=8, matmatmul=9, diag=10 | -| `ArrayRegion` | 153 | N-D region with start/stop/step per dim — `serialize()`:194, `shape()`:209, `compose()`:212, `overlaps()`:232, `covers()`:239, `intersect()`:246 | -| `Array(FrontendObject)` | 618 | Main array class — `wire_dtype()`:639, `get()`:683, `get_region()`/`__getitem__()`:690/697, `__setitem__()`:708, `dot()`:724, `matvec()`:727, `copy()`:730, `fill()`:735, `astype()`:738; operators: +,-,*,/,@ at 713-722 | -| `ArrayView(FrontendObjectView)` | 747 | View with same ops as Array — `get()`:796, `__getitem__()`:801, `__setitem__()`:812, `dot()`:828 | -| `_full_region()` | 258 | Get/cache full region for shape | -| `_expand_key()` | 272 | Normalize indexing keys (handles newaxis, ellipsis) | -| `_parse_key()` | 308 | Convert key → ArrayRegion | -| `_array_add/sub/mul/truediv()` | 394-489 | Binary elementwise operator implementations | -| `_array_matmul()` | 492 | Matrix multiply dispatch (dot for 1D×1D, matvec for 2D×1D, matmatmul for 2D×2D) | -| `_matvec()` | 536 | Matrix-vector multiply | -| `_matmatmul()` | 595 | Matrix-matrix multiply (SUMMA) | -| `create_array()` | 851 | Array creation factory | -| `empty/zeros/ones/full()` | 931-971 | Array constructors | -| `empty_like/zeros_like/ones_like/full_like()` | 975-994 | Like-constructors | -| `asarray/array()` | 997-1028 | Conversion functions | -| `arange()` | 1042 | Range array | -| `eye()/identity()` | 1048/1061 | Identity matrix | - -### `charmnumeric/operations.py` — Standalone operation functions (180 lines) -| Symbol | Line | Description | -|--------|------|-------------| -| `tanh()` | 38 | Activation function | -| `exp()` | 45 | Exponential | -| `tile()` | 52 | Tile/repeat | -| `add/subtract/multiply/divide()` | 75-96 | Binary elementwise | -| `matmul()` | 99 | Matrix multiply | -| `dot()` | 105 | Dot product | -| `norm2()` | 116 | Squared L2 norm (1D) | -| `diag()` | 125 | Diagonal extraction/construction | - -### `charmnumeric/interface.py` — Cluster management (130 lines) -| Symbol | Line | Description | -|--------|------|-------------| -| `CharmNumericInterface(CCSInterface)` | 12 | Subclass with `from_bytes()` deserializer | -| `LocalCluster` | 26 | Local server launcher — `_run_server()`:74, `_connect_with_retry()`:94, `close()`:111 | - -### `charmnumeric/random.py` — Random arrays (6 lines) -| Symbol | Line | Description | -|--------|------|-------------| -| `randn()` | 4 | Random normal array | - -## C++ Backend (`example/charmnumeric/src/`) - -### `backend.hpp` — Core backend types (673 lines) -| Symbol | Line | Description | -|--------|------|-------------| -| `CT_MIN_TILE_1D/2D/3D` | 11-17 | Minimum tile sizes | -| `PartitionTraits` | 22 | Maps N → Charm++ proxy/index types (specializations for 1,2,3) | -| `proxy_at()` | 44 | Access chare array element by index | -| `CTArrayBase` | 93 | Type-erased N-D array base — name, region, decomp, global_shape, dtype | -| `Array` | 113 | Typed array — Kokkos views (d_view, h_view), `data_ptr()`, `copyToHost/Device()` | -| `RemoteBuffer` | 206 | Remote data buffer with region | -| `PendingComm` | 213 | Pending communication state | -| `ArrayDAGGroup` | 240 | Charm++ group extending `DAGGroup` | -| — `PartitionGrid` | 245 | Grid dimensions + epoch | -| — `ArrayMetadata` | 267 | Per-array metadata with decomp — `decomp()`:276 | -| — Key methods | | `receive_dag()`:303, `execute_node_nd()`:308, `compute_decompositions()`:314, `compile_node()`:317, `compile()`:318, `gather()`:305 | -| `ArrayDAGExecutorND` | 323 | N-D DAG executor — `execute_dag_node()`, `execute_matmul_node()`, `execute_matmatmul_node()`, `execute_reduce_node()`, `execute_diag_node()`, `execute_tile_node()`, `ast_visitor()` | -| `PartitionImpl` | 352 | Partition logic (Charm++-independent) | -| — `FreeKey`/`FreeKeyHash` | 370/387 | Array reuse key (dtype+region+shape) | -| — Key methods | | `init()`:567, `create()`:571, `run()`:572, `retire_array()`:428, `try_reuse()`:456, `allocate_or_reuse()`:502, `ensure_array()`:534, `process_get()`:582 | -| `Partition1D/2D/3D` | 587-672 | Thin Charm++ chare wrappers around PartitionImpl | - -### `backend_internal.hpp` — Compute kernels & helpers (491 lines) -| Symbol | Line | Description | -|--------|------|-------------| -| `eigen_gemv/gemv_sub/gemv_sub_3d()` | 15/28/48 | Eigen GEMV kernels | -| `eigen_dot()` | 88 | Dot product | -| `eigen_gemm/gemm_sub()` | 98/114 | Eigen GEMM kernels | -| `ct_min_tile()` | 151 | Min tile by dimension | -| `array_tile()` | 168 | Tile size from metadata | -| `extract_regions_nd()` | 179 | Extract I/O regions from AST | -| `determine_dtype()` | 185 | Get dtype from node/arrays | -| `cross_matmul_send_result_3d_to_1d()` | 188 | Cross-dim matmul send | -| `cross_matmul_send_result_2d_to_1d()` | 223 | Cross-dim matmul send | -| `cross_set_region_send()` | 320 | Cross-dim SET_REGION | -| `ReduceContrib` | 264 | 24-byte dot reduction contribution | -| `reduce_dot_sum` | 274 | Custom reducer for dot products | - -### `decomposition_solver.hpp` — Standalone decomposition solver (1716 lines) -| Symbol | Line | Description | -|--------|------|-------------| -| `decomposition_solver::compute_decompositions()` | 95 | Charm-free decomposition solver used by the runtime and standalone tests | -| `detail::for_each_node_topo()` | 59 | Topological DAG traversal helper for solver passes | -| `collect_unionable_nonfixed_leaves()` | 515 | Union-find guard: only unions zero-shift, same-tile, non-fixed temp leaves | - -### `array_region.hpp` — N-D regions & decomposition (449+ lines) -| Symbol | Line | Description | -|--------|------|-------------| -| `ArrayRegion` | 22 | N-D region (start/stop/step per dim) — `size()`:110, `deserialize()`:121, `overlaps()`:136, `covers()`:158, `intersect()`:186 | -| `ArrayDecomp` | 294 | Tile decomposition — `default_decomp()`:303, `offset_decomp()`:312, `to_global/local()`:322/339, `num_chares()`:350, `owning_chare()`:357, `chare_region_global/local()`:360/371 | -| `MemRef` | 407 | MLIR memref descriptor (allocated, aligned, offset, sizes, strides) | -| `ChareIndex` | 437 | N-D chare array index | - -### `opcodes.hpp` — Operation codes (50 lines) -| Symbol | Line | Description | -|--------|------|-------------| -| `Opcode` | 3 | COPY=-5..DIAG=10 | -| `is_binary_elementwise()` | 23 | ADD/SUB/MUL/DIV | -| `is_unary_elementwise()` | 36 | TANH/EXP | -| `is_elementwise()` | 47 | Binary or unary elementwise | - -### `jit.hpp` / `jit.cpp` — MLIR JIT compiler (150+ lines each) -| Symbol | Line | Description | -|--------|------|-------------| -| `BinaryOpEntry` | 80 | Registry: float_body, int_body function pointers | -| `UnaryOpEntry` | 87 | Registry: float_body, int_body function pointers | -| `MLIRJitCompiler` | 101 | AST → MLIR → executable — `buildFromAST()`, `optimizeAndFuse()`, `generateCPU()`, `loadCPU()`, `generateNVIDIA/AMD/Intel()` | - -### `dispatch.hpp` — Kernel dispatch & GPU callbacks (200+ lines) -| Symbol | Line | Description | -|--------|------|-------------| -| `CT_COMPUTE_POLICY/CT_COMM_POLICY` | 11-19 | Kokkos execution policies | -| `dispatch_kernel()` | 138 | Dispatch JIT kernel with memref descriptors | -| `compute_done_cb()` | 42 | (GPU) HAPI callback for compute completion | -| `deferred_send_cb()` | 83 | (GPU) HAPI callback for deferred send | - -### Key C++ source files -| File | Description | -|------|-------------| -| `server.cpp` | Main chare, creates ArrayDAGGroup | -| `dag_group_runtime.cpp` | ArrayDAGGroup constructor, `compute_decompositions()`, set_proxies, receive_get_request, gather | -| `dag_group_compile.cpp` | JIT compilation orchestration | -| `dag_group_receive.cpp` | DAG reception and dispatch | -| `execute_node.cpp` | `execute_node_nd()` — dispatches JIT kernels, handles MATMUL/REDUCE/TILE/DIAG | -| `execute_node_regions.cpp` | Region extraction from AST | -| `executor_core.cpp` | `execute_dag_node()`, remote input sends | -| `executor_matmul.cpp` | Matrix-vector multiply execution | -| `executor_matmatmul.cpp` | Matrix-matrix multiply (SUMMA) execution | -| `executor_reduce.cpp` | Reduce node execution | -| `executor_reducer.cpp` | Custom reduction operations | -| `executor_transfer.cpp` | Cross-partition data transfer | -| `executor_incremental.cpp` | Incremental computation | -| `executor_ast_visitor.cpp` | AST visitor for elementwise node execution | -| `partition_lifecycle.cpp` | `init()`, `create()`, `run()`, destructor | -| `partition_comm.cpp` | `receive_data()` inter-chare communication | -| `partition_get.cpp` | `process_get()` — gather results to PE 0 | -| `jit.cpp` | MLIRJitCompiler implementation — buildFromAST, emitOp | - -## Standalone C++ Tests (`example/charmnumeric/tests/`) - -### `test_compute_decompositions.cpp` — Standalone decomposition coverage (416 lines) -| Symbol | Line | Description | -|--------|------|-------------| -| `test_shifted_copy_uses_absolute_offset()` | 128 | Checks lifted absolute offset for a simple shifted copy | -| `test_shifted_expression_temp_uses_absolute_representative()` | 177 | Verifies shifted RHS temp keeps absolute representative `150` | -| `test_jacobi1d_temp_offset()` | 207 | 1D Jacobi temp regression: offset should be `1` | -| `test_jacobi2d_temp_offset()` | 259 | 2D Jacobi temp regression: offset should be `(1,1)` | -| `test_jacobi3d_temp_offset()` | 317 | 3D Jacobi temp regression: offset should be `(1,1,1)` | -| `test_strided_copy_reduces_tile_and_sets_offset()` | 374 | Stride-2 access regression: tile reduction and absolute offset lifting | - -### `backend.ci` — Charm++ interface (86 lines) -Module `charmnumeric` (depends on `charmtyles`). Mainchare `Main`. Group `ArrayDAGGroup : DAGGroup` with entries: `set_proxies()`, `proxies_ready()` [reduction], `receive_dag()`, `receive_get_request()`, `gather()`. Arrays `Partition1D/2D/3D` with entries: `run()`, `receive_data()` [nocopy/nocopydevice], `send_complete()`, `comm_done()`, `reduce_result()`. diff --git a/README.rst b/README.rst new file mode 100644 index 0000000..9e4a5b9 --- /dev/null +++ b/README.rst @@ -0,0 +1,104 @@ +charmtiles +========== + +:code:`charmtiles` is a python interface to a C++ distributed array library +implemented using Charm++ [#charm]_. +charmtiles uses a client-server model with a client-side python +interface and a Charm++ server on the backend. The client and server +are connected using CCS [#ccs]_. +The server maintains a symbol table of distributed arrays which +are then looked up for computation when a CCS message is +received. + +:code:`charmtiles.array` +---------------------- + +.. highlight:: python + +:code:`charmtiles.array.ndarray`, analogous to :code:`numpy.ndarray`, is a proxy +object that wraps the name of the corresponding array on the server. +We use a lazy evaluation scheme for array computations. +The array operations incrementally build an AST which is stored in a buffer in the +:code:`ndarray` object. This AST is encoded into a CCS message when +either the data from the array is requested on the frontend or +when the size of the AST grows beyond a user configurable +threshold. +The server side Charm++ program decodes the CCS message and +rebuilds the AST which is then executed. + +The lazy evaluation scheme reduces the number of CCS messages required to +be sent from the client to the server. +It also helps in reducing the number of temporary arrays created on the +server side by accessing the reference counts of the frontend arrays in +the python runtime. For example:: + + v = ndarray(1, 10, np.float64) + b = ndarray(1, 10, np.float64) + c = ndarray(1, 10, np.float64) + w = c + for i in range(2): + y = v + b + w + z = v - y + w = 2 * (c - z) + b + w.evaluate() + +The above code snippet generates the following AST. Nodes with labels +starting with the letter :code:`a` are arrays. Nodes with an operation +label that are colored blue are operations that generate a temporary +array. Nodes with an operation label that are colored red are operations +that generate arrays that are to be stored on the server side. +The red node labels also show the name of the resulting array. +Note that the arrays :code:`y`, :code:`z` and :code:`w` for the first iteration +of the loop are considered to be temporary because they are overwritten +in the next iteration. These operations will be executed inplace +on the server side. + +.. figure:: docs/images/simple_ast.png + :alt: simple_ast + + *AST generated by the above code snippet* + +Here's another example of a conjugate gradient solver:: + + def solve(A, b): + x = ndarray(1, 1000, np.float64) + r = b - A @ x + p = r.copy() + rsold = r @ r + + for i in range(1000): + Ap = A @ p + alpha = rsold / (p @ Ap) + + x = lg.axpy(alpha, p, x) + r = lg.axpy(alpha, Ap, r, multiplier=-1.) + + rsnew = r @ r + + if np.sqrt(rsnew.get()) < 1e-8: + print("Converged in %i iterations" % (i + 1)) + break + + p = lg.axpy(rsnew / rsold, p, r) + rsold = rsnew + + return x + +This generates the following AST, + +.. figure:: docs/images/conj_ast.png + :alt: conj_ast + + *AST generated by the conjugate gradient example* + +Here the green nodes are arrays that do not exist on the server when the AST is +sent, but will be created and stored as a result of an operation in the current +AST before being referenced. + + +References +---------- + +.. [#charm] Charm++ Documentation - https://charm.readthedocs.io/en/latest/ +.. [#ccs] CCS Documentation - https://charm.readthedocs.io/en/latest/converse/manual.html?converse-client-server-interface#converse-client-server-interface + diff --git a/charmnumeric/__init__.py b/charmnumeric/__init__.py deleted file mode 100644 index 00c90d2..0000000 --- a/charmnumeric/__init__.py +++ /dev/null @@ -1,85 +0,0 @@ -import numpy as np - -from . import random -from .charmnumeric import ( - Array, - ArrayView, - DType, - arange, - array, - asarray, - copy, - copyto, - create_array, - empty, - empty_like, - eye, - full, - full_like, - identity, - ones, - ones_like, - zeros, - zeros_like, -) -from .interface import CharmNumericInterface, LocalCluster -from .operations import add, diag, divide, dot, exp, matmul, multiply, norm2, subtract, tanh, tile - - -__version__ = "0.1.dev" - -ndarray = Array -float32 = np.float32 -float64 = np.float64 -int32 = np.int32 -int64 = np.int64 -dtype = np.dtype -newaxis = None -pi = np.pi -e = np.e - -__all__ = [ - "__version__", - "Array", - "ArrayView", - "CharmNumericInterface", - "DType", - "LocalCluster", - "add", - "arange", - "array", - "asarray", - "copy", - "copyto", - "create_array", - "diag", - "divide", - "dot", - "dtype", - "e", - "empty", - "empty_like", - "exp", - "eye", - "float32", - "float64", - "full", - "full_like", - "identity", - "int32", - "int64", - "matmul", - "multiply", - "ndarray", - "newaxis", - "norm2", - "ones", - "ones_like", - "pi", - "random", - "subtract", - "tanh", - "tile", - "zeros", - "zeros_like", -] diff --git a/charmnumeric/_native_region.pyx b/charmnumeric/_native_region.pyx deleted file mode 100644 index dcaaafd..0000000 --- a/charmnumeric/_native_region.pyx +++ /dev/null @@ -1,166 +0,0 @@ -# distutils: language = c++ -# cython: language_level=3 - -from charmtyles._native_region cimport CPPRegion, NativeRegion - -cdef extern from "array_region.hpp": - void* make_array_region_handle_1 "make_array_region_handle<1>"(const int*, const int*, const int*, bint) - void* make_array_region_handle_2 "make_array_region_handle<2>"(const int*, const int*, const int*, bint) - void* make_array_region_handle_3 "make_array_region_handle<3>"(const int*, const int*, const int*, bint) - void* make_array_region_handle_4 "make_array_region_handle<4>"(const int*, const int*, const int*, bint) - void* make_array_region_handle_5 "make_array_region_handle<5>"(const int*, const int*, const int*, bint) - void* make_array_region_handle_6 "make_array_region_handle<6>"(const int*, const int*, const int*, bint) - void* make_array_region_handle_7 "make_array_region_handle<7>"(const int*, const int*, const int*, bint) - void* make_array_region_handle_8 "make_array_region_handle<8>"(const int*, const int*, const int*, bint) - - void delete_array_region_handle_1 "delete_array_region_handle<1>"(void*) - void delete_array_region_handle_2 "delete_array_region_handle<2>"(void*) - void delete_array_region_handle_3 "delete_array_region_handle<3>"(void*) - void delete_array_region_handle_4 "delete_array_region_handle<4>"(void*) - void delete_array_region_handle_5 "delete_array_region_handle<5>"(void*) - void delete_array_region_handle_6 "delete_array_region_handle<6>"(void*) - void delete_array_region_handle_7 "delete_array_region_handle<7>"(void*) - void delete_array_region_handle_8 "delete_array_region_handle<8>"(void*) - - bint overlaps_array_region_handle_1 "overlaps_array_region_handle<1>"(void*, void*) - bint overlaps_array_region_handle_2 "overlaps_array_region_handle<2>"(void*, void*) - bint overlaps_array_region_handle_3 "overlaps_array_region_handle<3>"(void*, void*) - bint overlaps_array_region_handle_4 "overlaps_array_region_handle<4>"(void*, void*) - bint overlaps_array_region_handle_5 "overlaps_array_region_handle<5>"(void*, void*) - bint overlaps_array_region_handle_6 "overlaps_array_region_handle<6>"(void*, void*) - bint overlaps_array_region_handle_7 "overlaps_array_region_handle<7>"(void*, void*) - bint overlaps_array_region_handle_8 "overlaps_array_region_handle<8>"(void*, void*) - - bint covers_array_region_handle_1 "covers_array_region_handle<1>"(void*, void*) - bint covers_array_region_handle_2 "covers_array_region_handle<2>"(void*, void*) - bint covers_array_region_handle_3 "covers_array_region_handle<3>"(void*, void*) - bint covers_array_region_handle_4 "covers_array_region_handle<4>"(void*, void*) - bint covers_array_region_handle_5 "covers_array_region_handle<5>"(void*, void*) - bint covers_array_region_handle_6 "covers_array_region_handle<6>"(void*, void*) - bint covers_array_region_handle_7 "covers_array_region_handle<7>"(void*, void*) - bint covers_array_region_handle_8 "covers_array_region_handle<8>"(void*, void*) - - bint intersect_array_region_handle_1 "intersect_array_region_handle<1>"(void*, void*, int*, int*, int*) - bint intersect_array_region_handle_2 "intersect_array_region_handle<2>"(void*, void*, int*, int*, int*) - bint intersect_array_region_handle_3 "intersect_array_region_handle<3>"(void*, void*, int*, int*, int*) - bint intersect_array_region_handle_4 "intersect_array_region_handle<4>"(void*, void*, int*, int*, int*) - bint intersect_array_region_handle_5 "intersect_array_region_handle<5>"(void*, void*, int*, int*, int*) - bint intersect_array_region_handle_6 "intersect_array_region_handle<6>"(void*, void*, int*, int*, int*) - bint intersect_array_region_handle_7 "intersect_array_region_handle<7>"(void*, void*, int*, int*, int*) - bint intersect_array_region_handle_8 "intersect_array_region_handle<8>"(void*, void*, int*, int*, int*) - - -cdef int _MAX_NDIMS = 8 - - -cdef inline void _fill_buffer(tuple values, int* out, int ndims): - cdef int i - for i in range(ndims): - out[i] = values[i] - - -cdef class NativeArrayRegion(NativeRegion): - cdef void* _ptr - cdef int _ndims - - def __cinit__(self, tuple start, tuple stop, tuple step, bint is_global=False): - cdef int start_buf[8] - cdef int stop_buf[8] - cdef int step_buf[8] - - self._ptr = NULL - self._ndims = len(start) - if self._ndims < 1 or self._ndims > _MAX_NDIMS: - raise ValueError(f"Native ArrayRegion wrapper supports 1-{_MAX_NDIMS} dims") - if len(stop) != self._ndims or len(step) != self._ndims: - raise ValueError("start, stop, and step must have the same rank") - - _fill_buffer(start, start_buf, self._ndims) - _fill_buffer(stop, stop_buf, self._ndims) - _fill_buffer(step, step_buf, self._ndims) - - if self._ndims == 1: - self._ptr = make_array_region_handle_1(start_buf, stop_buf, step_buf, is_global) - elif self._ndims == 2: - self._ptr = make_array_region_handle_2(start_buf, stop_buf, step_buf, is_global) - elif self._ndims == 3: - self._ptr = make_array_region_handle_3(start_buf, stop_buf, step_buf, is_global) - elif self._ndims == 4: - self._ptr = make_array_region_handle_4(start_buf, stop_buf, step_buf, is_global) - elif self._ndims == 5: - self._ptr = make_array_region_handle_5(start_buf, stop_buf, step_buf, is_global) - elif self._ndims == 6: - self._ptr = make_array_region_handle_6(start_buf, stop_buf, step_buf, is_global) - elif self._ndims == 7: - self._ptr = make_array_region_handle_7(start_buf, stop_buf, step_buf, is_global) - else: - self._ptr = make_array_region_handle_8(start_buf, stop_buf, step_buf, is_global) - self._region_ptr = self._ptr - - def __dealloc__(self): - if self._ptr == NULL: - return - if self._ndims == 1: - delete_array_region_handle_1(self._ptr) - elif self._ndims == 2: - delete_array_region_handle_2(self._ptr) - elif self._ndims == 3: - delete_array_region_handle_3(self._ptr) - elif self._ndims == 4: - delete_array_region_handle_4(self._ptr) - elif self._ndims == 5: - delete_array_region_handle_5(self._ptr) - elif self._ndims == 6: - delete_array_region_handle_6(self._ptr) - elif self._ndims == 7: - delete_array_region_handle_7(self._ptr) - else: - delete_array_region_handle_8(self._ptr) - self._region_ptr = NULL - self._ptr = NULL - - cdef void _check_compatible(self, NativeArrayRegion other): - if self._ndims != other._ndims: - raise ValueError("ArrayRegion ranks must match") - - cpdef object intersect(self, NativeArrayRegion other): - cdef int out_start[8] - cdef int out_stop[8] - cdef int out_step[8] - cdef int i - cdef bint ok - cdef list start_values - cdef list stop_values - cdef list step_values - - self._check_compatible(other) - - if self._ndims == 1: - ok = intersect_array_region_handle_1(self._ptr, other._ptr, out_start, out_stop, out_step) - elif self._ndims == 2: - ok = intersect_array_region_handle_2(self._ptr, other._ptr, out_start, out_stop, out_step) - elif self._ndims == 3: - ok = intersect_array_region_handle_3(self._ptr, other._ptr, out_start, out_stop, out_step) - elif self._ndims == 4: - ok = intersect_array_region_handle_4(self._ptr, other._ptr, out_start, out_stop, out_step) - elif self._ndims == 5: - ok = intersect_array_region_handle_5(self._ptr, other._ptr, out_start, out_stop, out_step) - elif self._ndims == 6: - ok = intersect_array_region_handle_6(self._ptr, other._ptr, out_start, out_stop, out_step) - elif self._ndims == 7: - ok = intersect_array_region_handle_7(self._ptr, other._ptr, out_start, out_stop, out_step) - else: - ok = intersect_array_region_handle_8(self._ptr, other._ptr, out_start, out_stop, out_step) - - if not ok: - return None - - start_values = [0] * self._ndims - stop_values = [0] * self._ndims - step_values = [0] * self._ndims - for i in range(self._ndims): - start_values[i] = out_start[i] - stop_values[i] = out_stop[i] - step_values[i] = out_step[i] - - return (tuple(start_values), tuple(stop_values), tuple(step_values)) diff --git a/charmnumeric/charmnumeric.py b/charmnumeric/charmnumeric.py deleted file mode 100644 index 90cd0ae..0000000 --- a/charmnumeric/charmnumeric.py +++ /dev/null @@ -1,1062 +0,0 @@ -import numbers - -import numpy as np - -from charmtyles.core import ( - OverlapType, - Region, - FrontendObject, - FrontendObjectView, - create_object, - execute, - _shape_size, -) -from charmtyles.operation import Operation, fusible_op, nonfusible_op -from charmtyles.interface import to_bytes - -from ._native_region import NativeArrayRegion - - -_FULL_REGION_CACHE = {} -_DEFAULT_DTYPE = np.dtype("float64") - - -class _CallableInt(int): - """Int-like object that still supports legacy ``obj.size()`` calls.""" - - def __new__(cls, value): - return int.__new__(cls, int(value)) - - def __call__(self): - return int(self) - - -class DType(object): - """Element type encoding shared with the C++ wire protocol.""" - - FLOAT32 = 0 - FLOAT64 = 1 - INT32 = 2 - INT64 = 3 - - _from_numpy = { - np.dtype("float32"): FLOAT32, - np.dtype("float64"): FLOAT64, - np.dtype("int32"): INT32, - np.dtype("int64"): INT64, - } - - @staticmethod - def from_numpy(np_dtype): - """Convert a supported numpy dtype to the wire encoding.""" - - dtype = _normalize_dtype(np_dtype) - return DType._from_numpy[dtype] - - _to_numpy = { - FLOAT32: np.dtype("float32"), - FLOAT64: np.dtype("float64"), - INT32: np.dtype("int32"), - INT64: np.dtype("int64"), - } - - @staticmethod - def to_numpy(wire_dtype): - """Convert a wire dtype encoding back to a numpy dtype.""" - - if wire_dtype not in DType._to_numpy: - raise TypeError(f"Unsupported wire dtype {wire_dtype!r}") - return DType._to_numpy[wire_dtype] - - @staticmethod - def promote(a, b): - """Return the promoted wire dtype for operands *a* and *b*.""" - - return DType.from_numpy(np.result_type(DType.to_numpy(a), DType.to_numpy(b))) - - -def _normalize_dtype(dtype, *, allow_none=False): - if dtype is None: - if allow_none: - return None - return _DEFAULT_DTYPE - - np_dtype = np.dtype(dtype) - if np_dtype not in DType._from_numpy: - raise TypeError( - "charmnumeric only supports float32, float64, int32, and int64; " - f"got {np_dtype}" - ) - return np_dtype - - -def _normalize_shape(shape): - if isinstance(shape, numbers.Integral): - normalized = (int(shape),) - else: - normalized = tuple(int(dim) for dim in shape) - - for dim in normalized: - if dim < 0: - raise ValueError(f"negative dimensions are not allowed: {normalized}") - return normalized - - -def _coerce_scalar(value): - if isinstance(value, np.ndarray): - if value.ndim != 0: - return value - value = value.item() - elif isinstance(value, np.generic): - value = value.item() - - if isinstance(value, bool): - return int(value) - if isinstance(value, numbers.Integral): - return int(value) - if isinstance(value, numbers.Real): - return float(value) - return value - - -def _is_scalar_like(value): - if isinstance(value, np.ndarray): - return value.ndim == 0 - if isinstance(value, np.generic): - return True - return isinstance(value, (bool, numbers.Real)) - - -def _apply_ndmin(shape, ndmin): - ndmin = int(ndmin) - if ndmin < 0: - raise ValueError("ndmin must be non-negative") - if len(shape) >= ndmin: - return tuple(shape) - return (1,) * (ndmin - len(shape)) + tuple(shape) - - -class ArrayOperation(Operation): - add = 0 - sub = 1 - mul = 2 - div = 3 - matmul = 4 - tanh = 5 - exp = 6 - tile = 7 - reduce = 8 - matmatmul = 9 - diag = 10 - - -class ArrayRegion(Region): - __slots__ = ( - "start", - "stop", - "step", - "_shape", - "_hash", - "_serialized", - "_native", - ) - - def __init__(self, start, stop, step): - super().__init__() - self.start = tuple(start) - self.stop = tuple(stop) - self.step = tuple(step) - self._shape = tuple( - (self.stop[i] - self.start[i] + self.step[i] - 1) // self.step[i] - for i in range(len(self.start)) - ) - self._hash = hash((self.start, self.stop, self.step)) - self._serialized = None - self._native = NativeArrayRegion(self.start, self.stop, self.step, False) - - def __str__(self): - return f"({self.start}, {self.stop}, {self.step})" - - def __hash__(self): - return self._hash - - def __eq__(self, value): - if not isinstance(value, Region): - return False - if self.is_global or value.is_global: - return True - return ( - self.start == value.start - and self.stop == value.stop - and self.step == value.step - ) - - def serialize(self): - if self._serialized is not None: - return self._serialized - - payload = bytearray() - payload.extend(to_bytes(1 if self.is_global else 0, "i")) - if not self.is_global: - payload.extend(to_bytes(len(self.start), "i")) - for i in range(len(self.start)): - payload.extend(to_bytes(self.start[i], "i")) - payload.extend(to_bytes(self.stop[i], "i")) - payload.extend(to_bytes(self.step[i], "i")) - self._serialized = bytes(payload) - return self._serialized - - def shape(self): - return self._shape - - def compose(self, inner): - """Compose an inner sub-region within this outer region.""" - - if ( - inner.start == (0,) * len(inner.start) - and inner.step == (1,) * len(inner.step) - and inner.stop == self._shape - ): - return self - - ndims = len(inner.start) - new_start = [0] * ndims - new_stop = [0] * ndims - new_step = [0] * ndims - for d in range(ndims): - new_start[d] = self.start[d] + inner.start[d] * self.step[d] - new_stop[d] = self.start[d] + inner.stop[d] * self.step[d] - new_step[d] = self.step[d] * inner.step[d] - return ArrayRegion(new_start, new_stop, new_step) - - def overlaps(self, other): - if not isinstance(other, ArrayRegion): - raise TypeError( - f"ArrayRegion overlap requires ArrayRegion, got {type(other).__name__}" - ) - return self._native.overlaps(other._native) - - def covers(self, other): - if not isinstance(other, ArrayRegion): - raise TypeError( - f"ArrayRegion cover check requires ArrayRegion, got {type(other).__name__}" - ) - return self._native.covers(other._native) - - def intersect(self, other): - if not isinstance(other, ArrayRegion): - raise TypeError( - f"ArrayRegion intersection requires ArrayRegion, got {type(other).__name__}" - ) - result = self._native.intersect(other._native) - if result is None: - return None - start, stop, step = result - return ArrayRegion(start, stop, step) - - -def _full_region(shape): - key = tuple(shape) - region = _FULL_REGION_CACHE.get(key) - if region is None: - ndims = len(key) - region = ArrayRegion( - start=(0,) * ndims, - stop=key, - step=(1,) * ndims, - ) - _FULL_REGION_CACHE[key] = region - return region - - -def _expand_key(key, ndims): - if key is Ellipsis: - key = (Ellipsis,) - elif not isinstance(key, tuple): - key = (key,) - - normalized = [] - saw_ellipsis = False - explicit_dims = sum(1 for item in key if item is not Ellipsis and item is not None) - - for item in key: - if item is Ellipsis: - if saw_ellipsis: - raise IndexError("an index can only have a single ellipsis") - saw_ellipsis = True - fill = ndims - explicit_dims - if fill < 0: - raise IndexError("too many indices for array") - normalized.extend(slice(None) for _ in range(fill)) - elif item is None: - raise NotImplementedError("np.newaxis is not supported yet") - else: - normalized.append(item) - - if len(normalized) > ndims: - raise IndexError( - f"too many indices for array: array is {ndims}-dimensional, " - f"but {len(normalized)} were indexed" - ) - - if len(normalized) < ndims: - normalized.extend(slice(None) for _ in range(ndims - len(normalized))) - - return tuple(normalized) - - -def _parse_key(key, shape, ndims): - """Normalize a ``__getitem__`` or ``__setitem__`` key into an ArrayRegion.""" - - key = _expand_key(key, ndims) - start = [0] * ndims - stop = [0] * ndims - step = [1] * ndims - - for i, item in enumerate(key): - if isinstance(item, slice): - step_i = 1 if item.step is None else int(item.step) - if step_i == 0: - raise ValueError("slice step cannot be zero") - if step_i < 0: - raise NotImplementedError("negative slicing is not supported yet") - start_i, stop_i, step_i = item.indices(shape[i]) - elif isinstance(item, numbers.Integral): - start_i = int(item) - if start_i < 0: - start_i += shape[i] - if start_i < 0 or start_i >= shape[i]: - raise IndexError( - f"index {item} is out of bounds for axis {i} with size {shape[i]}" - ) - stop_i = start_i + 1 - step_i = 1 - else: - raise TypeError( - "only integers, slices, and ellipsis are valid charmnumeric indices" - ) - - start[i] = start_i - stop[i] = stop_i - step[i] = step_i - - return ArrayRegion(start, stop, step) - - -def _broadcast_shape(lhs_shape, rhs): - if isinstance(rhs, (Array, ArrayView)): - rhs_shape = rhs.shape - else: - return tuple(lhs_shape) - - try: - return tuple(np.broadcast_shapes(lhs_shape, rhs_shape)) - except ValueError as exc: - raise ValueError( - f"operands could not be broadcast together with shapes " - f"{lhs_shape} and {rhs_shape}" - ) from exc - - -def _operand_dtype(value): - if isinstance(value, (Array, ArrayView)): - return value.dtype - # Scalars do not promote the array dtype (matches numpy value-based casting: - # a Python float does not upcast a float32 array to float64). - return None - - -def _result_dtype(lhs, rhs): - other_dtype = _operand_dtype(rhs) - if other_dtype is None: - return lhs.dtype - return _normalize_dtype(np.result_type(lhs.dtype, other_dtype)) - - -def _coerce_operand(value): - if isinstance(value, (Array, ArrayView)): - return value - if _is_scalar_like(value): - return _coerce_scalar(value) - if isinstance(value, (str, bytes)): - return value - - try: - host = np.asarray(value) - except Exception: - return value - - if host.ndim == 0: - return _coerce_scalar(host.item()) - return asarray(host) - - -def _array_add(self, other): - other = _coerce_operand(other) - return fusible_op( - self, - other, - operation=ArrayOperation.add, - shape=_broadcast_shape(self.shape, other), - dtype=_result_dtype(self, other), - ) - - -def _array_radd(self, other): - other = _coerce_operand(other) - return fusible_op( - other, - self, - operation=ArrayOperation.add, - shape=_broadcast_shape(self.shape, other), - dtype=_result_dtype(self, other), - ) - - -def _array_sub(self, other): - other = _coerce_operand(other) - return fusible_op( - self, - other, - operation=ArrayOperation.sub, - shape=_broadcast_shape(self.shape, other), - dtype=_result_dtype(self, other), - ) - - -def _array_neg(self): - return fusible_op( - self, - -1, - operation=ArrayOperation.mul, - shape=self.shape, - dtype=self.dtype, - ) - - -def _array_rsub(self, other): - other = _coerce_operand(other) - return fusible_op( - other, - self, - operation=ArrayOperation.sub, - shape=_broadcast_shape(self.shape, other), - dtype=_result_dtype(self, other), - ) - - -def _array_mul(self, other): - other = _coerce_operand(other) - return fusible_op( - self, - other, - operation=ArrayOperation.mul, - shape=_broadcast_shape(self.shape, other), - dtype=_result_dtype(self, other), - ) - - -def _array_rmul(self, other): - other = _coerce_operand(other) - return fusible_op( - other, - self, - operation=ArrayOperation.mul, - shape=_broadcast_shape(self.shape, other), - dtype=_result_dtype(self, other), - ) - - -def _array_truediv(self, other): - other = _coerce_operand(other) - return fusible_op( - self, - other, - operation=ArrayOperation.div, - shape=_broadcast_shape(self.shape, other), - dtype=_result_dtype(self, other), - ) - - -def _array_rtruediv(self, other): - other = _coerce_operand(other) - return fusible_op( - other, - self, - operation=ArrayOperation.div, - shape=_broadcast_shape(self.shape, other), - dtype=_result_dtype(self, other), - ) - - -def _array_matmul(self, other): - other = _coerce_operand(other) - if not isinstance(other, (Array, ArrayView)): - raise ValueError("matmul does not support scalar operands") - - if self.ndims == 1 and other.ndims == 1: - if self.shape[0] != other.shape[0]: - raise ValueError( - f"Shape mismatch for dot product: ({self.shape[0]},) @ ({other.shape[0]},)" - ) - return nonfusible_op( - self, - other, - operation=ArrayOperation.reduce, - shape=(1,), - dtype=_result_dtype(self, other), - ) - - if self.ndims == 3: - dropped = [i for i in range(3) if self.shape[i] == 1] - if len(dropped) == 1: - if other.ndims == 1: - return _matvec(self, other) - if other.ndims >= 2: - return _matmatmul(self, other) - raise ValueError( - f"3D matmul requires exactly one singleton dim, got shape {self.shape}" - ) - - if self.ndims != 2: - raise ValueError(f"matmul requires a 2D matrix, got {self.ndims}D") - - if other.ndims == 1: - return _matvec(self, other) - if other.ndims >= 2: - if other.ndims == 2 and other.shape[1] == 1: - temp = create_array((other.shape[0],), dtype=other.dtype) - temp[:] = other - return _matvec(self, temp) - return _matmatmul(self, other) - - raise ValueError(f"matmul requires a 2D operand, got {other.ndims}D") - - -def _matvec(mat_obj, vec): - """Matrix-vector multiply.""" - - if not isinstance(vec, (Array, ArrayView)): - raise ValueError("matvec requires an Array vector") - - if vec.ndims == 2 and vec.shape[1] == 1: - temp = create_array((vec.shape[0],), dtype=vec.dtype) - temp[:] = vec - return _matvec(mat_obj, temp) - - if vec.ndims != 1: - raise ValueError("matvec requires a 1D vector (N,) or 2D column vector (N, 1)") - - if mat_obj.ndims == 2: - mat_rows, mat_cols = mat_obj.shape[0], mat_obj.shape[1] - elif mat_obj.ndims == 3: - dropped = [i for i in range(3) if mat_obj.shape[i] == 1] - if len(dropped) != 1: - raise ValueError( - "3D matvec requires exactly one singleton dimension, " - f"got shape {mat_obj.shape}" - ) - remaining = [i for i in range(3) if i not in dropped] - mat_rows = mat_obj.shape[remaining[0]] - mat_cols = mat_obj.shape[remaining[1]] - else: - raise ValueError(f"matvec requires a 2D or 3D matrix, got {mat_obj.ndims}D") - - if mat_cols != vec.shape[0]: - raise ValueError( - f"Shape mismatch: matrix cols {mat_cols} != vector size {vec.shape[0]}" - ) - - return nonfusible_op( - mat_obj, - vec, - operation=ArrayOperation.matmul, - shape=(mat_rows,), - dtype=_result_dtype(mat_obj, vec), - ) - - -def _matmatmul_shape(obj): - """Return ``(rows, cols)`` for a matmatmul operand.""" - - if obj.ndims == 2: - return obj.shape[0], obj.shape[1] - if obj.ndims == 3: - dropped = [i for i in range(3) if obj.shape[i] == 1] - if len(dropped) != 1: - raise ValueError( - f"3D matmatmul requires exactly one singleton dim, got shape {obj.shape}" - ) - remaining = [i for i in range(3) if i not in dropped] - return obj.shape[remaining[0]], obj.shape[remaining[1]] - raise ValueError(f"matmatmul requires 2D or 3D operand, got {obj.ndims}D") - - -def _matmatmul(lhs, rhs): - """Matrix-matrix multiply with inline region serialization.""" - - if not isinstance(rhs, (Array, ArrayView)): - raise ValueError("matmatmul requires Array operands") - - lhs_rows, lhs_cols = _matmatmul_shape(lhs) - rhs_rows, rhs_cols = _matmatmul_shape(rhs) - - if lhs_cols != rhs_rows: - raise ValueError( - f"Shape mismatch: ({lhs_rows},{lhs_cols}) @ ({rhs_rows},{rhs_cols})" - ) - - return nonfusible_op( - lhs, - rhs, - operation=ArrayOperation.matmatmul, - shape=(lhs_rows, rhs_cols), - dtype=_result_dtype(lhs, rhs), - ) - - -class Array(FrontendObject): - __array_priority__ = 1000 - - def __init__(self, **kwargs): - dtype = _normalize_dtype(kwargs.pop("dtype", None), allow_none=True) - self.dtype = dtype - self._wire_dtype = DType.from_numpy(self.dtype) if self.dtype is not None else 0 - - shape = kwargs.pop("shape", None) - if shape is not None: - shape = _normalize_shape(shape) - - super().__init__(**kwargs) - - self._shape = tuple(shape) if shape is not None else tuple(self.region.shape()) - self.ndims = len(self._shape) - self._size = _shape_size(self._shape) - - if self.region.is_global and self._shape is not None: - self.region = _full_region(self._shape) - - def wire_dtype(self): - return self._wire_dtype - - @property - def shape(self): - return self._shape - - @property - def ndim(self): - return self.ndims - - @property - def size(self): - return _CallableInt(self._size) - - @property - def itemsize(self): - return 0 if self.dtype is None else self.dtype.itemsize - - @property - def nbytes(self): - return int(self.size) * self.itemsize - - def __len__(self): - if self.ndims == 0: - raise TypeError("len() of unsized object") - return self._shape[0] - - def __bool__(self): - raise ValueError( - "The truth value of a charmnumeric array is ambiguous. " - "Use .get(interface) to materialize it first." - ) - - def __array__(self, dtype=None, copy=None): - raise TypeError( - "charmnumeric arrays are lazy frontend objects. " - "Call .get(interface) before converting to numpy." - ) - - def __repr__(self): - dtype_name = None if self.dtype is None else self.dtype.name - return f"Array(shape={self.shape}, dtype={dtype_name}, name={self.name})" - - def get(self, interface): - execute(interface) - data = interface.get(self.ndims, self.name, self.size(), dtype=self.dtype) - if self.ndims > 1: - return data.reshape(self._shape) - return data - - def get_region(self, region, **kwargs): - kwargs.setdefault("shape", region.shape()) - kwargs.setdefault("dtype", self.dtype) - kwargs.setdefault("wire_dtype", self._wire_dtype) - kwargs.setdefault("size", _shape_size(kwargs["shape"])) - return ArrayView(self, region, **kwargs) - - def __getitem__(self, key): - region = _parse_key(key, self._shape, self.ndims) - region_shape = region.shape() - return self.get_region( - region, - shape=region_shape, - dtype=self.dtype, - wire_dtype=self._wire_dtype, - size=_shape_size(region_shape), - ) - - def __setitem__(self, key, rhs): - region = _parse_key(key, self._shape, self.ndims) - rhs = _coerce_assignment_operand(rhs, region.shape(), self.dtype) - self.set_region(region, rhs, shape=region.shape(), dtype=self.dtype) - - __add__ = _array_add - __radd__ = _array_radd - __sub__ = _array_sub - __neg__ = _array_neg - __rsub__ = _array_rsub - __mul__ = _array_mul - __rmul__ = _array_rmul - __truediv__ = _array_truediv - __rtruediv__ = _array_rtruediv - __matmul__ = _array_matmul - - def dot(self, other): - return _array_matmul(self, other) - - def matvec(self, vec): - return _matvec(self, vec) - - def copy(self): - result = empty(self.shape, dtype=self.dtype) - result[...] = self - return result - - def fill(self, value): - self[...] = value - - def astype(self, dtype, copy=True): - dtype = _normalize_dtype(dtype) - if dtype != self.dtype: - raise NotImplementedError( - "dtype conversion for existing charmnumeric arrays is not implemented yet" - ) - return self.copy() if copy else self - - -class ArrayView(FrontendObjectView): - """Array-specific view with arithmetic operators and slicing.""" - - __array_priority__ = 1000 - - @property - def ndim(self): - return self.ndims - - @property - def size(self): - return _CallableInt(self._size) - - @property - def itemsize(self): - return 0 if self.dtype is None else self.dtype.itemsize - - @property - def nbytes(self): - return int(self.size) * self.itemsize - - def __len__(self): - if self.ndims == 0: - raise TypeError("len() of unsized object") - return self.shape[0] - - def __bool__(self): - raise ValueError( - "The truth value of a charmnumeric array is ambiguous. " - "Use .get(interface) to materialize it first." - ) - - def __array__(self, dtype=None, copy=None): - raise TypeError( - "charmnumeric arrays are lazy frontend objects. " - "Call .get(interface) before converting to numpy." - ) - - def __repr__(self): - dtype_name = None if self.dtype is None else self.dtype.name - return f"ArrayView(shape={self.shape}, dtype={dtype_name}, name={self.name})" - - def get_region(self, region, **kwargs): - kwargs.setdefault("shape", region.shape()) - kwargs.setdefault("dtype", self.dtype) - kwargs.setdefault("wire_dtype", self.wire_dtype()) - kwargs.setdefault("size", _shape_size(kwargs["shape"])) - return ArrayView(self, region, **kwargs) - - def get(self, interface): - materialized = empty(self.shape, dtype=self.dtype) - materialized[...] = self - return materialized.get(interface) - - def __getitem__(self, key): - region = _parse_key(key, self.shape, self.ndims) - region_shape = region.shape() - return self.get_region( - region, - shape=region_shape, - dtype=self.dtype, - wire_dtype=self.wire_dtype(), - size=_shape_size(region_shape), - ) - - def __setitem__(self, key, rhs): - region = _parse_key(key, self.shape, self.ndims) - rhs = _coerce_assignment_operand(rhs, region.shape(), self.dtype) - self.set_region(region, rhs, shape=region.shape(), dtype=self.dtype) - - __add__ = _array_add - __radd__ = _array_radd - __sub__ = _array_sub - __neg__ = _array_neg - __rsub__ = _array_rsub - __mul__ = _array_mul - __rmul__ = _array_rmul - __truediv__ = _array_truediv - __rtruediv__ = _array_rtruediv - __matmul__ = _array_matmul - - def dot(self, other): - return _array_matmul(self, other) - - def matvec(self, vec): - return _matvec(self, vec) - - def copy(self): - result = empty(self.shape, dtype=self.dtype) - result[...] = self - return result - - def fill(self, value): - self[...] = value - - def astype(self, dtype, copy=True): - dtype = _normalize_dtype(dtype) - if dtype != self.dtype: - raise NotImplementedError( - "dtype conversion for existing charmnumeric arrays is not implemented yet" - ) - return self.copy() if copy else self - - -def create_array(shape, dtype=_DEFAULT_DTYPE, **kwargs): - shape = _normalize_shape(shape) - dtype = _normalize_dtype(dtype) - return create_object( - Array, - region=ArrayRegion(start=(0,) * len(shape), stop=shape, step=(1,) * len(shape)), - dtype=dtype, - shape=shape, - **kwargs, - ) - - -def _assign_host_data(target, values): - host = np.asarray(values, dtype=target.dtype) - if host.shape != target.shape: - raise ValueError(f"cannot assign data with shape {host.shape} into shape {target.shape}") - - if host.size == 0: - return target - - first = host.reshape(-1)[0] - if host.size == 1 or np.all(host == first): - target[...] = _coerce_scalar(first) - return target - - for index in np.ndindex(host.shape): - target[index] = _coerce_scalar(host[index]) - return target - - -def _array_from_host_data(values, dtype=None, copy=True, ndmin=0): - requested_dtype = _normalize_dtype(dtype, allow_none=True) - - if copy: - host = np.array(values, dtype=requested_dtype, copy=True) - else: - host = np.asarray(values, dtype=requested_dtype) - - while host.ndim < ndmin: - host = np.expand_dims(host, axis=0) - - if host.ndim == 0: - raise NotImplementedError("0D charmnumeric arrays are not supported yet") - - host_dtype = _normalize_dtype(host.dtype) - result = create_array(host.shape, dtype=host_dtype) - return _assign_host_data(result, host) - - -def _shape_dtype_from_like(obj): - if isinstance(obj, (Array, ArrayView)): - return obj.shape, obj.dtype - - host = np.asarray(obj) - if host.ndim == 0: - raise NotImplementedError("0D charmnumeric arrays are not supported yet") - return host.shape, _normalize_dtype(host.dtype) - - -def _coerce_assignment_operand(rhs, shape, dtype): - if isinstance(rhs, (Array, ArrayView)): - return rhs - - if _is_scalar_like(rhs): - return _coerce_scalar(rhs) - - host = np.asarray(rhs, dtype=_normalize_dtype(dtype)) - if host.ndim == 0: - return _coerce_scalar(host.item()) - - try: - host = np.broadcast_to(host, shape) - except ValueError as exc: - raise ValueError( - f"could not broadcast input from shape {host.shape} into shape {shape}" - ) from exc - - return _array_from_host_data(host, dtype=host.dtype, copy=False) - - -def empty(shape, dtype=_DEFAULT_DTYPE): - return create_array(shape, dtype=dtype) - - -def zeros(shape, dtype=_DEFAULT_DTYPE): - return empty(shape, dtype=dtype) - - -def full(shape, fill_value, dtype=None): - shape = _normalize_shape(shape) - - if _is_scalar_like(fill_value): - target_dtype = _normalize_dtype(dtype, allow_none=True) - if target_dtype is None: - target_dtype = _normalize_dtype(np.asarray(fill_value).dtype) - result = empty(shape, dtype=target_dtype) - result[...] = _coerce_scalar(fill_value) - return result - - host = np.asarray(fill_value, dtype=_normalize_dtype(dtype, allow_none=True)) - if host.ndim == 0: - return full(shape, host.item(), dtype=dtype) - - target_dtype = _normalize_dtype(host.dtype) - try: - host = np.broadcast_to(host, shape) - except ValueError as exc: - raise ValueError( - f"could not broadcast fill_value from shape {host.shape} into shape {shape}" - ) from exc - - return _array_from_host_data(host, dtype=target_dtype, copy=False) - - -def ones(shape, dtype=_DEFAULT_DTYPE): - return full(shape, 1, dtype=dtype) - - -def empty_like(a, dtype=None, shape=None): - base_shape, base_dtype = _shape_dtype_from_like(a) - target_shape = base_shape if shape is None else _normalize_shape(shape) - target_dtype = base_dtype if dtype is None else _normalize_dtype(dtype) - return empty(target_shape, dtype=target_dtype) - - -def zeros_like(a, dtype=None, shape=None): - base_shape, base_dtype = _shape_dtype_from_like(a) - target_shape = base_shape if shape is None else _normalize_shape(shape) - target_dtype = base_dtype if dtype is None else _normalize_dtype(dtype) - return zeros(target_shape, dtype=target_dtype) - - -def ones_like(a, dtype=None, shape=None): - base_shape, base_dtype = _shape_dtype_from_like(a) - target_shape = base_shape if shape is None else _normalize_shape(shape) - target_dtype = base_dtype if dtype is None else _normalize_dtype(dtype) - return ones(target_shape, dtype=target_dtype) - - -def full_like(a, fill_value, dtype=None, shape=None): - base_shape, base_dtype = _shape_dtype_from_like(a) - target_shape = base_shape if shape is None else _normalize_shape(shape) - target_dtype = base_dtype if dtype is None else _normalize_dtype(dtype) - return full(target_shape, fill_value, dtype=target_dtype) - - -def asarray(a, dtype=None): - requested_dtype = _normalize_dtype(dtype, allow_none=True) - - if isinstance(a, (Array, ArrayView)): - if requested_dtype is None or requested_dtype == a.dtype: - return a - raise NotImplementedError( - "dtype conversion for existing charmnumeric arrays is not implemented yet" - ) - - return _array_from_host_data(a, dtype=requested_dtype, copy=False) - - -def array(a, dtype=None, copy=True, ndmin=0): - requested_dtype = _normalize_dtype(dtype, allow_none=True) - - if isinstance(a, (Array, ArrayView)): - target_dtype = a.dtype if requested_dtype is None else requested_dtype - if target_dtype != a.dtype: - raise NotImplementedError( - "dtype conversion for existing charmnumeric arrays is not implemented yet" - ) - - target_shape = _apply_ndmin(a.shape, ndmin) - if not copy and target_shape == a.shape: - return a - - result = empty(target_shape, dtype=target_dtype) - result[...] = a - return result - - return _array_from_host_data(a, dtype=requested_dtype, copy=copy, ndmin=ndmin) - - -def copy(a, order="K"): - return array(a, copy=True) - - -def copyto(dst, src): - if not isinstance(dst, (Array, ArrayView)): - raise TypeError("copyto destination must be a charmnumeric Array or ArrayView") - dst[...] = src - return dst - - -def arange(start, stop=None, step=1, dtype=None): - requested_dtype = _normalize_dtype(dtype, allow_none=True) - host = np.arange(start, stop=stop, step=step, dtype=requested_dtype) - return _array_from_host_data(host, dtype=host.dtype, copy=False) - - -def eye(N, M=None, k=0, dtype=_DEFAULT_DTYPE): - rows = int(N) - cols = rows if M is None else int(M) - result = zeros((rows, cols), dtype=dtype) - - row = max(0, -int(k)) - col = max(0, int(k)) - length = max(0, min(rows - row, cols - col)) - for offset in range(length): - result[row + offset, col + offset] = 1 - return result - - -def identity(n, dtype=_DEFAULT_DTYPE): - return eye(n, dtype=dtype) diff --git a/charmnumeric/interface.py b/charmnumeric/interface.py deleted file mode 100644 index c4db495..0000000 --- a/charmnumeric/interface.py +++ /dev/null @@ -1,129 +0,0 @@ -from __future__ import annotations - -from pathlib import Path -import subprocess -import time - -import numpy as np - -from charmtyles.interface import CCSInterface - - -class CharmNumericInterface(CCSInterface): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - - def from_bytes(self, payload, **kwargs): - reply_type = kwargs.pop("reply_type", "array") - if reply_type == "array": - dtype = kwargs.pop("dtype", int) - return np.frombuffer(payload, dtype=dtype) - if reply_type is None: - return None - return super().from_bytes(payload, reply_type=reply_type, **kwargs) - - -class LocalCluster(CharmNumericInterface): - """Launch a local charmnumeric backend and connect to it.""" - - def __init__( - self, - server_binary, - *, - server_ip="127.0.0.1", - server_port=1234, - odf=4, - max_pes=1, - charmrun=None, - workdir=None, - startup_timeout=20.0, - ): - self.server_binary = Path(server_binary).resolve() - self.server_ip = server_ip - self.server_port = int(server_port) - self.odf = int(odf) - self.max_pes = int(max_pes) - self.workdir = Path(workdir or self.server_binary.parent).resolve() - self.workdir.mkdir(parents=True, exist_ok=True) - self.nodelist_path = self.workdir / "localnodelist" - self.log_path = self.workdir / "server.log" - self.charmrun = self._resolve_charmrun(charmrun) - self.process = None - self._log_handle = None - - self._write_nodelist() - self._run_server() - super().__init__() - self._connect_with_retry(timeout=float(startup_timeout)) - - def _resolve_charmrun(self, charmrun): - if charmrun: - if Path(charmrun).is_absolute() or "/" in charmrun: - return str(Path(charmrun).resolve()) - return charmrun - - sibling = self.server_binary.parent / "charmrun" - if sibling.is_file(): - return str(sibling) - return "charmrun" - - def _write_nodelist(self): - data = "".join("host localhost\n" for _ in range(self.max_pes)) - self.nodelist_path.write_text(data) - - def _run_server(self): - cmd = [ - self.charmrun, - f"+p{self.max_pes}", - str(self.server_binary), - "++server", - "++server-port", - str(self.server_port), - "++nodelist", - str(self.nodelist_path), - ] - self._log_handle = self.log_path.open("w") - self.process = subprocess.Popen( - cmd, - cwd=self.workdir, - stdout=self._log_handle, - stderr=subprocess.STDOUT, - text=True, - ) - - def _connect_with_retry(self, timeout): - deadline = time.monotonic() + timeout - last_error = None - while time.monotonic() < deadline: - try: - self.connect(self.server_ip, self.server_port, self.odf) - return - except Exception as exc: # pragma: no cover - depends on local runtime - last_error = exc - time.sleep(1.0) - - self.close() - raise RuntimeError( - f"Timed out connecting to charmnumeric backend at " - f"{self.server_ip}:{self.server_port}. See {self.log_path}." - ) from last_error - - def close(self): - if self.process is not None and self.process.poll() is None: - self.process.terminate() - try: - self.process.wait(timeout=20) - except subprocess.TimeoutExpired: - self.process.kill() - self.process.wait(timeout=20) - self.process = None - - if self._log_handle is not None and not self._log_handle.closed: - self._log_handle.close() - self._log_handle = None - - def __enter__(self): - return self - - def __exit__(self, exc_type, exc, tb): - self.close() diff --git a/charmnumeric/operations.py b/charmnumeric/operations.py deleted file mode 100644 index 7359e1b..0000000 --- a/charmnumeric/operations.py +++ /dev/null @@ -1,179 +0,0 @@ -import numpy as np - -from charmtyles.operation import fusible_op, nonfusible_op - -from .charmnumeric import ( - Array, - ArrayOperation, - ArrayView, - _coerce_scalar, - _is_scalar_like, - arange, - array, - asarray, - copy, - copyto, - create_array, - empty, - empty_like, - eye, - full, - full_like, - identity, - ones, - ones_like, - zeros, - zeros_like, -) - - -def _coerce_binary_operand(value): - if isinstance(value, (Array, ArrayView)): - return value - if _is_scalar_like(value): - return _coerce_scalar(value) - return asarray(value) - - -def tanh(x): - if _is_scalar_like(x): - return np.tanh(_coerce_scalar(x)) - x = asarray(x) - return fusible_op(x, operation=ArrayOperation.tanh, shape=x.shape, dtype=x.dtype) - - -def exp(x): - if _is_scalar_like(x): - return np.exp(_coerce_scalar(x)) - x = asarray(x) - return fusible_op(x, operation=ArrayOperation.exp, shape=x.shape, dtype=x.dtype) - - -def tile(x, reps): - if _is_scalar_like(x): - return array(np.tile(_coerce_scalar(x), reps)) - - x = asarray(x) - if isinstance(reps, int): - reps = (reps,) - reps = tuple(int(rep) for rep in reps) - - out_ndims = max(x.ndims, len(reps)) - padded_shape = (1,) * (out_ndims - x.ndims) + x.shape - padded_reps = (1,) * (out_ndims - len(reps)) + reps - - out_shape = tuple(s * r for s, r in zip(padded_shape, padded_reps)) - return nonfusible_op( - x, - *padded_reps, - operation=ArrayOperation.tile, - shape=out_shape, - dtype=x.dtype, - ) - - -def add(x1, x2): - if _is_scalar_like(x1) and _is_scalar_like(x2): - return _coerce_scalar(x1) + _coerce_scalar(x2) - return _coerce_binary_operand(x1) + _coerce_binary_operand(x2) - - -def subtract(x1, x2): - if _is_scalar_like(x1) and _is_scalar_like(x2): - return _coerce_scalar(x1) - _coerce_scalar(x2) - return _coerce_binary_operand(x1) - _coerce_binary_operand(x2) - - -def multiply(x1, x2): - if _is_scalar_like(x1) and _is_scalar_like(x2): - return _coerce_scalar(x1) * _coerce_scalar(x2) - return _coerce_binary_operand(x1) * _coerce_binary_operand(x2) - - -def divide(x1, x2): - if _is_scalar_like(x1) and _is_scalar_like(x2): - return _coerce_scalar(x1) / _coerce_scalar(x2) - return _coerce_binary_operand(x1) / _coerce_binary_operand(x2) - - -def matmul(x1, x2): - if _is_scalar_like(x1) or _is_scalar_like(x2): - raise ValueError("matmul does not support scalar operands") - return asarray(x1) @ asarray(x2) - - -def dot(x1, x2): - if _is_scalar_like(x1) and _is_scalar_like(x2): - return _coerce_scalar(x1) * _coerce_scalar(x2) - - lhs = _coerce_binary_operand(x1) - rhs = _coerce_binary_operand(x2) - if _is_scalar_like(lhs) or _is_scalar_like(rhs): - return lhs * rhs - return lhs.dot(rhs) - - -def norm2(x): - """Compute the squared L2 norm of a 1D array.""" - - x = asarray(x) - if x.ndims != 1: - raise ValueError(f"norm2 requires a 1D array, got {x.ndims}D") - return x.dot(x) - - -def diag(v, k=0): - """Construct a diagonal matrix or extract a diagonal.""" - - v = asarray(v) - - if v.ndims == 1: - n = v.shape[0] + abs(k) - return nonfusible_op(v, k, operation=ArrayOperation.diag, shape=(n, n), dtype=v.dtype) - if v.ndims == 2: - rows, cols = v.shape - if k >= 0: - diag_len = max(0, min(rows, cols - k)) - else: - diag_len = max(0, min(rows + k, cols)) - if diag_len == 0: - raise ValueError(f"k={k} is out of range for shape ({rows}, {cols})") - return nonfusible_op( - v, - k, - operation=ArrayOperation.diag, - shape=(diag_len,), - dtype=v.dtype, - ) - raise ValueError(f"diag requires a 1D or 2D array, got {v.ndims}D") - - -__all__ = [ - "add", - "arange", - "array", - "asarray", - "copy", - "copyto", - "create_array", - "diag", - "divide", - "dot", - "empty", - "empty_like", - "exp", - "eye", - "full", - "full_like", - "identity", - "matmul", - "multiply", - "norm2", - "ones", - "ones_like", - "subtract", - "tanh", - "tile", - "zeros", - "zeros_like", -] diff --git a/charmnumeric/random.py b/charmnumeric/random.py deleted file mode 100644 index 148b29e..0000000 --- a/charmnumeric/random.py +++ /dev/null @@ -1,5 +0,0 @@ -from .charmnumeric import create_array -import numpy as np - -def randn(*args): - return create_array(args, dtype=np.float32) diff --git a/charmtiles/__init__.py b/charmtiles/__init__.py new file mode 100644 index 0000000..a1c1976 --- /dev/null +++ b/charmtiles/__init__.py @@ -0,0 +1 @@ +__version__ = '0.1.dev' diff --git a/charmtiles/array.py b/charmtiles/array.py new file mode 100644 index 0000000..c961004 --- /dev/null +++ b/charmtiles/array.py @@ -0,0 +1,200 @@ +import sys +import warnings +import numpy as np +from charmtiles.ast import get_max_depth, ASTNode +from charmtiles.ccs import to_bytes, from_bytes, send_command_raw, send_command, \ + send_command_async, connect, get_creation_command, \ + get_name, get_fetch_command, Handlers, OPCODES, is_debug + + +def create_ndarray(ndim, dtype, shape=None, name=None, command_buffer=None): + return ndarray(ndim, dtype=dtype, shape=shape, name=name, + command_buffer=command_buffer) + + +def from_numpy(nparr): + return ndarray(nparr.ndim, dtype=nparr.dtype, shape=nparr.shape, + nparr=nparr) + + +class ndarray: + def __init__(self, ndim, shape=None, dtype=np.float64, init_value=None, + nparr=None, name=None, command_buffer=None): + """ + This is the wrapper class for AUM array objects. + The argument 'name' should be None except when wrapping + an array that already exists on the AUM backend server + """ + if ndim > 2: + raise NotImplementedError("Arrays of dimensionality greater than" + "2 not supported yet") + self.dtype = dtype + self.ndim = ndim + self.itemsize = np.dtype(dtype).itemsize + self.init_value = init_value + self.command_buffer = command_buffer + if isinstance(shape, np.ndarray) or isinstance(shape, list) or \ + isinstance(shape, tuple): + self._shape = np.asarray(shape, dtype=np.int32) + elif shape is not None: + self._shape = np.asarray([shape], dtype=np.int32) + else: + self._shape = np.zeros(self.ndim, dtype=np.int32) + self.valid = False + if command_buffer is None: + self.valid = True + if name: + self.name = name + self.command_buffer = ASTNode(self.name, 0, [self]) + else: + self.name = get_name() + if nparr is not None: + buf = nparr.tobytes() + else: + buf = None + cmd = get_creation_command(self, self.name, self._shape, buf=buf) + if not send_command(Handlers.creation_handler, cmd, + reply_type='?'): + warnings.warn("Error creating array on server", RuntimeWarning) + self.command_buffer = ASTNode(self.name, 0, [self]) + else: + self.name = name + max_depth = get_max_depth() + if self.command_buffer.depth > max_depth: + if is_debug(): + print("Maximum AST depth exceeded for %i, " + "flushing buffer" % self.name) + self._flush_command_buffer() + + @property + def shape(self): + self._flush_command_buffer() + return self._shape + + def __del__(self): + if self.valid: + cmd = to_bytes(self.name, 'L') + send_command_async(Handlers.delete_handler, cmd) + + def __len__(self): + return self.shape[0] + + def __str__(self): + print(self.get()) + + #def __repr__(self): + # #self._flush_command_buffer() + # # FIXME add repr + # pass + + def __neg__(self): + return self * -1 + + def __add__(self, other): + res = get_name() + cmd_buffer = ASTNode(res, OPCODES.get('+'), [self, other]) + return create_ndarray(self.ndim, self.dtype, + name=res, command_buffer=cmd_buffer) + + def __radd__(self, other): + return self + other + + def __sub__(self, other): + res = get_name() + cmd_buffer = ASTNode(res, OPCODES.get('-'), [self, other]) + return create_ndarray(self.ndim, self.dtype, + name=res, command_buffer=cmd_buffer) + + def __rsub__(self, other): + return -1 * (self - other) + + def __mul__(self, other): + if self.ndim > 0 or (isinstance(other, ndarray) and other.ndim > 0): + RuntimeError("Cannote multiply two arrays") + res = get_name() + cmd_buffer = ASTNode(res, OPCODES.get('*'), [self, other]) + return create_ndarray(self.ndim, self.dtype, + name=res, command_buffer=cmd_buffer) + + def __rmul__(self, other): + return self * other + + def __truediv__(self, other): + if self.ndim > 0 or other.ndim > 0: + RuntimeError("Cannote divide two arrays") + res = get_name() + cmd_buffer = ASTNode(res, OPCODES.get('/'), [self, other]) + return create_ndarray(self.ndim, self.dtype, + name=res, command_buffer=cmd_buffer) + + def __matmul__(self, other): + if self.ndim == 2 and other.ndim == 2: + res_ndim = 2 + elif self.ndim == 2 and other.ndim == 1: + res_ndim = 1 + elif self.ndim == 1 and other.ndim == 1: + res_ndim = 0 + else: + raise RuntimeError("Dimension mismatch") + res = get_name() + cmd_buffer = ASTNode(res, OPCODES.get('@'), [self, other]) + return create_ndarray(res_ndim, self.dtype, + name=res, command_buffer=cmd_buffer) + + def _flush_command_buffer(self): + # send the command to server + # finally set command buffer to array name + debug = is_debug() + if debug: + self.command_buffer.plot_graph() + if self.valid: + return + validated_arrays = {self.name : self} + cmd = self.command_buffer.get_command(validated_arrays) + reply_size = 0 + for name, arr in validated_arrays.items(): + reply_size += 8 + 8 * arr.ndim + if not debug: + metadata = send_command_raw(Handlers.operation_handler, + cmd, reply_size=reply_size) + # traverse metadata + offset = 0 + for i in range(len(validated_arrays)): + name = from_bytes(metadata[offset : offset + 8], 'L') + offset += 8 + arr = validated_arrays[name] + for d in range(arr.ndim): + arr._shape[d] = from_bytes(metadata[offset : offset + 8], 'L') + offset += 8 + arr.validate() + else: + for name, arr in validated_arrays.items(): + arr.validate() + self.validate() + + def get(self): + self._flush_command_buffer() + cmd = get_fetch_command(self) + if self.ndim == 0: + total_size = self.itemsize + data_bytes = send_command_raw(Handlers.fetch_handler, cmd, reply_size=total_size) + return from_bytes(data_bytes, np.dtype(self.dtype).char) + else: + total_size = self.size * self.itemsize + data_ptr = send_command(Handlers.fetch_handler, cmd, reply_size=total_size) + data = cast(memoryview, data_ptr) + return np.frombuffer(data, np.dtype(self.dtype)).copy() + + def evaluate(self): + self._flush_command_buffer() + + def validate(self): + self.valid = True + self.command_buffer = ASTNode(self.name, 0, [self]) + + def copy(self): + res = get_name() + cmd_buffer = ASTNode(res, OPCODES.get('copy'), [self]) + return create_ndarray(self.ndim, self.dtype, + name=res, command_buffer=cmd_buffer) + diff --git a/charmtiles/ast.py b/charmtiles/ast.py new file mode 100644 index 0000000..622c1a5 --- /dev/null +++ b/charmtiles/ast.py @@ -0,0 +1,120 @@ +import numpy as np +import networkx as nx +import matplotlib.pyplot as plt +from ctypes import c_long +from networkx.drawing.nx_pydot import graphviz_layout +from charmtiles.ccs import OPCODES, INV_OPCODES, to_bytes + + +max_depth = 10 + + +def set_max_depth(d): + global max_depth + max_depth = d + + +def get_max_depth(): + global max_depth + return max_depth + + +class ASTNode(object): + def __init__(self, name, opcode, operands): + from charmtiles.array import ndarray + # contains opcode, operands + # operands are ndarrays + self.name = name + self.opcode = opcode + self.operands = operands + self.depth = 0 + if self.opcode != 0: + for op in self.operands: + if isinstance(op, ndarray): + self.depth = max(self.depth, 1 + op.command_buffer.depth) + + def get_command(self, validated_arrays, save=True): + from charmtiles.array import ndarray + if self.opcode == 0: + cmd = to_bytes(self.opcode, 'L') + cmd += to_bytes(False, '?') + cmd += to_bytes(self.operands[0].name, 'L') + return cmd + cmd = to_bytes(self.opcode, 'L') + to_bytes(self.name, 'L') + cmd += to_bytes(save, '?') + to_bytes(len(self.operands), 'B') + for op in self.operands: + # an operand can also be a double + if isinstance(op, ndarray): + if op.name in validated_arrays: + opcmd = to_bytes(0, 'L') + opcmd += to_bytes(False, '?') + opcmd += to_bytes(op.name, 'L') + cmd += to_bytes(len(opcmd), 'I') + cmd += opcmd + else: + save_op = True if c_long.from_address(id(op)).value - 2 > 0 else False + opcmd = op.command_buffer.get_command(validated_arrays, + save=save_op) + if not op.valid and save_op: + validated_arrays[op.name] = op + cmd += to_bytes(len(opcmd), 'I') + cmd += opcmd + elif isinstance(op, float) or isinstance(op, int): + opcmd = to_bytes(0, 'L') + opcmd += to_bytes(True, '?') + opcmd += to_bytes(op, 'd') + #print(opcmd, len(opcmd), len(cmd)) + cmd += to_bytes(len(opcmd), 'I') + cmd += opcmd + return cmd + + def plot_graph(self, validated_arrays={}, G=None, node_map={}, + color_map={}, next_id=0, parent=None, save=True): + from charmtiles.array import ndarray + if G is None: + G = nx.Graph() + if self.opcode == 0: + node_map[next_id] = 'a' + str(self.operands[0].name) + G.add_node(next_id) + if parent is not None: + G.add_edge(parent, next_id) + return next_id + 1 + opnode = next_id + G.add_node(next_id) + if parent is not None: + G.add_edge(parent, next_id) + node_map[next_id] = INV_OPCODES.get(self.opcode, '?') + if save: + color_map[next_id] = 'tab:red' + node_map[next_id] += (': a%i' % self.name) + next_id += 1 + for op in self.operands: + # an operand can also be a double + if isinstance(op, ndarray): + if op.name in validated_arrays: + G.add_node(next_id) + G.add_edge(opnode, next_id) + node_map[next_id] = 'a' + str(op.name) + color_map[next_id] = 'tab:green' + next_id += 1 + else: + save_op = True if c_long.from_address(id(op)).value - 2 > 0 else False + if not op.valid and save_op: + #color_map[next_id] = 'tab:red' + validated_arrays[op.name] = op + next_id = op.command_buffer.plot_graph( + validated_arrays, G, node_map, color_map, next_id, + opnode, save_op) + elif isinstance(op, float) or isinstance(op, int): + G.add_node(next_id) + G.add_edge(opnode, next_id) + node_map[next_id] = op + next_id += 1 + if parent is None: + pos = graphviz_layout(G, prog='dot') + color_map_list = [color_map.get(node, 'tab:blue') for node in G] + nx.draw(G, pos, labels=node_map, node_color=color_map_list, + node_size=600, font_size=10) + plt.show() + return next_id + diff --git a/charmtiles/ccs.py b/charmtiles/ccs.py new file mode 100644 index 0000000..3063d79 --- /dev/null +++ b/charmtiles/ccs.py @@ -0,0 +1,99 @@ +import struct +import atexit +from pyccs import Server +from charmtiles import array + +debug = False +server = None +client_id = 0 +next_name = 0 + +OPCODES = {'+': 1, '-': 2, '*': 3 ,'/': 4, '@': 5, 'copy': 6, 'axpy': 7, + 'axpy_multiplier': 8} + +INV_OPCODES = {v: k for k, v in OPCODES.items()} + +def enable_debug(): + global debug + debug = True + +def disable_debug(): + global debug + debug = False + +def is_debug(): + global debug + return debug + +def get_name(): + global next_name + curr_name = next_name + next_name += 1 + return (client_id << 56) + curr_name + +def to_bytes(value, dtype='I'): + return struct.pack(dtype, value) + +def from_bytes(bvalue, dtype='I'): + return struct.unpack(dtype, bvalue)[0] + +def send_command_raw(handler, msg, reply_size): + if server is not None: + server.send_request(handler, 0, msg) + return server.receive_response(reply_size) + +def send_command(handler, msg, reply_size=1, reply_type='B'): + global server + if server is not None: + return from_bytes(send_command_raw(handler, msg, reply_size), reply_type) + +def send_command_async(handler, msg): + global server + if server is not None: + server.send_request(handler, 0, msg) + +def connect(server_ip, server_port): + global server, client_id, debug + if not debug: + server = Server(server_ip, server_port) + server.connect() + client_id = send_command(Handlers.connection_handler, "") + atexit.register(disconnect) + +def disconnect(): + global client_id + cmd = to_bytes(client_id, 'B') + send_command_async(Handlers.disconnection_handler, cmd) + +def get_creation_command(arr, name, shape, buf=None): + """ + Generate array creation CCS command + """ + cmd = to_bytes(name, 'L') + cmd += to_bytes(arr.ndim, 'I') + cmd += to_bytes(buf is not None, '?') + cmd += to_bytes(arr.init_value is not None, '?') + for s in shape: + cmd += to_bytes(int(s), 'L') + if buf is not None: + cmd += buf + elif arr.init_value is not None: + cmd += to_bytes(arr.init_value, 'd') + return cmd + +def get_fetch_command(arr): + """ + Generate CCS command to fetch entire array data + """ + cmd = to_bytes(arr.name, 'L') + return cmd + +class Handlers(object): + connection_handler = b'aum_connect' + disconnection_handler = b'aum_disconnect' + creation_handler = b'aum_creation' + operation_handler = b'aum_operation' + fetch_handler = b'aum_fetch' + delete_handler = b'aum_delete' + exit_handler = b'aum_exit' + diff --git a/charmtiles/linalg.py b/charmtiles/linalg.py new file mode 100644 index 0000000..5449a69 --- /dev/null +++ b/charmtiles/linalg.py @@ -0,0 +1,22 @@ +import sys +import struct +import numpy as np +from pyccs import Server +from charmtiles.ccs import OPCODES, get_name, send_command, Handlers +from charmtiles.array import create_ndarray +from charmtiles.ast import ASTNode + + +def axpy(a, x, y, multiplier=None): + operands = [a, x, y] + if multiplier is not None: + operands.append(multiplier) + operation = 'axpy_multiplier' + else: + operation = 'axpy' + res = get_name() + cmd_buffer = ASTNode(res, OPCODES.get(operation), operands) + return create_ndarray(x.ndim, x.dtype, + name=res, command_buffer=cmd_buffer) + + diff --git a/docs/images/conj_ast.png b/docs/images/conj_ast.png new file mode 100644 index 0000000..e20d740 Binary files /dev/null and b/docs/images/conj_ast.png differ diff --git a/docs/images/simple_ast.png b/docs/images/simple_ast.png new file mode 100644 index 0000000..1e1f52c Binary files /dev/null and b/docs/images/simple_ast.png differ diff --git a/examples/cg.py b/examples/cg.py deleted file mode 100644 index a36ba31..0000000 --- a/examples/cg.py +++ /dev/null @@ -1,73 +0,0 @@ -"""Conjugate Gradient solver using charmnumeric. - -Solves the linear system A x = b where A is symmetric positive-definite. -We build a simple SPD matrix A = I + ones (the identity plus a constant -matrix) so the answer is easy to verify with NumPy. -""" - -import charmnumeric as cnp -import numpy as np - -N = 128 -max_iter = 2 - -# --- Build a symmetric positive-definite matrix A = (N+1)*I + ones -------- -# This is SPD because eigenvalues are (N+1) (multiplicity N-1) and (2N+1). -A = cnp.ones((N, N), dtype=cnp.float64) -for i in range(N): - A[i, i] = N + 1.0 # add N to diagonal -> diag = N+1 - -# --- Right-hand side b = [1, 2, ..., N] ----------------------------------- -# (Not an eigenvector of A, so CG needs multiple iterations.) -b = cnp.arange(1, N + 1, dtype=cnp.float64) - -# --- CG iteration --------------------------------------------------------- -# x0 = 0, r0 = b - A*x0 = b, p0 = r0 -x = cnp.zeros((N,), dtype=cnp.float64) - -r = b.copy() - -p = b.copy() - -rtr = r @ r # r^T r (dot product via matvec on column vec) - -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.115', 1234, 4) - -tol = 1e-12 -check_every = 1 - -for k in range(max_iter): - Ap = A @ p # matrix-vector product - pAp = p @ Ap # p^T A p - alpha = rtr / pAp # step length (scalar) - - x = x + alpha * p # update solution - r = r - alpha * Ap # update residual - - rtr_new = r @ r # new r^T r - beta = rtr_new / rtr # improvement ratio - p = r + beta * p # update search direction - - rtr = rtr_new - - if (k + 1) % check_every == 0 or k == 0: - rtr_val = rtr.get(interface) - print(f" iter {k+1}: rtr = {rtr_val.item():.6e}") - if rtr_val.item() < tol: - print(f" Converged at iteration {k+1}") - break - -result = x.get(interface) -print("CG solution x:") -print(result.flatten()) - -# --- Verify against NumPy -------------------------------------------------- -A_np = np.ones((N, N), dtype=np.float64) -np.fill_diagonal(A_np, N + 1.0) -b_np = np.arange(1, N + 1, dtype=np.float64) -x_expected = np.linalg.solve(A_np, b_np) - -print("\nExpected (numpy):") -print(x_expected) -print("\nMatch:", np.allclose(result.flatten(), x_expected)) diff --git a/examples/conjugate_gradient.py b/examples/conjugate_gradient.py new file mode 100644 index 0000000..6eac734 --- /dev/null +++ b/examples/conjugate_gradient.py @@ -0,0 +1,44 @@ +from charmtiles.array import connect, ndarray +import charmtiles.linalg as lg +from charmtiles.ccs import enable_debug +import numpy as np + +import time + +enable_debug() + +def solve(A, b): + x = ndarray(1, 1000, np.float64) + r = b - A @ x + p = r.copy() + rsold = r @ r + + for i in range(1000): + Ap = A @ p + alpha = rsold / (p @ Ap) + + x = lg.axpy(alpha, p, x) + r = lg.axpy(alpha, Ap, r, multiplier=-1.) + + rsnew = r @ r + + if np.sqrt(rsnew.get()) < 1e-8: + print("Converged in %i iterations" % (i + 1)) + break + + p = lg.axpy(rsnew / rsold, p, r) + rsold = rsnew + + return x + +if __name__ == '__main__': + connect("172.17.0.1", 10000) + + A = ndarray(2, (1000, 1000), np.float64, init_value=1.) + b = ndarray(1, 1000, np.float64, init_value=1.) + + start = time.time() + x = solve(A, b) + print("Execution time = %.6f" % (time.time() - start)) + + diff --git a/examples/diag.py b/examples/diag.py deleted file mode 100644 index ffc5102..0000000 --- a/examples/diag.py +++ /dev/null @@ -1,132 +0,0 @@ -"""Test diag function with numpy.diag semantics. - -Tests both modes: - 1D → 2D: construct diagonal matrix from vector - 2D → 1D: extract diagonal from matrix - -Also tests the k offset parameter for super/sub-diagonals. -""" - -import charmnumeric as cnp -from charmtyles.core import execute -import numpy as np - -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.114', 1234, 4) - -N = 128 - -# ------------------------------------------------------------------ -# Test 1: 1D → 2D — diag(v) constructs diagonal matrix (k=0) -# ------------------------------------------------------------------ -v = cnp.ones((N,), dtype=cnp.float64) -execute(interface) - -D = cnp.diag(v) -execute(interface) -result1 = D.get(interface) -v_np = np.ones(N, dtype=np.float64) -expected1 = np.diag(v_np) -print("Test 1: diag(v) — 1D->2D, k=0") -print(f" Match: {np.allclose(result1, expected1)}") -if not np.allclose(result1, expected1): - print(f" Max error: {np.max(np.abs(result1 - expected1))}") - -# ------------------------------------------------------------------ -# Test 2: 2D → 1D — diag(A) extracts main diagonal (k=0) -# ------------------------------------------------------------------ -A = cnp.ones((N, N), dtype=cnp.float64) -execute(interface) - -d = cnp.diag(A) -execute(interface) -result2 = d.get(interface) -A_np = np.ones((N, N), dtype=np.float64) -expected2 = np.diag(A_np) -print("\nTest 2: diag(A) — 2D->1D, k=0") -print(f" Match: {np.allclose(result2, expected2)}") -if not np.allclose(result2, expected2): - print(f" Max error: {np.max(np.abs(result2 - expected2))}") - -# ------------------------------------------------------------------ -# Test 3: 1D → 2D with k=1 (super-diagonal) -# ------------------------------------------------------------------ -v2 = cnp.full((N,), 2.0, dtype=cnp.float64) -execute(interface) - -D3 = cnp.diag(v2, k=1) -execute(interface) -result3 = D3.get(interface) -v2_np = np.full(N, 2.0, dtype=np.float64) -expected3 = np.diag(v2_np, k=1) -print("\nTest 3: diag(v, k=1) — 1D->2D, super-diagonal") -print(f" Match: {np.allclose(result3, expected3)}") -if not np.allclose(result3, expected3): - print(f" Max error: {np.max(np.abs(result3 - expected3))}") - -# ------------------------------------------------------------------ -# Test 4: 1D → 2D with k=-1 (sub-diagonal) -# ------------------------------------------------------------------ -D4 = cnp.diag(v2, k=-1) -execute(interface) -result4 = D4.get(interface) -expected4 = np.diag(v2_np, k=-1) -print("\nTest 4: diag(v, k=-1) — 1D->2D, sub-diagonal") -print(f" Match: {np.allclose(result4, expected4)}") -if not np.allclose(result4, expected4): - print(f" Max error: {np.max(np.abs(result4 - expected4))}") - -# ------------------------------------------------------------------ -# Test 5: 2D → 1D with k=1 (extract super-diagonal) -# ------------------------------------------------------------------ -d5 = cnp.diag(A, k=1) -execute(interface) -result5 = d5.get(interface) -expected5 = np.diag(A_np, k=1) -print("\nTest 5: diag(A, k=1) — 2D->1D, extract super-diagonal") -print(f" Match: {np.allclose(result5, expected5)}") -if not np.allclose(result5, expected5): - print(f" Max error: {np.max(np.abs(result5 - expected5))}") - -# ------------------------------------------------------------------ -# Test 6: 2D → 1D with k=-1 (extract sub-diagonal) -# ------------------------------------------------------------------ -d6 = cnp.diag(A, k=-1) -execute(interface) -result6 = d6.get(interface) -expected6 = np.diag(A_np, k=-1) -print("\nTest 6: diag(A, k=-1) — 2D->1D, extract sub-diagonal") -print(f" Match: {np.allclose(result6, expected6)}") -if not np.allclose(result6, expected6): - print(f" Max error: {np.max(np.abs(result6 - expected6))}") - -# ------------------------------------------------------------------ -# Test 7: Non-square matrix — extract diagonal -# ------------------------------------------------------------------ -M, K = 64, 128 -B = cnp.full((M, K), 3.0, dtype=cnp.float64) -execute(interface) - -d7 = cnp.diag(B) -execute(interface) -result7 = d7.get(interface) -B_np = np.full((M, K), 3.0, dtype=np.float64) -expected7 = np.diag(B_np) -print("\nTest 7: diag(B) — non-square (64x128), k=0") -print(f" Match: {np.allclose(result7, expected7)}") -if not np.allclose(result7, expected7): - print(f" Max error: {np.max(np.abs(result7 - expected7))}") - -# ------------------------------------------------------------------ -# Summary -# ------------------------------------------------------------------ -all_pass = all([ - np.allclose(result1, expected1), - np.allclose(result2, expected2), - np.allclose(result3, expected3), - np.allclose(result4, expected4), - np.allclose(result5, expected5), - np.allclose(result6, expected6), - np.allclose(result7, expected7), -]) -print(f"\n{'All tests passed!' if all_pass else 'SOME TESTS FAILED'}") diff --git a/examples/gauss_seidel.py b/examples/gauss_seidel.py deleted file mode 100644 index c307a03..0000000 --- a/examples/gauss_seidel.py +++ /dev/null @@ -1,59 +0,0 @@ -"""Red-Black Gauss-Seidel solver for 2D Poisson equation using charmnumeric. - -Solves -laplacian(phi) = f on a unit square with zero Dirichlet BCs. -Uses red-black ordering via strided slicing for parallel-safe updates. -""" - -import charmnumeric as cnp -import numpy as np - -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.115', 1234, 4) - -N = 512 -h = 1.0 / (N - 1) -h2 = h ** 2 -max_iter = 10 - -# --- Setup arrays --- -phi = cnp.zeros((N, N), dtype=cnp.float64) - -f = cnp.zeros((N, N), dtype=cnp.float64) - -# Point source in the middle -f[N // 2, N // 2] = -100.0 - -# --- Red-Black Gauss-Seidel iterations --- -for it in range(max_iter): - # RED points: (i+j) even - # Pattern 1: odd rows, odd cols - phi[1:-1:2, 1:-1:2] = 0.25 * ( - phi[0:-2:2, 1:-1:2] + phi[2::2, 1:-1:2] - + phi[1:-1:2, 0:-2:2] + phi[1:-1:2, 2::2] - - h2 * f[1:-1:2, 1:-1:2] - ) - # Pattern 2: even rows, even cols - phi[2:-1:2, 2:-1:2] = 0.25 * ( - phi[1:-2:2, 2:-1:2] + phi[3::2, 2:-1:2] - + phi[2:-1:2, 1:-2:2] + phi[2:-1:2, 3::2] - - h2 * f[2:-1:2, 2:-1:2] - ) - - # BLACK points: (i+j) odd - # Pattern 1: odd rows, even cols - phi[1:-1:2, 2:-1:2] = 0.25 * ( - phi[0:-2:2, 2:-1:2] + phi[2::2, 2:-1:2] - + phi[1:-1:2, 1:-2:2] + phi[1:-1:2, 3::2] - - h2 * f[1:-1:2, 2:-1:2] - ) - # Pattern 2: even rows, odd cols - phi[2:-1:2, 1:-1:2] = 0.25 * ( - phi[1:-2:2, 1:-1:2] + phi[3::2, 1:-1:2] - + phi[2:-1:2, 0:-2:2] + phi[2:-1:2, 2::2] - - h2 * f[2:-1:2, 1:-1:2] - ) - -result = phi.get(interface) -print(f"Gauss-Seidel ({max_iter} iterations, N={N})") -print(f" phi max = {np.max(result):.6f}") -print(f" phi min = {np.min(result):.6f}") diff --git a/examples/graph.py b/examples/graph.py new file mode 100644 index 0000000..bd061b5 --- /dev/null +++ b/examples/graph.py @@ -0,0 +1,25 @@ +from charmtiles.array import connect, ndarray +from charmtiles.ast import set_max_depth +from charmtiles.ccs import enable_debug +import charmtiles.linalg as lg +import numpy as np + +#enable_debug() +set_max_depth(100) + +def f(): + v = ndarray(1, 10, np.float64) + b = ndarray(1, 10, np.float64, init_value=10) + c = ndarray(1, 10, np.float64) + w = c + for i in range(5): + y = v + b + w + z = v - y + w = 2 * (c - z) + b + w.evaluate() + + +if __name__ == '__main__': + connect("172.17.0.1", 10000) + s = f() + diff --git a/examples/jacobi.py b/examples/jacobi.py deleted file mode 100644 index 8b803ab..0000000 --- a/examples/jacobi.py +++ /dev/null @@ -1,29 +0,0 @@ -import charmnumeric as cnp -from charmtyles.core import execute - -u = cnp.zeros((128,), dtype=cnp.float32) -#b = cnp.zeros((100, 100), dtype=cnp.float32) - - -# u[0, :] = 1.0 -# u[-1, :] = 1.0 -# u[:, 0] = 1.0 -# u[:, -1] = 1.0 - -# for it in range(3): -# u[1:-1, 1:-1] = 0.25 * (u[:-2, 1:-1] + u[2:, 1:-1] + u[1:-1, :-2] + u[1:-1, 2:]) - -u[0] = 1.0 -u[-1] = 1.0 - -for it in range(3): - t = 0.5 * (u[:-2] + u[2:]) - u[1:-1] = t - #u, u2 = u2, u - -#plot_execution_state() -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.114', 1234, 4) -execute(interface) - -print(u.get(interface)) diff --git a/examples/jacobi2d.py b/examples/jacobi2d.py deleted file mode 100644 index 0927974..0000000 --- a/examples/jacobi2d.py +++ /dev/null @@ -1,20 +0,0 @@ -import charmnumeric as cnp -from charmtyles.core import execute, set_auto_flush - -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.115', 1234, 4) -set_auto_flush(interface, 1000) - -u = cnp.zeros((2048, 2048), dtype=cnp.float32) - -u[0, :] = 1.0 -u[-1, :] = 1.0 -u[:, 0] = 1.0 -u[:, -1] = 1.0 - -for it in range(20): - u[1:-1, 1:-1] = 0.25 * (u[:-2, 1:-1] + u[2:, 1:-1] + u[1:-1, :-2] + u[1:-1, 2:]) - -#plot_execution_state() - -print(u.get(interface)) diff --git a/examples/jacobi3d.py b/examples/jacobi3d.py deleted file mode 100644 index 6375375..0000000 --- a/examples/jacobi3d.py +++ /dev/null @@ -1,20 +0,0 @@ -import charmnumeric as cnp -from charmtyles.core import execute - -u = cnp.zeros((128, 128, 128), dtype=cnp.float32) - -u[0, :, :] = 1.0 -u[-1, :, :] = 1.0 -u[:, 0, :] = 1.0 -u[:, -1, :] = 1.0 -u[:, :, 0] = 1.0 -u[:, :, -1] = 1.0 - -for it in range(100): - u[1:-1, 1:-1, 1:-1] = 0.16666666666666666 * (u[:-2, 1:-1, 1:-1] + u[2:, 1:-1, 1:-1] + u[1:-1, :-2, 1:-1] + u[1:-1, 2:, 1:-1] + u[1:-1, 1:-1, :-2] + u[1:-1, 1:-1, 2:]) - -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.115', 1234, 4) -execute(interface) - -print(u.get(interface)) diff --git a/examples/jacobi_iteration.py b/examples/jacobi_iteration.py deleted file mode 100644 index 7fd348f..0000000 --- a/examples/jacobi_iteration.py +++ /dev/null @@ -1,73 +0,0 @@ -"""Jacobi iterative solver for a linear system A x = b. - -Uses the splitting A = D + R where D = diag(A) and R = A - D. -Each iteration: x_{k+1} = D^{-1} (b - R x_k) - -We build a diagonally dominant SPD matrix so convergence is guaranteed, -then verify the result against NumPy's direct solve. -""" - -import charmnumeric as cnp -from charmtyles.core import execute -import numpy as np - -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.115', 1234, 4) - -N = 128 -max_iter = 50 - -# --- Build a diagonally dominant matrix A ------------------------------------ -# A = 2*N*I + ones (all off-diagonal entries are 1, diagonal entries are 2*N+1) -# This is SPD and diagonally dominant, so Jacobi converges. -A = cnp.ones((N, N), dtype=cnp.float64) -A_diag_vec = cnp.full((N,), 2.0 * N, dtype=cnp.float64) -D_mat = cnp.diag(A_diag_vec) # diagonal matrix -execute(interface) - -A = A + D_mat # A = ones + 2N*I (diag = 2N+1) -execute(interface) - -# --- Right-hand side b ------------------------------------------------------- -b = cnp.ones((N,), dtype=cnp.float64) -execute(interface) - -# --- Extract diagonal and compute D_inv (element-wise reciprocal) ------------ -d = cnp.diag(A) # extract diagonal → 1D vector -execute(interface) - -d_np = d.get(interface) -d_inv_val = 1.0 / (2.0 * N + 1.0) -D_inv = cnp.full((N,), d_inv_val, dtype=cnp.float64) -execute(interface) - -# --- Jacobi iteration -------------------------------------------------------- -# x_{k+1} = D^{-1} * (b - (A x_k - D x_k)) -# = D^{-1} * (b - A x_k + d * x_k) -# = D^{-1} * (b - A x_k) + (1 - D^{-1} * d) ... simplified below -# -# Direct form: x = D_inv * (b - R x) where R = A - diag(A) -# Equivalently: x = D_inv * (b - A x + d * x) - -x = cnp.zeros((N,), dtype=cnp.float64) # initial guess -execute(interface) - -for it in range(max_iter): - Ax = A @ x # matrix-vector product - r = b - Ax # residual - # x_new = x + D_inv * r (Jacobi update: x + D^{-1}(b - Ax)) - x = x + D_inv * r - -# --- Retrieve and verify ----------------------------------------------------- -result = x.get(interface) - -# NumPy reference solution -A_np = np.ones((N, N), dtype=np.float64) + 2.0 * N * np.eye(N, dtype=np.float64) -b_np = np.ones(N, dtype=np.float64) -x_np = np.linalg.solve(A_np, b_np) - -print(f"Jacobi iteration ({max_iter} iterations, N={N})") -print(f" x[0:5] = {result[:5]}") -print(f" x_np[0:5] = {x_np[:5]}") -print(f" Max error = {np.max(np.abs(result - x_np)):.2e}") -print(f" Converged = {np.allclose(result, x_np, atol=1e-6)}") diff --git a/examples/lstm_forward.py b/examples/lstm_forward.py deleted file mode 100644 index b3947f3..0000000 --- a/examples/lstm_forward.py +++ /dev/null @@ -1,51 +0,0 @@ -import charmnumeric as cnp -from charmtyles.core import plot_execution_state -import numpy as np - - -batch_size = 32 -hidden_size = 10 -sentence_length = 4 -word_size = 10 - -X = cnp.random.randn(sentence_length, batch_size, hidden_size) -h0 = cnp.random.randn(1, hidden_size) -WLSTM = cnp.random.randn( - word_size + hidden_size, 4 * hidden_size -) / np.sqrt(word_size + hidden_size) - -xphpb = WLSTM.shape[0] -d = hidden_size -n = sentence_length -b = batch_size - -Hin = cnp.zeros((n, b, xphpb), dtype=cnp.float32) -Hout = cnp.zeros((n, b, d), dtype=cnp.float32) -IFOG = cnp.zeros((n, b, d * 4), dtype=cnp.float32) -IFOGf = cnp.zeros((n, b, d * 4), dtype=cnp.float32) -C = cnp.zeros((n, b, d), dtype=cnp.float32) -Ct = cnp.zeros((n, b, d), dtype=cnp.float32) - -for t in range(n): - if t == 0: - prev = cnp.tile(h0, (b, 1)) - else: - prev = Hout[t - 1] - - Hin[t, :, :word_size] = X[t] - Hin[t, :, word_size:] = prev - # compute all gate activations. dots: - IFOG[t] = Hin[t].dot(WLSTM) - # non-linearities - IFOGf[t, :, : 3 * d] = 1.0 / ( - 1.0 + cnp.exp(-IFOG[t, :, : 3 * d]) - ) # sigmoids these are the gates - IFOGf[t, :, 3 * d :] = cnp.tanh(IFOG[t, :, 3 * d :]) # tanh - # compute the cell activation - C[t] = IFOGf[t, :, :d] * IFOGf[t, :, 3 * d :] - if t > 0: - C[t] += IFOGf[t, :, d : 2 * d] * C[t - 1] - Ct[t] = cnp.tanh(C[t]) - Hout[t] = IFOGf[t, :, 2 * d : 3 * d] * Ct[t] - -plot_execution_state() diff --git a/examples/matmatmul.py b/examples/matmatmul.py deleted file mode 100644 index 14918c7..0000000 --- a/examples/matmatmul.py +++ /dev/null @@ -1,166 +0,0 @@ -"""Test matrix-matrix multiplication (SUMMA) with full and sliced arrays. - -Creates matrices and performs matmatmul in various configurations, -verifying against NumPy reference results. -""" - -import charmnumeric as cnp -from charmtyles.core import execute -import numpy as np - -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.115', 1234, 4) - -M, K, N = 128, 128, 128 - -# Build matrices with known patterns -# A[i, j] = i + j + 1, B[i, j] = i - j + 1 -A = cnp.ones((M, K), dtype=cnp.float64) -B = cnp.ones((K, N), dtype=cnp.float64) -# for i in range(M): -# for j in range(K): -# A[i:i+1, j:j+1] = float(i + j + 1) -# for i in range(K): -# for j in range(N): -# B[i:i+1, j:j+1] = float(i - j + 1) - -# NumPy reference -A_np = np.ones((M, K), dtype=np.float64) -B_np = np.ones((K, N), dtype=np.float64) - -execute(interface) - -# ------------------------------------------------------------------ -# Test 1: Full matmul — A @ B -# ------------------------------------------------------------------ -C = A @ B -execute(interface) -result = C.get(interface) -expected = A_np @ B_np -print("Test 1: Full matmul A @ B") -print(f" Match: {np.allclose(result, expected)}") -if not np.allclose(result, expected): - print(f" Max error: {np.max(np.abs(result - expected))}") - -# ------------------------------------------------------------------ -# Test 2: Non-square — A(128×64) @ B(64×256) -# ------------------------------------------------------------------ -M2, K2, N2 = 128, 64, 256 -A2 = cnp.ones((M2, K2), dtype=cnp.float64) -B2 = cnp.ones((K2, N2), dtype=cnp.float64) -execute(interface) - -C2 = A2 @ B2 -execute(interface) -result2 = C2.get(interface) -A2_np = np.ones((M2, K2), dtype=np.float64) -B2_np = np.ones((K2, N2), dtype=np.float64) -expected2 = A2_np @ B2_np -print("\nTest 2: Non-square (128x64) @ (64x256)") -print(f" Match: {np.allclose(result2, expected2)}") -if not np.allclose(result2, expected2): - print(f" Max error: {np.max(np.abs(result2 - expected2))}") - -# ------------------------------------------------------------------ -# Test 3: Row-sliced A — A[0:32, :] @ B -# ------------------------------------------------------------------ -A_row = A[0:32, :] -C3 = A_row @ B -execute(interface) -result3 = C3.get(interface) -expected3 = A_np[0:32, :] @ B_np -print("\nTest 3: Row-sliced A[0:32, :] @ B") -print(f" Match: {np.allclose(result3, expected3)}") -if not np.allclose(result3, expected3): - print(f" Max error: {np.max(np.abs(result3 - expected3))}") - -# ------------------------------------------------------------------ -# Test 4: Col-sliced B — A @ B[:, 16:48] -# ------------------------------------------------------------------ -B_col = B[:, 16:48] -C4 = A @ B_col -execute(interface) -result4 = C4.get(interface) -expected4 = A_np @ B_np[:, 16:48] -print("\nTest 4: Col-sliced A @ B[:, 16:48]") -print(f" Match: {np.allclose(result4, expected4)}") -if not np.allclose(result4, expected4): - print(f" Max error: {np.max(np.abs(result4 - expected4))}") - -# ------------------------------------------------------------------ -# Test 5: Both sliced — A[8:40, 10:50] @ B[10:50, 5:45] -# ------------------------------------------------------------------ -A_both = A[8:40, 10:50] -B_both = B[10:50, 5:45] -C5 = A_both @ B_both -execute(interface) -result5 = C5.get(interface) -expected5 = A_np[8:40, 10:50] @ B_np[10:50, 5:45] -print("\nTest 5: Both sliced A[8:40, 10:50] @ B[10:50, 5:45]") -print(f" Match: {np.allclose(result5, expected5)}") -if not np.allclose(result5, expected5): - print(f" Max error: {np.max(np.abs(result5 - expected5))}") - -# ------------------------------------------------------------------ -# Test 6: Non-tile-aligned — A[3:37, 5:47] @ B[5:47, 7:39] -# ------------------------------------------------------------------ -A_mis = A[3:37, 5:47] -B_mis = B[5:47, 7:39] -C6 = A_mis @ B_mis -execute(interface) -result6 = C6.get(interface) -expected6 = A_np[3:37, 5:47] @ B_np[5:47, 7:39] -print("\nTest 6: Non-tile-aligned A[3:37, 5:47] @ B[5:47, 7:39]") -print(f" Match: {np.allclose(result6, expected6)}") -if not np.allclose(result6, expected6): - print(f" Max error: {np.max(np.abs(result6 - expected6))}") - -# ------------------------------------------------------------------ -# Test 7: Chain — (A @ B) @ x (matmatmul then matvec) -# ------------------------------------------------------------------ -x = cnp.ones((N,), dtype=cnp.float64) -execute(interface) - -C_chain = A @ B -y = C_chain.matvec(x) -execute(interface) -result7 = y.get(interface) -x_np = np.ones(N, dtype=np.float64) -expected7 = (A_np @ B_np) @ x_np -print("\nTest 7: Chain (A @ B) @ x") -print(f" Match: {np.allclose(result7.flatten(), expected7)}") -if not np.allclose(result7.flatten(), expected7): - print(f" Max error: {np.max(np.abs(result7.flatten() - expected7))}") - -# ------------------------------------------------------------------ -# Test 8: 3D dim-drop — A3d[0, :, :] @ B -# ------------------------------------------------------------------ -Batch = 2 -A3d = cnp.ones((Batch, M, K), dtype=cnp.float64) -execute(interface) - -A3d_slice = A3d[0, :, :] # shape (1, M, K) — singleton dim 0 -C8 = A3d_slice @ B -execute(interface) -result8 = C8.get(interface) -A3d_np = np.ones((Batch, M, K), dtype=np.float64) -expected8 = A3d_np[0, :, :] @ B_np -print("\nTest 8: 3D dim-drop A3d[0, :, :] @ B") -print(f" Match: {np.allclose(result8, expected8)}") -if not np.allclose(result8, expected8): - print(f" Max error: {np.max(np.abs(result8 - expected8))}") - -# ------------------------------------------------------------------ -# Summary -# ------------------------------------------------------------------ -all_pass = all([ - np.allclose(result, expected), - np.allclose(result2, expected2), - np.allclose(result3, expected3), - np.allclose(result4, expected4), - np.allclose(result5, expected5), - np.allclose(result6, expected6), - np.allclose(result7.flatten(), expected7), - np.allclose(result8, expected8), -]) -print(f"\n{'All tests passed!' if all_pass else 'SOME TESTS FAILED'}") diff --git a/examples/matvec.py b/examples/matvec.py deleted file mode 100644 index 3508e83..0000000 --- a/examples/matvec.py +++ /dev/null @@ -1,55 +0,0 @@ -import charmnumeric as cnp -from charmtyles.core import execute -import numpy as np -import sys - -# --- Matrix-vector multiply example --- -# Supports two modes: -# --cross : cross-partition matvec (1D vector on 1D partition) -# (default) : same-partition matvec (2D column vector on 2D partition) - -cross_partition = '--cross' in sys.argv - -M, N = 48, 48 - -A = cnp.zeros((M, N), dtype=cnp.float64) - -# Set A to a known pattern: A[i, j] = i + j -for i in range(M): - A[i, :] = float(i) -col_vals = cnp.zeros((M, N), dtype=cnp.float64) -for j in range(N): - col_vals[:, j] = float(j) -A = A + col_vals - -if cross_partition: - # Cross-partition: 1D vector on a separate 1D partition - x = cnp.ones((N,), dtype=cnp.float64) -else: - # Same-partition: 2D column vector on the 2D partition - x = cnp.ones((N, 1), dtype=cnp.float64) - -# Compute y = A @ x using matvec -y = A.matvec(x) - -# Connect to backend and execute -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.114', 1234, 4) -execute(interface) - -# Retrieve results -result = y.get(interface) -print("y = A @ x:") -print(result) - -# Verify against numpy -A_np = np.zeros((M, N), dtype=np.float64) -for i in range(M): - for j in range(N): - A_np[i, j] = i + j -x_np = np.ones((N,), dtype=np.float64) -y_expected = A_np @ x_np - -print("\nExpected (numpy):") -print(y_expected) -print("\nMatch:", np.allclose(result.flatten(), y_expected)) diff --git a/examples/matvec_slice.py b/examples/matvec_slice.py deleted file mode 100644 index e95e8ca..0000000 --- a/examples/matvec_slice.py +++ /dev/null @@ -1,124 +0,0 @@ -"""Test matvec with sliced arrays (no temporary alignment copies). - -Creates a large matrix and vector, then performs matvec on slices -to verify that the slice-aware cross-partition matmul produces -correct results directly from view metadata. -""" - -import charmnumeric as cnp -from charmtyles.core import execute -import numpy as np - -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.114', 1234, 4) - -M, N = 128, 128 - -# Build a matrix with a known pattern: A[i, j] = 1 -A = cnp.ones((M, N), dtype=cnp.float64) - -# Build a 1D vector: x[i] = i + 1 -x = cnp.arange(1, N + 1, dtype=cnp.float64) - -# NumPy reference -A_np = np.ones(M * N, dtype=np.float64).reshape(M, N) -x_np = np.arange(1, N + 1, dtype=np.float64) - -# ------------------------------------------------------------------ -# Test 1: Full matvec (baseline, no slicing) -# ------------------------------------------------------------------ -y_full = A.matvec(x) -execute(interface) -result_full = y_full.get(interface) -expected_full = A_np @ x_np -print("Test 1: Full matvec") -print(f" Match: {np.allclose(result_full.flatten(), expected_full)}") -if not np.allclose(result_full.flatten(), expected_full): - print(f" Max error: {np.max(np.abs(result_full.flatten() - expected_full))}") - -# ------------------------------------------------------------------ -# Test 2: Row slice — A[0:32, :] @ x -# ------------------------------------------------------------------ -A_row_slice = A[0:32, :] -y_row = A_row_slice.matvec(x) -execute(interface) -result_row = y_row.get(interface) -expected_row = A_np[0:32, :] @ x_np -print("\nTest 2: Row slice A[0:32, :] @ x") -print(f" Match: {np.allclose(result_row.flatten(), expected_row)}") -if not np.allclose(result_row.flatten(), expected_row): - print(f" Max error: {np.max(np.abs(result_row.flatten() - expected_row))}") - -# ------------------------------------------------------------------ -# Test 3: Column slice — A[:, 16:48] @ x[16:48] -# ------------------------------------------------------------------ -A_col_slice = A[:, 16:48] -x_col_slice = x[16:48] -y_col = A_col_slice.matvec(x_col_slice) -execute(interface) -result_col = y_col.get(interface) -expected_col = A_np[:, 16:48] @ x_np[16:48] -print("\nTest 3: Column slice A[:, 16:48] @ x[16:48]") -print(f" Match: {np.allclose(result_col.flatten(), expected_col)}") -if not np.allclose(result_col.flatten(), expected_col): - print(f" Max error: {np.max(np.abs(result_col.flatten() - expected_col))}") - -# ------------------------------------------------------------------ -# Test 4: Both row and column slice — A[8:40, 10:50] @ x[10:50] -# ------------------------------------------------------------------ -A_both_slice = A[8:40, 10:50] -x_both_slice = x[10:50] -y_both = A_both_slice.matvec(x_both_slice) -execute(interface) -result_both = y_both.get(interface) -expected_both = A_np[8:40, 10:50] @ x_np[10:50] -print("\nTest 4: Both slices A[8:40, 10:50] @ x[10:50]") -print(f" Match: {np.allclose(result_both.flatten(), expected_both)}") -if not np.allclose(result_both.flatten(), expected_both): - print(f" Max error: {np.max(np.abs(result_both.flatten() - expected_both))}") - -# ------------------------------------------------------------------ -# Test 5: Non-tile-aligned slice — A[3:37, 5:47] @ x[5:47] -# (tile size is 16, so 3 and 5 are misaligned) -# ------------------------------------------------------------------ -A_misaligned = A[3:37, 5:47] -x_misaligned = x[5:47] -y_mis = A_misaligned.matvec(x_misaligned) -execute(interface) -result_mis = y_mis.get(interface) -expected_mis = A_np[3:37, 5:47] @ x_np[5:47] -print("\nTest 5: Non-tile-aligned A[3:37, 5:47] @ x[5:47]") -print(f" Match: {np.allclose(result_mis.flatten(), expected_mis)}") -if not np.allclose(result_mis.flatten(), expected_mis): - print(f" Max error: {np.max(np.abs(result_mis.flatten() - expected_mis))}") - -# ------------------------------------------------------------------ -# Test 6: 3D array with dimension dropping — A3d[0, :, :] @ x -# ------------------------------------------------------------------ -B = 2 # batch size -A3d = cnp.ones((B, M, N), dtype=cnp.float64) - -A3d_slice = A3d[0, :, :] # shape (1, M, N) — singleton dim 0 -y_3d = A3d_slice.matvec(x) -execute(interface) -result_3d = y_3d.get(interface) - -A3d_np = np.ones((B, M, N), dtype=np.float64) -expected_3d = A3d_np[0, :, :] @ x_np -print("\nTest 6: 3D dim-drop A3d[0, :, :] @ x") -print(f" Match: {np.allclose(result_3d.flatten(), expected_3d)}") -if not np.allclose(result_3d.flatten(), expected_3d): - print(f" Max error: {np.max(np.abs(result_3d.flatten() - expected_3d))}") - -# ------------------------------------------------------------------ -# Summary -# ------------------------------------------------------------------ -all_pass = all([ - np.allclose(result_full.flatten(), expected_full), - np.allclose(result_row.flatten(), expected_row), - np.allclose(result_col.flatten(), expected_col), - np.allclose(result_both.flatten(), expected_both), - np.allclose(result_mis.flatten(), expected_mis), - np.allclose(result_3d.flatten(), expected_3d), -]) -print(f"\n{'All tests passed!' if all_pass else 'SOME TESTS FAILED'}") diff --git a/examples/multigrid2d.py b/examples/multigrid2d.py deleted file mode 100644 index 3749358..0000000 --- a/examples/multigrid2d.py +++ /dev/null @@ -1,677 +0,0 @@ -"""2D Multigrid V-cycle solver for the Poisson equation. - -Solves -Laplacian(u) = f on the unit square [0,1]^2 with Dirichlet -boundary conditions using a geometric multigrid V-cycle with: - - Weighted Jacobi smoothing - - Full-weighting restriction (fine -> coarse) - - Bilinear prolongation (coarse -> fine) - -Usage: - python multigrid2d.py -""" - -import argparse -import charmnumeric as cnp -from charmtyles.core import execute, set_auto_flush -import numpy as np - - -def jacobi_smooth(u, f, h2, n_smooth, omega=2.0 / 3.0): - """Apply n_smooth weighted-Jacobi iterations to smooth u.""" - for _ in range(n_smooth): - jacobi_update = 0.25 * ( - u[:-2, 1:-1] + u[2:, 1:-1] - + u[1:-1, :-2] + u[1:-1, 2:] - + h2 * f[1:-1, 1:-1] - ) - u[1:-1, 1:-1] = ( - (1.0 - omega) * u[1:-1, 1:-1] + omega * jacobi_update - ) - - -def compute_residual(u, f, h2): - """Return r = f - (-Lap(u)/h2), i.e. r = f + Lap(u)/h2. - - Discrete Laplacian: Lap(u)[i,j] = (u[i-1,j]+u[i+1,j]+u[i,j-1]+u[i,j+1]-4*u[i,j]) / h^2 - For -Lap(u) = f, the residual is r = f - (-Lap(u)/h2) = f - (4*u - neighbors) / h2. - """ - r = cnp.zeros(u.shape, dtype=cnp.float64) - r[1:-1, 1:-1] = f[1:-1, 1:-1] - (1.0 / h2) * ( - 4.0 * u[1:-1, 1:-1] - - u[:-2, 1:-1] - u[2:, 1:-1] - - u[1:-1, :-2] - u[1:-1, 2:] - ) - return r - - -def compute_residual_numpy(u, f, h2): - """NumPy reference for the residual operator.""" - r = np.zeros_like(u) - r[1:-1, 1:-1] = f[1:-1, 1:-1] - (1.0 / h2) * ( - 4.0 * u[1:-1, 1:-1] - - u[:-2, 1:-1] - u[2:, 1:-1] - - u[1:-1, :-2] - u[1:-1, 2:] - ) - return r - - -def restrict_full_weighting(r_fine, r_coarse): - """Restrict a fine-grid residual with the standard 9-point stencil.""" - r_coarse[:, :] = 0.0 - r_coarse[1:-1, 1:-1] = ( - 0.25 * r_fine[2:-2:2, 2:-2:2] - + 0.125 * ( - r_fine[1:-3:2, 2:-2:2] + r_fine[3:-1:2, 2:-2:2] - + r_fine[2:-2:2, 1:-3:2] + r_fine[2:-2:2, 3:-1:2] - ) - + 0.0625 * ( - r_fine[1:-3:2, 1:-3:2] + r_fine[1:-3:2, 3:-1:2] - + r_fine[3:-1:2, 1:-3:2] + r_fine[3:-1:2, 3:-1:2] - ) - ) - - -def restrict_full_weighting_numpy(r_fine): - """NumPy reference for the 9-point full-weighting restriction.""" - n_coarse = (r_fine.shape[0] - 1) // 2 + 1 - r_coarse = np.zeros((n_coarse, n_coarse), dtype=r_fine.dtype) - r_coarse[1:-1, 1:-1] = ( - 0.25 * r_fine[2:-2:2, 2:-2:2] - + 0.125 * ( - r_fine[1:-3:2, 2:-2:2] + r_fine[3:-1:2, 2:-2:2] - + r_fine[2:-2:2, 1:-3:2] + r_fine[2:-2:2, 3:-1:2] - ) - + 0.0625 * ( - r_fine[1:-3:2, 1:-3:2] + r_fine[1:-3:2, 3:-1:2] - + r_fine[3:-1:2, 1:-3:2] + r_fine[3:-1:2, 3:-1:2] - ) - ) - return r_coarse - - -def validation_pattern_ops(n): - """A deterministic pattern that crosses chare boundaries and parities.""" - quarter = n // 4 - three_quarter = 3 * n // 4 + 1 - eighth = n // 8 - three_eighth = 3 * n // 8 - seven_eighth = 7 * n // 8 + 1 - mid = n // 2 - - return [ - ((slice(None), slice(None)), 0.0), - ((slice(1, n - 1), slice(1, n - 1)), 1.0), - ((slice(1, n - 1, 2), slice(1, n - 1, 2)), 2.0), - ((slice(2, n - 2, 2), slice(1, n - 1, 2)), -1.0), - ((slice(1, n - 1, 2), slice(2, n - 2, 2)), 0.5), - ((slice(2, n - 2, 2), slice(2, n - 2, 2)), -0.25), - ((slice(quarter, three_quarter), slice(quarter, three_quarter)), 3.0), - ((slice(three_eighth, seven_eighth, 2), slice(eighth, 5 * n // 8 + 1, 2)), -2.0), - ((slice(eighth, seven_eighth, 4), slice(5 * n // 16, 15 * n // 16, 2)), 0.75), - ((slice(mid - 2, mid + 3), slice(mid - 2, mid + 3)), -4.0), - ] - - -def apply_validation_pattern_numpy(arr): - """Fill a NumPy array with the deterministic validation pattern.""" - for key, value in validation_pattern_ops(arr.shape[0]): - arr[key] = value - - -def apply_validation_pattern_backend(arr): - """Fill a backend array with the deterministic validation pattern.""" - for key, value in validation_pattern_ops(arr.shape[0]): - arr[key] = float(value) - - -def summarize_diff(label, backend, reference): - """Print max/L2 error stats for two arrays and return the max abs error.""" - diff = backend - reference - max_abs = float(np.max(np.abs(diff))) - l2_err = float(np.linalg.norm(diff)) - worst = np.unravel_index(np.argmax(np.abs(diff)), diff.shape) - print(label) - print(f" max abs error = {max_abs:.6e}") - print(f" l2 error = {l2_err:.6e}") - print( - f" worst entry = {worst}: backend={backend[worst]:.12e}, " - f"numpy={reference[worst]:.12e}, diff={diff[worst]:.12e}" - ) - return max_abs, l2_err - - -def prolongate_and_correct(e_coarse, u_fine): - """Prolongate a coarse correction with bilinear interpolation.""" - n_fine = u_fine.shape[0] - e_fine = cnp.zeros((n_fine, n_fine), dtype=cnp.float64) - e_fine[2:-2:2, 2:-2:2] = e_coarse[1:-1, 1:-1] - e_fine[1:-1:2, 2:-2:2] = 0.5 * ( - e_coarse[:-1, 1:-1] + e_coarse[1:, 1:-1] - ) - e_fine[2:-2:2, 1:-1:2] = 0.5 * ( - e_coarse[1:-1, :-1] + e_coarse[1:-1, 1:] - ) - e_fine[1:-1:2, 1:-1:2] = 0.25 * ( - e_coarse[:-1, :-1] + e_coarse[1:, :-1] - + e_coarse[:-1, 1:] + e_coarse[1:, 1:] - ) - u_fine[1:-1, 1:-1] = u_fine[1:-1, 1:-1] + e_fine[1:-1, 1:-1] - - -def prolongate_numpy(e_coarse): - """NumPy reference for bilinear prolongation from coarse to fine grid.""" - n_fine = 2 * (e_coarse.shape[0] - 1) + 1 - e_fine = np.zeros((n_fine, n_fine), dtype=e_coarse.dtype) - e_fine[2:-2:2, 2:-2:2] = e_coarse[1:-1, 1:-1] - e_fine[1:-1:2, 2:-2:2] = 0.5 * ( - e_coarse[:-1, 1:-1] + e_coarse[1:, 1:-1] - ) - e_fine[2:-2:2, 1:-1:2] = 0.5 * ( - e_coarse[1:-1, :-1] + e_coarse[1:-1, 1:] - ) - e_fine[1:-1:2, 1:-1:2] = 0.25 * ( - e_coarse[:-1, :-1] + e_coarse[1:, :-1] - + e_coarse[:-1, 1:] + e_coarse[1:, 1:] - ) - return e_fine - - -def jacobi_smooth_numpy(u, f, h2, n_smooth, omega=2.0 / 3.0): - """NumPy reference for weighted Jacobi smoothing.""" - out = np.array(u, copy=True) - for _ in range(n_smooth): - jacobi_update = 0.25 * ( - out[:-2, 1:-1] + out[2:, 1:-1] - + out[1:-1, :-2] + out[1:-1, 2:] - + h2 * f[1:-1, 1:-1] - ) - out[1:-1, 1:-1] = ( - (1.0 - omega) * out[1:-1, 1:-1] + omega * jacobi_update - ) - return out - - -def snapshot_backend(arr): - """Materialize a full-array backend snapshot without flushing early.""" - snap = cnp.zeros(arr.shape, dtype=cnp.float64) - snap[:, :] = arr[:, :] - return snap - - -def vcycle_numpy(u, f, h, n_pre, n_post, coarse_sweeps=400): - """NumPy reference for one multigrid V-cycle.""" - out = np.array(u, copy=True) - n = out.shape[0] - h2 = h * h - omega = 2.0 / 3.0 - - if n <= 5: - return jacobi_smooth_numpy(out, f, h2, coarse_sweeps, omega) - - out = jacobi_smooth_numpy(out, f, h2, n_pre, omega) - r = compute_residual_numpy(out, f, h2) - r_coarse = restrict_full_weighting_numpy(r) - e_coarse = np.zeros_like(r_coarse) - e_coarse = vcycle_numpy(e_coarse, r_coarse, 2.0 * h, n_pre, n_post, coarse_sweeps) - out = out + prolongate_numpy(e_coarse) - out = jacobi_smooth_numpy(out, f, h2, n_post, omega) - return out - - -def vcycles_numpy(u, f, h, cycles, n_pre, n_post, coarse_sweeps=400): - """NumPy reference for multiple V-cycles applied back-to-back.""" - out = np.array(u, copy=True) - for _ in range(cycles): - out = vcycle_numpy(out, f, h, n_pre, n_post, coarse_sweeps) - return out - - -def vcycle_numpy_capture(u, f, h, n_pre, n_post, coarse_sweeps=400): - """NumPy reference for one V-cycle plus top-level intermediate snapshots.""" - out = np.array(u, copy=True) - n = out.shape[0] - h2 = h * h - omega = 2.0 / 3.0 - - if n <= 5: - coarse = jacobi_smooth_numpy(out, f, h2, coarse_sweeps, omega) - return { - "u_pre": coarse, - "r": np.zeros_like(coarse), - "r_coarse": np.zeros((0, 0), dtype=coarse.dtype), - "e_coarse": coarse, - "u_corr": coarse, - "u_final": coarse, - } - - out = jacobi_smooth_numpy(out, f, h2, n_pre, omega) - u_pre = np.array(out, copy=True) - r = compute_residual_numpy(out, f, h2) - r_coarse = restrict_full_weighting_numpy(r) - e_coarse = np.zeros_like(r_coarse) - e_coarse = vcycle_numpy(e_coarse, r_coarse, 2.0 * h, n_pre, n_post, coarse_sweeps) - out = out + prolongate_numpy(e_coarse) - u_corr = np.array(out, copy=True) - out = jacobi_smooth_numpy(out, f, h2, n_post, omega) - return { - "u_pre": u_pre, - "r": r, - "r_coarse": r_coarse, - "e_coarse": e_coarse, - "u_corr": u_corr, - "u_final": out, - } - - -def vcycle(u, f, h, n_pre, n_post, coarse_sweeps=400): - """Perform one V-cycle for -Lap(u) = f on a grid with spacing h.""" - n = u.shape[0] - h2 = h * h - omega = 2.0 / 3.0 - - if n <= 5: - jacobi_smooth(u, f, h2, coarse_sweeps, omega) - return - - # Pre-smoothing - jacobi_smooth(u, f, h2, n_pre, omega) - - # Compute residual on fine grid - #r = cnp.zeros(u.shape, dtype=cnp.float64) - r = compute_residual(u, f, h2) - - # Restrict residual to coarse grid - n_coarse = (n - 1) // 2 + 1 - r_coarse = cnp.zeros((n_coarse, n_coarse), dtype=cnp.float64) - restrict_full_weighting(r, r_coarse) - - # Solve the coarse error equation recursively. - e_coarse = cnp.zeros((n_coarse, n_coarse), dtype=cnp.float64) - vcycle(e_coarse, r_coarse, 2.0 * h, n_pre, n_post, coarse_sweeps) - - # Prolongate and correct - prolongate_and_correct(e_coarse, u) - - # Post-smoothing - jacobi_smooth(u, f, h2, n_post, omega) - - -def vcycle_capture_backend(u, f, h, n_pre, n_post, coarse_sweeps=400): - """Build one backend V-cycle DAG and retain top-level intermediate arrays.""" - n = u.shape[0] - h2 = h * h - omega = 2.0 / 3.0 - - if n <= 5: - jacobi_smooth(u, f, h2, coarse_sweeps, omega) - return { - "u_pre": snapshot_backend(u), - "r": None, - "r_coarse": None, - "e_coarse": snapshot_backend(u), - "u_corr": snapshot_backend(u), - "u_final": u, - } - - jacobi_smooth(u, f, h2, n_pre, omega) - u_pre = snapshot_backend(u) - - r = compute_residual(u, f, h2) - - n_coarse = (n - 1) // 2 + 1 - r_coarse = cnp.zeros((n_coarse, n_coarse), dtype=cnp.float64) - restrict_full_weighting(r, r_coarse) - - e_coarse = cnp.zeros((n_coarse, n_coarse), dtype=cnp.float64) - vcycle(e_coarse, r_coarse, 2.0 * h, n_pre, n_post, coarse_sweeps) - e_coarse_snapshot = snapshot_backend(e_coarse) - - prolongate_and_correct(e_coarse, u) - u_corr = snapshot_backend(u) - - jacobi_smooth(u, f, h2, n_post, omega) - return { - "u_pre": u_pre, - "r": r, - "r_coarse": r_coarse, - "e_coarse": e_coarse_snapshot, - "u_corr": u_corr, - "u_final": u, - } - - -def validate_restriction(interface, n): - """Compare backend restriction against a NumPy reference.""" - fine_ref = np.zeros((n, n), dtype=np.float64) - apply_validation_pattern_numpy(fine_ref) - coarse_ref = restrict_full_weighting_numpy(fine_ref) - - fine_probe = cnp.zeros((n, n), dtype=cnp.float64) - apply_validation_pattern_backend(fine_probe) - fine_backend = fine_probe.get(interface) - print("") - fine_max_abs, fine_l2 = summarize_diff( - "Fine-grid validation pattern against NumPy", fine_backend, fine_ref - ) - - coarse_from_backend_fine = restrict_full_weighting_numpy(fine_backend) - - fine = cnp.zeros((n, n), dtype=cnp.float64) - coarse = cnp.zeros(((n - 1) // 2 + 1, (n - 1) // 2 + 1), dtype=cnp.float64) - apply_validation_pattern_backend(fine) - restrict_full_weighting(fine, coarse) - coarse_backend = coarse.get(interface) - - print("") - coarse_max_abs, coarse_l2 = summarize_diff( - "Restriction validation against NumPy", coarse_backend, coarse_ref - ) - print("") - summarize_diff( - "Restriction validation against backend fine-grid snapshot", - coarse_backend, - coarse_from_backend_fine, - ) - - return max(fine_max_abs, coarse_max_abs), max(fine_l2, coarse_l2) - - -def validate_prolongation(interface, n): - """Compare backend prolongation/correction against a NumPy reference.""" - if (n - 1) % 2 != 0: - raise ValueError("validate-prolong requires n = 2^k + 1") - - n_coarse = (n - 1) // 2 + 1 - coarse_ref = np.zeros((n_coarse, n_coarse), dtype=np.float64) - apply_validation_pattern_numpy(coarse_ref) - fine_ref = prolongate_numpy(coarse_ref) - - e_coarse = cnp.zeros((n_coarse, n_coarse), dtype=cnp.float64) - u_fine = cnp.zeros((n, n), dtype=cnp.float64) - apply_validation_pattern_backend(e_coarse) - u_fine[:, :] = 0.0 - prolongate_and_correct(e_coarse, u_fine) - fine_backend = u_fine.get(interface) - - print("") - return summarize_diff( - "Prolongation validation against NumPy", fine_backend, fine_ref - ) - - -def validate_jacobi(interface, n, n_smooth=1): - """Compare backend weighted Jacobi against a NumPy reference.""" - h = 1.0 / (n - 1) - h2 = h * h - omega = 2.0 / 3.0 - - u_ref = np.zeros((n, n), dtype=np.float64) - f_ref = np.zeros((n, n), dtype=np.float64) - apply_validation_pattern_numpy(u_ref) - apply_validation_pattern_numpy(f_ref) - u_expected = jacobi_smooth_numpy(u_ref, f_ref, h2, n_smooth, omega) - - u_backend = cnp.zeros((n, n), dtype=cnp.float64) - f_backend = cnp.zeros((n, n), dtype=cnp.float64) - apply_validation_pattern_backend(u_backend) - apply_validation_pattern_backend(f_backend) - jacobi_smooth(u_backend, f_backend, h2, n_smooth, omega) - u_actual = u_backend.get(interface) - - print("") - return summarize_diff( - f"Jacobi validation ({n_smooth} sweep{'s' if n_smooth != 1 else ''}) against NumPy", - u_actual, - u_expected, - ) - - -def validate_residual(interface, n): - """Compare backend residual formation against a NumPy reference.""" - h = 1.0 / (n - 1) - h2 = h * h - - u_ref = np.zeros((n, n), dtype=np.float64) - f_ref = np.zeros((n, n), dtype=np.float64) - apply_validation_pattern_numpy(u_ref) - apply_validation_pattern_numpy(f_ref) - r_expected = compute_residual_numpy(u_ref, f_ref, h2) - - u_backend = cnp.zeros((n, n), dtype=cnp.float64) - f_backend = cnp.zeros((n, n), dtype=cnp.float64) - apply_validation_pattern_backend(u_backend) - apply_validation_pattern_backend(f_backend) - r_backend = compute_residual(u_backend, f_backend, h2) - r_actual = r_backend.get(interface) - - print("") - return summarize_diff( - "Residual validation against NumPy", - r_actual, - r_expected, - ) - - -def validate_vcycle(interface, n, n_pre, n_post, coarse_sweeps): - """Compare one backend V-cycle against a NumPy reference.""" - h = 1.0 / (n - 1) - - u_ref = np.zeros((n, n), dtype=np.float64) - f_ref = np.zeros((n, n), dtype=np.float64) - f_ref[1:-1, 1:-1] = 1.0 - u_expected = vcycle_numpy(u_ref, f_ref, h, n_pre, n_post, coarse_sweeps) - - u_backend = cnp.zeros((n, n), dtype=cnp.float64) - f_backend = cnp.zeros((n, n), dtype=cnp.float64) - u_backend[:, :] = 0.0 - f_backend[:, :] = 0.0 - f_backend[1:-1, 1:-1] = 1.0 - vcycle(u_backend, f_backend, h, n_pre, n_post, coarse_sweeps) - u_actual = u_backend.get(interface) - - print("") - return summarize_diff( - "One V-cycle validation against NumPy", - u_actual, - u_expected, - ) - - -def validate_vcycles(interface, n, cycles, n_pre, n_post, coarse_sweeps): - """Compare multiple backend V-cycles in one flushed DAG against NumPy.""" - h = 1.0 / (n - 1) - - u_ref = np.zeros((n, n), dtype=np.float64) - f_ref = np.zeros((n, n), dtype=np.float64) - f_ref[1:-1, 1:-1] = 1.0 - u_expected = vcycles_numpy(u_ref, f_ref, h, cycles, n_pre, n_post, coarse_sweeps) - - u_backend = cnp.zeros((n, n), dtype=cnp.float64) - f_backend = cnp.zeros((n, n), dtype=cnp.float64) - u_backend[:, :] = 0.0 - f_backend[:, :] = 0.0 - f_backend[1:-1, 1:-1] = 1.0 - for _ in range(cycles): - vcycle(u_backend, f_backend, h, n_pre, n_post, coarse_sweeps) - u_actual = u_backend.get(interface) - - print("") - return summarize_diff( - f"{cycles} batched V-cycle{'s' if cycles != 1 else ''} validation against NumPy", - u_actual, - u_expected, - ) - - -def validate_vcycle_stages(interface, n, n_pre, n_post, coarse_sweeps): - """Compare top-level intermediate arrays from one backend V-cycle.""" - h = 1.0 / (n - 1) - - u_ref = np.zeros((n, n), dtype=np.float64) - f_ref = np.zeros((n, n), dtype=np.float64) - f_ref[1:-1, 1:-1] = 1.0 - expected = vcycle_numpy_capture(u_ref, f_ref, h, n_pre, n_post, coarse_sweeps) - - u_backend = cnp.zeros((n, n), dtype=cnp.float64) - f_backend = cnp.zeros((n, n), dtype=cnp.float64) - u_backend[:, :] = 0.0 - f_backend[:, :] = 0.0 - f_backend[1:-1, 1:-1] = 1.0 - actual_arrays = vcycle_capture_backend(u_backend, f_backend, h, n_pre, n_post, coarse_sweeps) - - stage_order = ("u_pre", "r", "r_coarse", "e_coarse", "u_corr", "u_final") - execute(interface) - print("") - worst_max = 0.0 - worst_l2 = 0.0 - for stage in stage_order: - backend_arr = actual_arrays[stage] - reference = expected[stage] - if backend_arr is None or reference is None or reference.size == 0: - continue - backend = backend_arr.get(interface) - max_abs, l2_err = summarize_diff( - f"V-cycle stage `{stage}` against NumPy", - backend, - reference, - ) - print("") - worst_max = max(worst_max, max_abs) - worst_l2 = max(worst_l2, l2_err) - return worst_max, worst_l2 - - -def parse_args(): - parser = argparse.ArgumentParser(description="Run or validate the 2D multigrid example.") - parser.add_argument("--mode", choices=( - "solve", - "validate-restrict", - "validate-prolong", - "validate-residual", - "validate-jacobi", - "validate-vcycle", - "validate-vcycles", - "validate-vcycle-stages", - "validate-all", - "both"), - default="solve", - help="Run the full V-cycle solve or targeted NumPy validation checks.") - parser.add_argument("--host", default="192.168.1.115", - help="Charm++ server host") - parser.add_argument("--port", type=int, default=1234, - help="Charm++ server port") - parser.add_argument("--odf", type=int, default=4, - help="Object decomposition factor") - parser.add_argument("--n", type=int, default=129, - help="Grid size (must be 2^k + 1)") - parser.add_argument("--pre", type=int, default=5, - help="Number of pre-smoothing sweeps") - parser.add_argument("--post", type=int, default=5, - help="Number of post-smoothing sweeps") - parser.add_argument("--coarse-sweeps", type=int, default=400, - help="Jacobi sweeps on the coarsest grid") - parser.add_argument("--cycles", type=int, default=10, - help="Number of V-cycles to run") - parser.add_argument("--check-every", type=int, default=2, - help="Fetch the solution and print a residual every N cycles") - return parser.parse_args() - - -def main(): - args = parse_args() - # Grid parameters: n x n grid on [0,1]^2 - # n must be 2^k + 1 for multigrid coarsening to work cleanly - n = args.n - h = 1.0 / (n - 1) - - # Solution array with zero initial guess - u = cnp.zeros((n, n), dtype=cnp.float64) - - # Boundary conditions: u = 0 on all boundaries (already set) - # Could set non-trivial BCs here, e.g.: - # u[0, :] = 1.0 # top boundary = 1 - - # Right-hand side: f = 2*pi^2 * sin(pi*x) * sin(pi*y) - # (exact solution is u = sin(pi*x) * sin(pi*y)) - f = cnp.zeros((n, n), dtype=cnp.float64) - - # Since we can't fill f point-by-point efficiently in this DSL, - # we set a uniform RHS for demonstration. - # f = 1.0 gives Poisson equation -Lap(u) = 1 with u=0 on boundary. - f[1:-1, 1:-1] = 1.0 - - # Connect to backend - interface = cnp.CharmNumericInterface() - interface.connect(args.host, args.port, args.odf) - set_auto_flush(interface, 1000) - - if args.mode in ("validate-restrict", "both"): - validate_restriction(interface, n) - if args.mode == "validate-restrict": - return - - if args.mode in ("validate-prolong", "validate-all"): - validate_prolongation(interface, n) - if args.mode == "validate-prolong": - return - - if args.mode in ("validate-residual", "validate-all"): - validate_residual(interface, n) - if args.mode == "validate-residual": - return - - if args.mode in ("validate-jacobi", "validate-all"): - validate_jacobi(interface, n, n_smooth=args.pre) - if args.mode == "validate-jacobi": - return - - if args.mode in ("validate-vcycle", "validate-all"): - validate_vcycle(interface, n, args.pre, args.post, args.coarse_sweeps) - if args.mode == "validate-vcycle": - return - - if args.mode in ("validate-vcycles", "validate-all"): - validate_vcycles(interface, n, args.cycles, args.pre, args.post, args.coarse_sweeps) - if args.mode == "validate-vcycles": - return - - if args.mode in ("validate-vcycle-stages", "validate-all"): - validate_vcycle_stages(interface, n, args.pre, args.post, args.coarse_sweeps) - if args.mode == "validate-vcycle-stages": - return - - if args.mode == "validate-all": - return - - # V-cycle iterations - # Slightly stronger smoothing tends to make the residual trend more monotone. - n_pre = args.pre - n_post = args.post - coarse_sweeps = args.coarse_sweeps - n_cycles = args.cycles - check_every = max(1, args.check_every) - - for cycle in range(n_cycles): - vcycle(u, f, h, n_pre, n_post, coarse_sweeps) - - # Execute and check solution periodically - if (cycle + 1) % check_every == 0: - result = u.get(interface) - # Compute residual norm using numpy on retrieved data - r = np.zeros_like(result) - r[1:-1, 1:-1] = result[:-2, 1:-1] + result[2:, 1:-1] + \ - result[1:-1, :-2] + result[1:-1, 2:] - \ - 4 * result[1:-1, 1:-1] - r[1:-1, 1:-1] = 1.0 - (-r[1:-1, 1:-1] / (h * h)) # f - (-Lap u) - res_norm = np.linalg.norm(r) - print(f"V-cycle {cycle + 1}: residual norm = {res_norm:.6e}") - - # Final result - execute(interface) - result = u.get(interface) - print(f"\nFinal solution: min = {result.min():.6f}, max = {result.max():.6f}") - print(f"Grid size: {n} x {n}, h = {h:.6f}") - - -if __name__ == '__main__': - main() diff --git a/examples/reduction.py b/examples/reduction.py deleted file mode 100644 index 7253f41..0000000 --- a/examples/reduction.py +++ /dev/null @@ -1,10 +0,0 @@ -import charmnumeric as cnp - -N = 64 -x = cnp.ones((N,), dtype=cnp.float32) -res = x @ x -x = x / res - -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.115', 1234, 4) -print(x.get(interface)) diff --git a/examples/simple.py b/examples/simple.py index 5b694da..463874b 100644 --- a/examples/simple.py +++ b/examples/simple.py @@ -1,25 +1,22 @@ -import charmnumeric as cnp -from charmtyles.core import execute +from charmtiles.array import connect, ndarray +import charmtiles.linalg as lg +import numpy as np -a = cnp.zeros((128,), dtype=cnp.float32) -b = cnp.zeros((128,), dtype=cnp.float32) -#c = cnp.zeros(10, dtype=cnp.float32) +def f(): + vnp = np.array([2, 1, 1], dtype=np.float64) + Anp = np.array([[1, 2, 3], [4, 5, 6]], dtype=np.float64) + bnp = np.array([1, 0], dtype=np.float64) + v = ndarray(1, 3, np.float64, nparr=vnp) + b = ndarray(1, 2, np.float64, nparr=bnp) + A = ndarray(2, (2, 3), np.float64, nparr=Anp) + Av = A @ v + Avnorm = Av @ b + print(Avnorm.get()) + #print("Actual =", xnp @ ynp) -x = 2 * (a + b + 3) -x[:50] = x[70:120] + 1 -#y = 2 * x -#x[1:50] = 1#3 * c[1:50] -#x[49:100] = 2#b[49:100] + 5 -#z = a + x -# a[0, :, :10] = c[0] -# a[0, :, 10:] = 2 +if __name__ == '__main__': + connect("172.17.0.1", 10000) + s = f() -# b[0] = a[0] * 3 -#plot_execution_state() -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.114', 1234, 4) -execute(interface) - -print(x.get(interface)) diff --git a/examples/test.py b/examples/test.py deleted file mode 100644 index f643622..0000000 --- a/examples/test.py +++ /dev/null @@ -1,31 +0,0 @@ -"""Conjugate Gradient solver using charmnumeric. - -Solves the linear system A x = b where A is symmetric positive-definite. -We build a simple SPD matrix A = I + ones (the identity plus a constant -matrix) so the answer is easy to verify with NumPy. -""" - -import charmnumeric as cnp -import numpy as np - -N = 65 - -# --- Build a symmetric positive-definite matrix A = (N+1)*I + ones -------- -# This is SPD because eigenvalues are (N+1) (multiplicity N-1) and (2N+1). -h = (N - 1) // 2 + 1 -x = cnp.zeros((h, h), dtype=cnp.float64) -y = cnp.ones((N, N), dtype=cnp.float64) -x[1:-1, 1:-1] = y[2:-2:2, 2:-2:2] - -interface = cnp.CharmNumericInterface() -interface.connect('192.168.1.115', 1234, 4) - -res = x.get(interface) -expected = np.zeros((h, h), dtype=np.float64) -expected[1:-1, 1:-1] = np.ones((N, N), dtype=np.float64)[2:-2:2, 2:-2:2] - -print("Result x:") -print(res.flatten()) -print("Expected x:") -print(expected.flatten()) -assert np.array_equal(res, expected), "charmnumeric result does not match NumPy" diff --git a/pyproject.toml b/pyproject.toml deleted file mode 100644 index 35219b8..0000000 --- a/pyproject.toml +++ /dev/null @@ -1,8 +0,0 @@ -[build-system] -requires = [ - "setuptools>=64", - "wheel", - "Cython", - "numpy", -] -build-backend = "setuptools.build_meta" diff --git a/pytest.ini b/pytest.ini deleted file mode 100644 index e43330f..0000000 --- a/pytest.ini +++ /dev/null @@ -1,4 +0,0 @@ -[pytest] -testpaths = tests -markers = - integration: requires a built charmnumeric backend and Charm++ runtime diff --git a/scripts/run_integration_tests.sh b/scripts/run_integration_tests.sh deleted file mode 100644 index 1ce8d54..0000000 --- a/scripts/run_integration_tests.sh +++ /dev/null @@ -1,43 +0,0 @@ -#!/usr/bin/env bash -set -euo pipefail - -ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../../.." && pwd)" -CHARMNUMERIC_ROOT="${ROOT}/example/charmnumeric" -CORE_BUILD_DIR="${CHARMTYLES_CORE_BUILD_DIR:-${ROOT}/src/charmtyles/core/build}" -CHARMNUMERIC_BUILD_DIR="${CHARMNUMERIC_BUILD_DIR:-${CHARMNUMERIC_ROOT}/src/build}" - -if [[ -z "${CHARM_HOME:-}" ]]; then - echo "CHARM_HOME must be set before running charmnumeric integration tests." >&2 - exit 1 -fi - -python - <<'PY' -import importlib - -for module in ("pyccs",): - try: - importlib.import_module(module) - except ModuleNotFoundError as exc: - raise SystemExit(f"Missing Python dependency: {module}") from exc -PY - -parallelism="${CMAKE_BUILD_PARALLEL_LEVEL:-$(getconf _NPROCESSORS_ONLN 2>/dev/null || echo 4)}" - -cmake -S "${ROOT}/src/charmtyles/core" -B "${CORE_BUILD_DIR}" -cmake --build "${CORE_BUILD_DIR}" --parallel "${parallelism}" - -cmake \ - -S "${CHARMNUMERIC_ROOT}/src" \ - -B "${CHARMNUMERIC_BUILD_DIR}" \ - -DCHARMTYLES_CORE_BUILD_DIR="${CORE_BUILD_DIR}" -cmake --build "${CHARMNUMERIC_BUILD_DIR}" --target server --parallel "${parallelism}" - -python -m pip install -e "${ROOT}" -python -m pip install -e "${CHARMNUMERIC_ROOT}[tests]" - -export CHARMNUMERIC_SERVER="${CHARMNUMERIC_BUILD_DIR}/server.out" -if [[ -x "${CHARMNUMERIC_BUILD_DIR}/charmrun" ]]; then - export CHARMNUMERIC_CHARMRUN="${CHARMNUMERIC_BUILD_DIR}/charmrun" -fi - -python -m pytest "${CHARMNUMERIC_ROOT}/tests" -m integration -v "$@" diff --git a/setup.py b/setup.py index 19f0fb2..347b391 100644 --- a/setup.py +++ b/setup.py @@ -1,65 +1,24 @@ +import sys import os -import re -from pathlib import Path - -import numpy as np -from Cython.Build import cythonize -from setuptools import Extension, setup, find_packages - - -ROOT = Path(__file__).resolve().parent -CHARMTYLES_HOME = ROOT.parent.parent -os.chdir(ROOT) +import subprocess +from setuptools import setup, find_packages def get_version(): data = {} - fname = ROOT / 'charmnumeric' / '__init__.py' - exec(compile(fname.read_text(), str(fname), 'exec'), data) + fname = os.path.join('charmtiles', '__init__.py') + exec(compile(open(fname).read(), fname, 'exec'), data) return data.get('__version__') -def find_charm_include(): - env_include = os.environ.get('CHARM_INCLUDE_DIR') - if env_include: - candidate = Path(env_include) - if (candidate / 'charm++.h').exists(): - return str(candidate) - - for key in ('CHARM_HOME', 'CHARM_DIR', 'CHARM_ROOT'): - value = os.environ.get(key) - if value: - candidate = Path(value) / 'include' - if (candidate / 'charm++.h').exists(): - return str(candidate) - - depfiles = [ - ROOT / 'src' / 'build' / 'CMakeFiles' / 'backend_obj.dir' / 'compiler_depend.make', - CHARMTYLES_HOME / 'src' / 'charmtyles' / 'core' / 'build' / 'CMakeFiles' / 'charmtyles_core.dir' / 'compiler_depend.make', - ] - pattern = re.compile(r'(/[^\s]*?/charm\+\+\.h)') - for depfile in depfiles: - if not depfile.exists(): - continue - match = pattern.search(depfile.read_text()) - if match: - return str(Path(match.group(1)).parent) - - home = Path.home() - for candidate in ( - home / 'charm' / 'mpi-linux-x86_64' / 'include', - home / 'charm-gpu' / 'include', - home / 'charm' / 'include', - ): - if (candidate / 'charm++.h').exists(): - return str(candidate) - - raise RuntimeError( - 'Could not locate Charm++ headers. Set CHARM_HOME or CHARM_INCLUDE_DIR before building charmnumeric.' - ) +def compile_server(): + charmc = os.environ.get('CHARMC', '~/charm/netlrts-linux-x86_64/bin/charmc') + aum_base = os.environ.get('AUM_HOME', '~/LibAum') + subprocess.run(["make", "-C", "src/", + "CHARMC=%s" % charmc, "BASE_DIR=%s" % aum_base]) -install_requires = ['numpy', 'Cython'] +install_requires = ['numpy', 'charm4py'] tests_require = ['pytest'] docs_require = ['sphinx'] @@ -79,33 +38,19 @@ def find_charm_include(): ''' classifiers = [x.strip() for x in classes.splitlines() if x] -extensions = [ - Extension( - 'charmnumeric._native_region', - [str(ROOT / 'charmnumeric' / '_native_region.pyx')], - language='c++', - include_dirs=[ - np.get_include(), - str(ROOT / 'src'), - str(CHARMTYLES_HOME / 'include'), - find_charm_include(), - ], - extra_compile_args=['-std=c++17', '-O3', '-DNDEBUG'], - ) -] +compile_server() setup( - name='charmnumeric', - #version=get_version(), + name='charmtiles', + version=get_version(), author='Aditya Bhosale', author_email='adityapb1546@gmail.com', - description='A framework for writing DSLs', - long_description=(ROOT / 'README.rst').read_text(), + description='A python library for distributed array computations', + long_description=open('README.rst').read(), license="BSD", - #url='https://github.com/UIUC-PPL/PyProject', + url='https://github.com/UIUC-PPL/PyProject', classifiers=classifiers, packages=find_packages(), - ext_modules=cythonize(extensions, language_level='3', include_path=[str(CHARMTYLES_HOME)]), install_requires=install_requires, extras_require={ "docs": docs_require, diff --git a/src/CMakeLists.txt b/src/CMakeLists.txt deleted file mode 100644 index dffceec..0000000 --- a/src/CMakeLists.txt +++ /dev/null @@ -1,334 +0,0 @@ -cmake_minimum_required(VERSION 3.20) -project(charmnumeric_backend LANGUAGES CXX) - -set(CMAKE_CXX_STANDARD 17) -set(CMAKE_CXX_STANDARD_REQUIRED ON) -option(NDEBUG "Enable NDEBUG preprocessor define (-DNDEBUG)" ON) - -# ===================================================================== -# Charm++ -- located via CHARM_HOME environment variable -# ===================================================================== -if(NOT DEFINED ENV{CHARM_HOME} AND NOT DEFINED CHARM_HOME) - message(FATAL_ERROR "CHARM_HOME is not set. Point it to your Charm++ installation.") -endif() -if(DEFINED ENV{CHARM_HOME} AND NOT DEFINED CHARM_HOME) - set(CHARM_HOME "$ENV{CHARM_HOME}") -endif() - -set(CHARMC "${CHARM_HOME}/bin/charmc") - -# Charmtyles root (three levels up from this CMakeLists.txt) -set(CHARMTYLES_HOME "${CMAKE_CURRENT_SOURCE_DIR}/../../..") - -# ===================================================================== -# MLIR / LLVM -# ===================================================================== -find_package(MLIR REQUIRED CONFIG) -find_package(LLVM REQUIRED CONFIG) - -message(STATUS "Using MLIRConfig.cmake in: ${MLIR_DIR}") -message(STATUS "Using LLVMConfig.cmake in: ${LLVM_DIR}") - -list(APPEND CMAKE_MODULE_PATH "${MLIR_CMAKE_DIR}") -list(APPEND CMAKE_MODULE_PATH "${LLVM_CMAKE_DIR}") - -include(TableGen) -include(AddLLVM) -include(AddMLIR) - -include_directories(${LLVM_INCLUDE_DIRS}) -include_directories(${MLIR_INCLUDE_DIRS}) - -# ===================================================================== -# Include paths -# ===================================================================== -include_directories(${CHARM_HOME}/include) -include_directories(${CHARMTYLES_HOME}/include) -include_directories(${CMAKE_CURRENT_SOURCE_DIR}) # for local headers -include_directories(${CMAKE_CURRENT_BINARY_DIR}) # for generated decl/def.h - -# ===================================================================== -# Eigen (header-only) -# ===================================================================== -find_package(Eigen3 REQUIRED NO_MODULE) -include_directories(${EIGEN3_INCLUDE_DIR}) - -# ===================================================================== -# Charm++ Interface (.ci) code generation -# ===================================================================== -# 1. Core charmtyles CI file → charmtyles.decl.h / charmtyles.def.h -set(CHARMTYLES_CI "${CHARMTYLES_HOME}/include/charmtyles/core/charmtyles.ci") -add_custom_command( - OUTPUT ${CMAKE_CURRENT_BINARY_DIR}/charmtyles.decl.h - ${CMAKE_CURRENT_BINARY_DIR}/charmtyles.def.h - COMMAND ${CHARMC} -E -module CommonLBs - ${CHARMTYLES_CI} - WORKING_DIRECTORY ${CMAKE_CURRENT_BINARY_DIR} - DEPENDS ${CHARMTYLES_CI} - COMMENT "Generating charmtyles.decl.h / charmtyles.def.h" -) - -# Pass -DUSE_KOKKOS to charmc when GPU backend is enabled (for nocopydevice in .ci) -set(CHARMC_CI_DEFS "") -if(DEFINED GPU_BACKEND) - set(CHARMC_CI_DEFS "-DUSE_KOKKOS") -endif() - -# 2. charmnumeric backend CI file → charmnumeric.decl.h / charmnumeric.def.h -set(BACKEND_CI "${CMAKE_CURRENT_SOURCE_DIR}/backend.ci") -add_custom_command( - OUTPUT ${CMAKE_CURRENT_BINARY_DIR}/charmnumeric.decl.h - ${CMAKE_CURRENT_BINARY_DIR}/charmnumeric.def.h - COMMAND ${CHARMC} -E -module CommonLBs - ${CHARMC_CI_DEFS} - ${BACKEND_CI} - WORKING_DIRECTORY ${CMAKE_CURRENT_BINARY_DIR} - DEPENDS ${BACKEND_CI} - ${CMAKE_CURRENT_BINARY_DIR}/charmtyles.decl.h - COMMENT "Generating charmnumeric.decl.h / charmnumeric.def.h" -) - -# Convenience target so other targets can depend on the generated headers -add_custom_target(ci_headers - DEPENDS ${CMAKE_CURRENT_BINARY_DIR}/charmtyles.decl.h - ${CMAKE_CURRENT_BINARY_DIR}/charmtyles.def.h - ${CMAKE_CURRENT_BINARY_DIR}/charmnumeric.decl.h - ${CMAKE_CURRENT_BINARY_DIR}/charmnumeric.def.h -) - -# ===================================================================== -# Build the server executable via charmc -# -# charmc is used as the linker because it injects the Charm++ runtime, -# main trampoline, and module registration. We compile the C++ source -# with CMake (so MLIR/LLVM flags are picked up automatically), then -# link the resulting object with charmc. -# ===================================================================== - -# Locate pre-built charmtyles_core static library. -# Set CHARMTYLES_CORE_BUILD_DIR to the core build directory. -if(NOT DEFINED CHARMTYLES_CORE_BUILD_DIR) - set(CHARMTYLES_CORE_BUILD_DIR "${CHARMTYLES_HOME}/src/charmtyles/core/build") -endif() -set(CHARMTYLES_CORE_LIB "${CHARMTYLES_CORE_BUILD_DIR}/libcharmtyles_core.a") -if(NOT EXISTS ${CHARMTYLES_CORE_LIB}) - message(FATAL_ERROR "charmtyles_core not found at ${CHARMTYLES_CORE_LIB}.\n" - "Build it first: cd ${CHARMTYLES_HOME}/src/charmtyles/core/build && cmake .. && make") -endif() -# Add core's generated headers to include path -include_directories(${CHARMTYLES_CORE_BUILD_DIR}) - -# -- JIT module (heaviest TU — all MLIR/LLVM codegen) -- -add_library(jit_obj OBJECT jit.cpp) -add_dependencies(jit_obj ci_headers) -set(SIMPLE_ARRAY_DEFS) -if(NDEBUG) - list(APPEND SIMPLE_ARRAY_DEFS -DNDEBUG) -endif() -target_compile_options(jit_obj PRIVATE -O3 -g ${SIMPLE_ARRAY_DEFS}) - -# -- Backend module (Charm++ classes + JIT integration) -- -set(BACKEND_GROUP_SOURCES - dag_group_runtime.cpp - dag_group_receive.cpp - dag_group_compile.cpp -) - -set(BACKEND_NODE_SOURCES - execute_node.cpp - execute_node_regions.cpp -) - -set(BACKEND_EXECUTOR_SOURCES - executor_reducer.cpp - executor_core.cpp - executor_reduce.cpp - executor_matmul.cpp - executor_matmatmul.cpp - executor_transfer.cpp - executor_incremental.cpp - executor_ast_visitor.cpp -) - -set(BACKEND_PARTITION_SOURCES - partition_lifecycle.cpp - partition_comm.cpp - partition_get.cpp -) - -add_library(backend_obj OBJECT - ${BACKEND_GROUP_SOURCES} - ${BACKEND_NODE_SOURCES} - ${BACKEND_EXECUTOR_SOURCES} - ${BACKEND_PARTITION_SOURCES} -) -add_dependencies(backend_obj ci_headers) -target_compile_options(backend_obj PRIVATE -O3 -g ${SIMPLE_ARRAY_DEFS}) - -# -- Server module (Main class, startup) -- -add_library(server_obj OBJECT server.cpp) -add_dependencies(server_obj ci_headers) -target_compile_options(server_obj PRIVATE -O3 -g ${SIMPLE_ARRAY_DEFS}) - -# GPU backend selection (pass -DGPU_BACKEND=NVIDIA|AMD|INTEL on cmake line) -if(DEFINED GPU_BACKEND) - string(TOUPPER "${GPU_BACKEND}" GPU_BACKEND_UPPER) - - # Kokkos provides portable memory management for all GPU backends. - # Pre-build Kokkos with the matching backend (Kokkos_ENABLE_CUDA, etc.) - find_package(Kokkos REQUIRED) - find_package(KokkosKernels REQUIRED) - target_link_libraries(backend_obj PRIVATE Kokkos::kokkos Kokkos::kokkoskernels) - target_link_libraries(server_obj PRIVATE Kokkos::kokkos) - target_compile_definitions(backend_obj PRIVATE USE_KOKKOS) - target_compile_definitions(server_obj PRIVATE USE_KOKKOS) - - # JIT dispatch still needs vendor driver APIs for kernel launch - if(GPU_BACKEND_UPPER STREQUAL "NVIDIA") - target_compile_definitions(jit_obj PRIVATE USE_NVIDIA JIT_ENABLE_GPU_BACKEND) - target_compile_definitions(backend_obj PRIVATE USE_NVIDIA JIT_ENABLE_GPU_BACKEND) - find_package(CUDAToolkit REQUIRED) - target_link_libraries(jit_obj PRIVATE CUDA::cuda_driver) - target_link_libraries(backend_obj PRIVATE CUDA::cuda_driver) - elseif(GPU_BACKEND_UPPER STREQUAL "AMD") - target_compile_definitions(jit_obj PRIVATE USE_AMD JIT_ENABLE_GPU_BACKEND) - target_compile_definitions(backend_obj PRIVATE USE_AMD JIT_ENABLE_GPU_BACKEND) - find_package(hip REQUIRED) - target_link_libraries(jit_obj PRIVATE hip::host) - target_link_libraries(backend_obj PRIVATE hip::host) - elseif(GPU_BACKEND_UPPER STREQUAL "INTEL") - target_compile_definitions(jit_obj PRIVATE USE_INTEL JIT_ENABLE_GPU_BACKEND) - target_compile_definitions(backend_obj PRIVATE USE_INTEL JIT_ENABLE_GPU_BACKEND) - find_library(LEVEL_ZERO_LIB NAMES ze_loader) - find_path(LEVEL_ZERO_INCLUDE_DIR NAMES level_zero/ze_api.h) - if(LEVEL_ZERO_LIB AND LEVEL_ZERO_INCLUDE_DIR) - target_include_directories(jit_obj PRIVATE ${LEVEL_ZERO_INCLUDE_DIR}) - target_include_directories(backend_obj PRIVATE ${LEVEL_ZERO_INCLUDE_DIR}) - target_link_libraries(jit_obj PRIVATE ${LEVEL_ZERO_LIB}) - target_link_libraries(backend_obj PRIVATE ${LEVEL_ZERO_LIB}) - else() - message(FATAL_ERROR "Level Zero not found. Install oneAPI Level Zero Loader.") - endif() - else() - message(WARNING "Unknown GPU_BACKEND '${GPU_BACKEND}', building CPU-only.") - endif() -endif() - -# MLIR/LLVM link libraries — needed by jit_obj and backend_obj -get_property(mlir_dialect_libs GLOBAL PROPERTY MLIR_DIALECT_LIBS) -get_property(mlir_conversion_libs GLOBAL PROPERTY MLIR_CONVERSION_LIBS) - -foreach(tgt jit_obj backend_obj) - target_link_libraries(${tgt} - PRIVATE - ${mlir_dialect_libs} - ${mlir_conversion_libs} - MLIRPass - MLIRTransforms - MLIRExecutionEngine - MLIRTargetLLVMIRExport - MLIRLLVMToLLVMIRTranslation - MLIRNVVMToLLVMIRTranslation - MLIRROCDLToLLVMIRTranslation - MLIRSPIRVSerialization - MLIRSPIRVTransforms - MLIRIR - MLIRSupport - MLIRParser - LLVMCore - LLVMSupport - ) -endforeach() - -# ===================================================================== -# Link the object file with charmc to produce the final Charm++ executable. -# -# charmc must be used as the linker (it injects the Charm++ runtime). -# MLIR/LLVM have deep transitive dependency chains, so we use -# llvm-config to get the complete set of LLVM libraries, and glob -# all MLIR static libraries from the install directory. -# ===================================================================== - -# -- Locate llvm-config -- -find_program(LLVM_CONFIG_EXE NAMES llvm-config - HINTS "${LLVM_TOOLS_BINARY_DIR}" "${LLVM_DIR}/../../../bin") -if(NOT LLVM_CONFIG_EXE) - message(FATAL_ERROR "llvm-config not found. Set LLVM_TOOLS_BINARY_DIR or add it to PATH.") -endif() -message(STATUS "Using llvm-config: ${LLVM_CONFIG_EXE}") - -# -- Get ALL LLVM link flags (resolves full transitive closure) -- -execute_process(COMMAND ${LLVM_CONFIG_EXE} --link-static --libfiles all - OUTPUT_VARIABLE LLVM_ALL_LIB_FILES - OUTPUT_STRIP_TRAILING_WHITESPACE) -separate_arguments(LLVM_ALL_LIB_FILES) - -execute_process(COMMAND ${LLVM_CONFIG_EXE} --system-libs - OUTPUT_VARIABLE LLVM_SYSTEM_LIBS - OUTPUT_STRIP_TRAILING_WHITESPACE) -separate_arguments(LLVM_SYSTEM_LIBS) - -execute_process(COMMAND ${LLVM_CONFIG_EXE} --ldflags - OUTPUT_VARIABLE LLVM_LD_FLAGS - OUTPUT_STRIP_TRAILING_WHITESPACE) -separate_arguments(LLVM_LD_FLAGS) - -# -- Glob all MLIR static libraries -- -file(GLOB MLIR_ALL_LIB_FILES "${LLVM_LIBRARY_DIR}/libMLIR*.a") - -# -- Combine into a single list for charmc -- -# Use --whole-archive for MLIR libs to avoid registration issues, -# then --no-whole-archive for LLVM libs. -set(ALL_STATIC_LIBS ${MLIR_ALL_LIB_FILES} ${LLVM_ALL_LIB_FILES}) - -# charmc does not forward absolute .a paths to the underlying linker. -# Write them to a response file and pass via -Wl,@file so they reach ld. -set(LINK_RESPONSE_FILE "${CMAKE_CURRENT_BINARY_DIR}/link_libs.rsp") -set(_rsp_content "--start-group\n${CHARMTYLES_CORE_LIB}\n") -foreach(lib ${ALL_STATIC_LIBS}) - string(APPEND _rsp_content "${lib}\n") -endforeach() -string(APPEND _rsp_content "--end-group\n") -foreach(lib ${LLVM_SYSTEM_LIBS}) - string(APPEND _rsp_content "${lib}\n") -endforeach() -file(WRITE ${LINK_RESPONSE_FILE} "${_rsp_content}") - -add_custom_command( - OUTPUT ${CMAKE_CURRENT_BINARY_DIR}/server.out - COMMAND ${CHARMC} - -c++-option -std=c++17 - -O3 -g $<$:-DNDEBUG> - -language charm++ - -module CommonLBs - ${LLVM_LD_FLAGS} - -o ${CMAKE_CURRENT_BINARY_DIR}/server.out - "$" - "$" - "$" - -Wl,@${LINK_RESPONSE_FILE} - -ldl - DEPENDS jit_obj backend_obj server_obj ci_headers - ${CHARMTYLES_CORE_LIB} - $ - $ - $ - COMMAND_EXPAND_LISTS - COMMENT "Linking server.out with charmc" -) - -add_custom_target(server ALL - DEPENDS ${CMAKE_CURRENT_BINARY_DIR}/server.out -) - -# ===================================================================== -# Convenience run target -# ===================================================================== -add_custom_target(run-server - COMMAND ${CMAKE_CURRENT_BINARY_DIR}/charmrun +p4 - ${CMAKE_CURRENT_BINARY_DIR}/server.out - ++server ++server-port 10000 - DEPENDS server - WORKING_DIRECTORY ${CMAKE_CURRENT_BINARY_DIR} - COMMENT "Running server with 4 PEs" -) diff --git a/src/Makefile b/src/Makefile new file mode 100644 index 0000000..4f6206f --- /dev/null +++ b/src/Makefile @@ -0,0 +1,20 @@ +CHARMC=/home/adityapb/charm/charm/netlrts-linux-x86_64/bin/charmc +BASE_DIR=/home/adityapb/charm/LibAum +LIBS_DIR=$(BASE_DIR) +OPTS=-c++-option -std=c++17 -O3 #-DNDEBUG + +all: server + +.PHONY: clean server.out + +server_ci: server.ci + $(CHARMC) -E server.ci + +server: server.cpp server_ci + $(CHARMC) $< -L$(LIBS_DIR)/aum -laum -I$(BASE_DIR) -I$(BASE_DIR)/aum/backend -o $@.out $(OPTS) + +run-server: server.out + ./charmrun +p4 ./server.out ++server ++server-port 10000 ++local + +clean: + rm *.decl.h *.def.h *.out charmrun diff --git a/src/array_region.hpp b/src/array_region.hpp deleted file mode 100644 index a5ed805..0000000 --- a/src/array_region.hpp +++ /dev/null @@ -1,866 +0,0 @@ -#pragma once - -#include "charm++.h" -#include -#include -#include -#include // extract -#include -#include -#include -#include - -#ifndef DBG_PRINT -#ifdef NDEBUG -#define DBG_PRINT(...) -#else -#define DBG_PRINT(...) CkPrintf(__VA_ARGS__) -#endif -#endif - -template -class ArrayRegion : public Region { - public: - std::array start, stop, step; - - ArrayRegion(std::array start_, std::array stop_, std::array step_) - : start(start_), stop(stop_), step(step_) {} - - ArrayRegion() : start{}, stop{}, step{} {} - - ArrayRegion(bool is_global_) : start{}, stop{}, step{} { is_global = is_global_; } - - static constexpr int ndims() { return N; } - - static int positive_mod(int value, int mod) { - int result = value % mod; - return result < 0 ? result + mod : result; - } - - static int gcd_int(int a, int b) { - a = std::abs(a); - b = std::abs(b); - while (b != 0) { - int t = a % b; - a = b; - b = t; - } - return a; - } - - static long long extended_gcd(long long a, long long b, long long& x, long long& y) { - if (b == 0) { - x = 1; - y = 0; - return a; - } - long long x1 = 0, y1 = 0; - long long g = extended_gcd(b, a % b, x1, y1); - x = y1; - y = x1 - (a / b) * y1; - return g; - } - - static int modular_inverse(int value, int mod) { - long long x = 0, y = 0; - long long g = extended_gcd(value, mod, x, y); - (void)g; - assert(g == 1 && "modular inverse requires coprime inputs"); - long long inv = x % mod; - if (inv < 0) - inv += mod; - return static_cast(inv); - } - - static bool first_aligned_point(int start_a, int step_a, int start_b, int step_b, int lo, - int hi, int& aligned, int& aligned_step) { - int g = gcd_int(step_a, step_b); - if (positive_mod(start_a - start_b, g) != 0) - return false; - - long long reduced_a = step_a / g; - long long reduced_b = step_b / g; - long long diff = (start_b - start_a) / g; - - long long t0 = 0; - if (reduced_b != 1) { - long long a_mod = reduced_a % reduced_b; - if (a_mod < 0) - a_mod += reduced_b; - t0 = (diff % reduced_b + reduced_b) % reduced_b; - t0 = (t0 * modular_inverse(static_cast(a_mod), static_cast(reduced_b))) % - reduced_b; - } - - long long first = static_cast(start_a) + static_cast(step_a) * t0; - long long period = (static_cast(step_a) / g) * step_b; - if (first < lo) { - long long k = (static_cast(lo) - first + period - 1) / period; - first += k * period; - } - if (first >= hi) - return false; - - aligned = static_cast(first); - aligned_step = static_cast(period); - return true; - } - - /// Number of elements along dimension d - int size(int d) const { return (stop[d] - start[d] + step[d] - 1) / step[d]; } - - /// Total number of elements across all dimensions - int size() const { - int total = 1; - for (int d = 0; d < N; ++d) - total *= size(d); - return total; - } - - /// Deserialize from byte buffer (matches Python ArrayRegion.serialize) - static ArrayRegion* deserialize(char*& msg) { - int is_global = extract(msg); - if (is_global) - return new ArrayRegion(true); - int nd = extract(msg); - assert(nd == N && "Deserialized ndims does not match template parameter N"); - std::array s{}, e{}, st{}; - for (int d = 0; d < N; ++d) { - s[d] = extract(msg); - e[d] = extract(msg); - st[d] = extract(msg); - } - return new ArrayRegion(s, e, st); - } - - bool overlaps(ArrayRegion const& other) const { - if (is_global || other.is_global) - return true; - if (start == other.start && stop == other.stop && step == other.step) - return true; - - for (int d = 0; d < N; ++d) { - int lo = std::max(start[d], other.start[d]); - int hi = std::min(stop[d], other.stop[d]); - if (lo >= hi) - return false; - - int aligned = 0; - int aligned_step = 0; - if (!first_aligned_point(start[d], step[d], other.start[d], other.step[d], lo, hi, - aligned, aligned_step)) - return false; - } - - return true; - } - - bool covers(ArrayRegion const& other) const { - if (is_global) - return true; - if (other.is_global) - return false; - if (start == other.start && stop == other.stop && step == other.step) - return true; - - for (int d = 0; d < N; ++d) { - if (other.start[d] < start[d] || other.stop[d] > stop[d]) - return false; - - if (step[d] == 1) - continue; - - if (positive_mod(other.start[d] - start[d], step[d]) != 0) - return false; - - if (other.size(d) <= 1) - continue; - - if (other.step[d] % step[d] != 0) - return false; - } - - return true; - } - - bool intersect(ArrayRegion const& other, ArrayRegion& result) const { - if (is_global && other.is_global) { - result = other; - return true; - } - if (is_global) { - result = other; - return true; - } - if (other.is_global) { - result = *this; - return true; - } - if (start == other.start && stop == other.stop && step == other.step) { - result = *this; - return true; - } - - for (int d = 0; d < N; ++d) { - int lo = std::max(start[d], other.start[d]); - int hi = std::min(stop[d], other.stop[d]); - if (lo >= hi) - return false; - - int aligned = 0; - int aligned_step = 0; - if (!first_aligned_point(start[d], step[d], other.start[d], other.step[d], lo, hi, - aligned, aligned_step)) - return false; - - result.start[d] = aligned; - result.stop[d] = hi; - result.step[d] = aligned_step; - } - - return true; - } - - bool overlaps(Region const& other) const override { - auto other_region = dynamic_cast const*>(&other); - return other_region != nullptr && overlaps(*other_region); - } - - bool covers(Region const& other) const override { - auto other_region = dynamic_cast const*>(&other); - return other_region != nullptr && covers(*other_region); - } - - bool intersect(Region const& other, Region& result) const override { - auto other_region = dynamic_cast const*>(&other); - auto result_region = dynamic_cast*>(&result); - return other_region != nullptr && result_region != nullptr && - intersect(*other_region, *result_region); - } - - void map_region(Region& other) override {} -}; - -template -ArrayRegion* make_array_region_handle(const int* start, const int* stop, const int* step, - bool is_global) { - if (is_global) - return new ArrayRegion(true); - - std::array s{}; - std::array e{}; - std::array st{}; - for (int d = 0; d < N; ++d) { - s[d] = start[d]; - e[d] = stop[d]; - st[d] = step[d]; - } - return new ArrayRegion(s, e, st); -} - -template -void delete_array_region_handle(void* ptr) { - delete reinterpret_cast*>(ptr); -} - -template -bool overlaps_array_region_handle(void* lhs, void* rhs) { - return reinterpret_cast*>(lhs)->overlaps(*reinterpret_cast*>(rhs)); -} - -template -bool covers_array_region_handle(void* lhs, void* rhs) { - return reinterpret_cast*>(lhs)->covers(*reinterpret_cast*>(rhs)); -} - -template -bool intersect_array_region_handle(void* lhs, void* rhs, int* out_start, int* out_stop, - int* out_step) { - ArrayRegion result; - bool ok = - reinterpret_cast*>(lhs)->intersect(*reinterpret_cast*>(rhs), result); - if (!ok) - return false; - - for (int d = 0; d < N; ++d) { - out_start[d] = result.start[d]; - out_stop[d] = result.stop[d]; - out_step[d] = result.step[d]; - } - return true; -} - -template -class ArrayDecomp { - public: - std::array offset; // per-dim offset into global tile grid - std::array global_shape; // array shape in each dimension - int tile; // tile size (same for all dims of this N) - - ArrayDecomp() : offset{}, global_shape{}, tile(0) {} - - /// Default decomposition (offset=0, standard tile). - static ArrayDecomp default_decomp(std::array shape, int tile_size) { - ArrayDecomp d; - d.offset = {}; - d.global_shape = shape; - d.tile = tile_size; - return d; - } - - /// Offset decomposition. - static ArrayDecomp offset_decomp(std::array shape, std::array off, - int tile_size) { - ArrayDecomp d; - d.offset = off; - d.global_shape = shape; - d.tile = tile_size; - return d; - } - - /// Map local coordinate to global coordinate along dimension d. - int to_global(int d, int local_coord) const { return local_coord + offset[d]; } - - /// Map global coordinate to local coordinate along dimension d. - int to_local(int d, int global_coord) const { return global_coord - offset[d]; } - - /// Map a local-space region to global-space region. - ArrayRegion to_global(ArrayRegion const& r) const { - ArrayRegion result; - for (int d = 0; d < N; ++d) { - result.start[d] = r.start[d] + offset[d]; - result.stop[d] = r.stop[d] + offset[d]; - result.step[d] = r.step[d]; - } - return result; - } - - /// Map a global-space region to local-space region. - ArrayRegion to_local(ArrayRegion const& r) const { - ArrayRegion result; - for (int d = 0; d < N; ++d) { - result.start[d] = r.start[d] - offset[d]; - result.stop[d] = r.stop[d] - offset[d]; - result.step[d] = r.step[d]; - } - return result; - } - - /// Number of chares along dimension d for this array. - int num_chares(int d) const { - int gstart = offset[d]; - int gstop = offset[d] + global_shape[d]; - return (gstop + tile - 1) / tile - gstart / tile; - } - - /// Chare index that owns global coordinate g along dimension d. - int owning_chare(int d, int global_coord) const { return global_coord / tile; } - - /// Global-space region owned by chare ci (clipped to this array's extent). - ArrayRegion chare_region_global(std::array const& ci) const { - ArrayRegion r; - for (int d = 0; d < N; ++d) { - r.start[d] = std::max(ci[d] * tile, offset[d]); - r.stop[d] = std::min((ci[d] + 1) * tile, offset[d] + global_shape[d]); - r.step[d] = 1; - } - return r; - } - - /// Local-space region owned by chare ci (clipped to this array's extent). - ArrayRegion chare_region_local(std::array const& ci) const { - return to_local(chare_region_global(ci)); - } - - /// Global start coordinate for chare ci along dimension d. - int chare_start_global(int d, int ci_d) const { - return std::max(ci_d * tile, offset[d]); - } - - /// Whether this is a default (offset=0) decomposition. - bool is_default() const { - for (int d = 0; d < N; ++d) - if (offset[d] != 0) - return false; - return true; - } -}; - -/// Format an N-dimensional ArrayRegion as "[s0:e0, s1:e1, ...]" -template -std::string fmt_region(ArrayRegion const& r) { - std::string s = "["; - for (int d = 0; d < N; ++d) { - if (d) - s += ", "; - s += std::to_string(r.start[d]) + ":" + std::to_string(r.stop[d]); - if (r.step[d] != 1) - s += ":" + std::to_string(r.step[d]); - } - s += "]"; - return s; -} - -/// N-D dynamic memref descriptor matching MLIR's LLVM ABI for memref. -/// Layout: allocated, aligned, offset, sizes[N], strides[N]. -template -struct MemRef { - T* allocated; - T* aligned; - int64_t offset; - int64_t sizes[N]; - int64_t strides[N]; - - static constexpr int fields_per_memref() { return 3 + 2 * N; } -}; - -/// A fragment of data: a pointer to the buffer and the region it covers. -/// src_strides: actual strides of the source buffer (0 = packed, compute -/// from region dimensions). Needed when the fragment points into a larger -/// 2D array whose row stride differs from the fragment's column count. -template -struct FragmentData { - ArrayRegion region; - T* data; - int64_t src_strides[N] = {}; -}; - -/// An aligned fragment: the MemRef descriptor ready for kernel dispatch, -/// plus the region in global coordinates (needed for output offset computation). -template -struct AlignedMemRef { - MemRef memref; - ArrayRegion region; -}; - -template -class ChareIndex { - public: - int idx[N]; - - bool operator==(ChareIndex const& other) const { - for (int i = 0; i < N; ++i) - if (idx[i] != other.idx[i]) - return false; - return true; - } -}; - -template -struct ChareIndexHash { - std::size_t operator()(ChareIndex const& ci) const { - std::size_t h = 0; - for (int i = 0; i < N; ++i) - h ^= std::hash()(ci.idx[i]) + 0x9e3779b9 + (h << 6) + (h >> 2); - return h; - } -}; - -/// Intersection of two regions in the same coordinate space. -/// Returns {result, true} if the intersection is non-empty, {_, false} otherwise. -template -std::pair, bool> intersect(ArrayRegion const& r1, ArrayRegion const& r2) { - ArrayRegion result; - return {result, r1.intersect(r2, result)}; -} - -/// Subtract r2 from r1: returns the fragments of r1 not covered by r2. -/// Produces up to 2*N axis-aligned slabs by peeling one dimension at a time. -template -std::vector> subtract(ArrayRegion const& r1, ArrayRegion const& r2) { - auto [overlap, has_overlap] = intersect(r1, r2); - if (!has_overlap) - return {r1}; - - std::vector> fragments; - - // Current remainder starts as r1; we narrow it per dimension as we peel - ArrayRegion remainder = r1; - - for (int d = 0; d < N; ++d) { - // For strided regions, overlap.stop may be an exclusive bound that is - // not itself on remainder's lattice (for example [3:65:2] intersect - // [0:64) => overlap [3:64:2]). When peeling the right slab, align the - // covered stop back onto remainder's lattice so we do not synthesize - // bogus fragments like [64:65:2] that contain no real source element. - int covered_stop = overlap.stop[d]; - if (remainder.step[d] > 1) - covered_stop = overlap.start[d] + overlap.size(d) * remainder.step[d]; - - // Left slab: remainder up to the overlap start in dimension d - if (remainder.start[d] < overlap.start[d]) { - ArrayRegion slab = remainder; - slab.stop[d] = overlap.start[d]; - fragments.push_back(slab); - } - - // Right slab: remainder from overlap stop in dimension d - if (covered_stop < remainder.stop[d]) { - ArrayRegion slab = remainder; - slab.start[d] = covered_stop; - fragments.push_back(slab); - } - - // Narrow remainder to the overlap range in this dimension - remainder.start[d] = overlap.start[d]; - remainder.stop[d] = covered_stop; - } - - return fragments; -} - -/// Given sub-region r1 (a subset of parent1), find the corresponding -/// sub-region in parent2's coordinate space. The i-th element of parent1 -/// corresponds to the i-th element of parent2. -template -ArrayRegion map(ArrayRegion const& r1, ArrayRegion const& parent1, - ArrayRegion const& parent2) { - ArrayRegion result; - - for (int d = 0; d < N; ++d) { - // Position of r1 within parent1 in logical coordinates. - // Use the sub-region's element count to compute the mapped stop so - // stepped regions preserve their full extent even when the exclusive - // stop is not aligned to parent1.step (for example [64:127:2]). - int lo = (r1.start[d] - parent1.start[d]) / parent1.step[d]; - int ls = r1.step[d] / parent1.step[d]; - int count = r1.size(d); - - // Translate to parent2's coordinate space - result.start[d] = parent2.start[d] + lo * parent2.step[d]; - result.step[d] = ls * parent2.step[d]; - result.stop[d] = result.start[d] + count * result.step[d]; - } - - return result; -} - -/// Result of local_inputs: the input sub-regions this chare needs, plus -/// how many remote messages to expect. -template -struct LocalInputs { - std::vector> my_inputs; - int expected_msgs; -}; - -/// For a given output region and list of input regions, compute: -/// - the sub-region of each input needed to produce this chare's output -/// - how many messages to expect from remote chares for data not owned locally -/// -/// All regions must be in global space. -/// output_decomp: decomposition of the output array -/// input_decomps: decomposition of each input array (one per input) -template -LocalInputs local_inputs(ArrayRegion const& r_out, std::vector> const& inputs, - std::array const& nd_idx, - ArrayDecomp const& output_decomp, - std::vector> const& input_decomps) { - LocalInputs result; - result.expected_msgs = 0; - - ArrayRegion r_chare_out = output_decomp.chare_region_global(nd_idx); - - DBG_PRINT(" local_inputs: r_out=%s, r_chare_out=%s, %d inputs\n", fmt_region(r_out).c_str(), - fmt_region(r_chare_out).c_str(), (int)inputs.size()); - - // Portion of the output that belongs to this chare - auto [r_myout, has_output] = intersect(r_out, r_chare_out); - if (!has_output) { - DBG_PRINT(" local_inputs: no output overlap -> returning empty\n"); - return result; - } - DBG_PRINT(" local_inputs: r_myout=%s\n", fmt_region(r_myout).c_str()); - - for (int idx = 0; idx < (int)inputs.size(); ++idx) { - auto const& r_inp = inputs[idx]; - ArrayRegion r_chare_inp = input_decomps[idx].chare_region_global(nd_idx); - DBG_PRINT(" local_inputs: input[%d] r_inp=%s r_chare_inp=%s\n", idx, - fmt_region(r_inp).c_str(), fmt_region(r_chare_inp).c_str()); - - // Map my output slice to the corresponding input slice - ArrayRegion r_myinp = map(r_myout, r_out, r_inp); - DBG_PRINT(" local_inputs: r_myinp (mapped)=%s\n", fmt_region(r_myinp).c_str()); - - auto [local_input, has_local] = intersect(r_myinp, r_chare_inp); - if (has_local) - DBG_PRINT(" local_inputs: local_input=%s\n", fmt_region(local_input).c_str()); - else - DBG_PRINT(" local_inputs: no local overlap\n"); - // Always push the full mapped input region so the fragment path - // knows the complete input range (local + remote portions). - result.my_inputs.push_back(r_myinp); - - // The part of my input that is NOT local — need remote messages - auto remote_fragments = subtract(r_myinp, r_chare_inp); - DBG_PRINT(" local_inputs: %d remote fragments\n", (int)remote_fragments.size()); - for (auto const& frag : remote_fragments) { - DBG_PRINT(" local_inputs: remote frag=%s\n", fmt_region(frag).c_str()); - auto chare_map = decompose(frag, input_decomps[idx]); - DBG_PRINT(" local_inputs: decomposed into %d chare(s)\n", (int)chare_map.size()); - result.expected_msgs += static_cast(chare_map.size()); - } - } - - DBG_PRINT(" local_inputs: total expected_msgs=%d, %d my_inputs\n", result.expected_msgs, - (int)result.my_inputs.size()); - return result; -} - -/// A message to send: the input sub-region destined for a remote chare. -template -struct RemoteSend { - ChareIndex target; - ArrayRegion region; - int input_index; // which input this send corresponds to -}; - -/// For each input, determine what local data this chare must send to remote -/// chares that need it for their portion of the output. -/// -/// All regions must be in global space. -/// For each input: -/// 1. intersect(r_inp, r_chare_inp) — the part of this input that I own -/// 2. map to output space — what output does my local input contribute to -/// 3. subtract r_chare_out — the output pieces that belong to remote chares -/// 4. decompose with output_decomp — which remote chares own each piece -/// 5. map back to input space — the actual input sub-region to send -template -std::vector> send_remote_inputs(ArrayRegion const& r_out, - std::vector> const& inputs, - std::array const& nd_idx, - ArrayDecomp const& output_decomp, - std::vector> const& input_decomps) { - std::vector> sends; - - ArrayRegion r_chare_out = output_decomp.chare_region_global(nd_idx); - - for (int inp_idx = 0; inp_idx < (int)inputs.size(); ++inp_idx) { - auto const& r_inp = inputs[inp_idx]; - ArrayRegion r_chare_inp = input_decomps[inp_idx].chare_region_global(nd_idx); - - // Part of this input that I own - auto [r_myinp, has_input] = intersect(r_inp, r_chare_inp); - if (!has_input) - continue; - - // What output does my local input contribute to - ArrayRegion r_myout = map(r_myinp, r_inp, r_out); - - // Output pieces that belong to remote chares (not me) - auto remote_out_frags = subtract(r_myout, r_chare_out); - for (auto const& frag : remote_out_frags) { - // Which remote chares own each piece - auto chare_map = decompose(frag, output_decomp); - for (auto const& [index, r_outsend] : chare_map) { - // Map the output piece back to input coordinates - ArrayRegion r_inpsend = map(r_outsend, r_out, r_inp); - sends.push_back({index, r_inpsend, inp_idx}); - } - } - } - - return sends; -} - -/// Common refinement of K inputs. Each input is a list of disjoint -/// region fragments (with data pointers) whose combined shape is the same -/// across all inputs. parents[k] is the parent region for input k (used to -/// convert to logical indices). Returns K outputs where every output has -/// the same M fragments, each with matching shape across all K outputs. -/// Each returned MemRef has the correct data pointer and strides for a -/// zero-copy view into the original fragment's buffer. -/// -/// Works by collecting all fragment boundaries (in logical index space) per -/// dimension across all inputs, then splitting every fragment at that grid. -template -std::vector>> -align_fragments(std::vector>> const& inputs, - std::vector> const& parents) { - if (inputs.empty()) - return {}; - int K = static_cast(inputs.size()); - - // 1. Collect all unique logical boundaries per dimension - std::array, N> grid; - - for (int k = 0; k < K; ++k) { - for (auto const& fd : inputs[k]) { - for (int d = 0; d < N; ++d) { - int gs = parents[k].start[d]; - int gst = parents[k].step[d]; - int lo = (fd.region.start[d] - gs) / gst; - assert(fd.region.step[d] % gst == 0 && - "align_fragments: fragment step must align with parent step"); - int logical_step = fd.region.step[d] / gst; - int hi = lo + fd.region.size(d) * logical_step; - grid[d].push_back(lo); - grid[d].push_back(hi); - } - } - } - - for (int d = 0; d < N; ++d) { - std::sort(grid[d].begin(), grid[d].end()); - grid[d].erase(std::unique(grid[d].begin(), grid[d].end()), grid[d].end()); - } - - // 2. For each input, split every fragment at the grid boundaries - std::vector>> result(K); - - for (int k = 0; k < K; ++k) { - for (auto const& fd : inputs[k]) { - auto const& frag = fd.region; - int gs[N], gst[N], frag_lo[N], frag_hi[N]; - for (int d = 0; d < N; ++d) { - gs[d] = parents[k].start[d]; - gst[d] = parents[k].step[d]; - frag_lo[d] = (frag.start[d] - gs[d]) / gst[d]; - assert(frag.step[d] % gst[d] == 0 && - "align_fragments: fragment step must align with parent step"); - int logical_step = frag.step[d] / gst[d]; - frag_hi[d] = frag_lo[d] + frag.size(d) * logical_step; - } - - // Use explicit source strides if provided, otherwise compute - // row-major packed strides from the fragment dimensions. - int64_t orig_strides[N]; - bool has_explicit = false; - for (int d = 0; d < N; ++d) - if (fd.src_strides[d] != 0) { - has_explicit = true; - break; - } - if (has_explicit) { - for (int d = 0; d < N; ++d) - orig_strides[d] = fd.src_strides[d]; - } else { - orig_strides[N - 1] = 1; - for (int d = N - 2; d >= 0; --d) - orig_strides[d] = orig_strides[d + 1] * frag.size(d + 1); - } - - // Find grid index range per dimension - std::array i_lo, i_hi; - for (int d = 0; d < N; ++d) { - i_lo[d] = static_cast( - std::lower_bound(grid[d].begin(), grid[d].end(), frag_lo[d]) - grid[d].begin()); - i_hi[d] = static_cast( - std::lower_bound(grid[d].begin(), grid[d].end(), frag_hi[d]) - grid[d].begin()); - } - - // Iterate over all grid cells within this fragment (odometer) - std::array idx; - for (int d = 0; d < N; ++d) - idx[d] = i_lo[d]; - - while (true) { - ArrayRegion cell; - for (int d = 0; d < N; ++d) { - cell.start[d] = gs[d] + grid[d][idx[d]] * gst[d]; - cell.stop[d] = gs[d] + grid[d][idx[d] + 1] * gst[d]; - cell.step[d] = frag.step[d]; - } - - // Compute data pointer offset into the original fragment - int64_t data_offset = 0; - for (int d = 0; d < N; ++d) { - int delta = cell.start[d] - frag.start[d]; - int logical_delta = delta / frag.step[d]; - data_offset += (has_explicit ? delta : logical_delta) * orig_strides[d]; - } - T* ptr = fd.data + data_offset; - - MemRef mr; - mr.allocated = ptr; - mr.aligned = ptr; - mr.offset = 0; - for (int d = 0; d < N; ++d) { - mr.sizes[d] = cell.size(d); - mr.strides[d] = has_explicit ? orig_strides[d] * cell.step[d] - : orig_strides[d]; - } - result[k].push_back({mr, cell}); - - // Advance odometer - int d = N - 1; - while (d >= 0) { - idx[d]++; - if (idx[d] < i_hi[d]) - break; - idx[d] = i_lo[d]; - --d; - } - if (d < 0) - break; - } - } - } - - // Sort each input's sub-fragments by logical cell index to ensure - // consistent ordering across all inputs. Different inputs may have - // different original fragments (local vs remote) that cover different - // regions, so the odometer traversal can produce sub-cells in different - // orders. Sorting by logical position makes the pairing deterministic. - for (int k = 0; k < K; ++k) { - std::sort(result[k].begin(), result[k].end(), - [&](AlignedMemRef const& a, AlignedMemRef const& b) { - for (int d = 0; d < N; ++d) { - int la = (a.region.start[d] - parents[k].start[d]) / parents[k].step[d]; - int lb = (b.region.start[d] - parents[k].start[d]) / parents[k].step[d]; - if (la != lb) - return la < lb; - } - return false; - }); - } - - return result; -} - -/// Decompose a global-space region into per-chare sub-regions. -/// Uses the ArrayDecomp's tile size for chare boundaries. -template -std::unordered_map, ArrayRegion, ChareIndexHash> -decompose(ArrayRegion const& region, ArrayDecomp const& decomp) { - std::unordered_map, ArrayRegion, ChareIndexHash> decomposition; - - // Compute chare index ranges per dimension - std::array chare_start, chare_stop; - for (int d = 0; d < N; ++d) { - chare_start[d] = region.start[d] / decomp.tile; - chare_stop[d] = (region.stop[d] + decomp.tile - 1) / decomp.tile; - } - - // Iterate over all chare indices in the N-dimensional range - ChareIndex ci{}; - for (int d = 0; d < N; ++d) - ci.idx[d] = chare_start[d]; - - while (true) { - // Intersect with the actual chare region so stepped fragments stay on - // the source lattice and offset decompositions are respected. Raw - // max/min clipping can synthesize bogus fragments at tile boundaries - // (for example [65:129:2] clipped to a [128:129) sliver). - std::array ci_arr{}; - for (int d = 0; d < N; ++d) - ci_arr[d] = ci.idx[d]; - ArrayRegion chare_reg = decomp.chare_region_global(ci_arr); - auto [subreg, valid] = intersect(region, chare_reg); - if (valid) - decomposition[ci] = subreg; - - // Advance the N-dimensional chare index (odometer-style) - int d = N - 1; - while (d >= 0) { - ci.idx[d]++; - if (ci.idx[d] < chare_stop[d]) - break; - ci.idx[d] = chare_start[d]; - --d; - } - if (d < 0) - break; - } - - return decomposition; -} diff --git a/src/ast.hpp b/src/ast.hpp new file mode 100644 index 0000000..1937ec6 --- /dev/null +++ b/src/ast.hpp @@ -0,0 +1,97 @@ +#include +#include +#include + + +template +inline T extract(char* &msg, bool increment=true) +{ + T arg = *(reinterpret_cast(msg)); + if (increment) + msg += sizeof(T); + return arg; +} + + +enum class operation : uint64_t +{ + noop = 0, + add = 1, + sub = 2, + mul = 3, + div = 4, + matmul = 5, + copy = 6, + axpy = 7, + axpy_multiplier = 8 +}; + + +class astnode +{ +public: + bool store; + bool is_scalar; + // FIXME double scalars fit into name, but should probably + // handle this better + uint64_t name; + operation oper; + std::vector operands; +}; + + +operation lookup_operation(uint64_t opcode) +{ + return static_cast(opcode); +} + + +astnode* decode(char* cmd) +{ + uint64_t opcode = extract(cmd); + astnode* node = new astnode; + node->oper = lookup_operation(opcode); + if (opcode == 0) + { + node->is_scalar = extract(cmd); + // if leaf is a scalar + if (node->is_scalar) + { + double value = extract(cmd); + memcpy(&(node->name), &value, sizeof(double)); + } + else + node->name = extract(cmd); + return node; + } + + node->is_scalar = false; + node->name = extract(cmd); + node->store = extract(cmd); + uint8_t num_operands = extract(cmd); + for (uint8_t i = 0; i < num_operands; i++) + { + uint32_t operand_size = extract(cmd); + astnode* opnode = decode(cmd); + node->operands.push_back(opnode); + cmd += operand_size; + } + + return node; +} + + +void delete_ast(astnode* node) +{ + if (node->oper == operation::noop) + { + delete node; + return; + } + + for (astnode* n : node->operands) + delete_ast(n); + + delete node; +} + diff --git a/src/backend.ci b/src/backend.ci deleted file mode 100644 index eb6def2..0000000 --- a/src/backend.ci +++ /dev/null @@ -1,85 +0,0 @@ -mainmodule charmnumeric -{ - extern module charmtyles; - - mainchare [migratable] Main - { - entry Main(CkArgMsg*); - }; - - group ArrayDAGGroup : DAGGroup - { - entry ArrayDAGGroup(); - entry void set_proxies(CkArrayID p1, CkArrayID p2, CkArrayID p3); - entry [reductiontarget] void proxies_ready(); - entry void receive_dag(int epoch, int size, char msg[size]); - entry void receive_get_request(int ndims, int epoch, int name, int size, int dtype); - entry void gather(int epoch, int name, int64_t offset, int64_t size, char data[size]); - }; - - array [1D] Partition1D - { - entry Partition1D(CProxy_ArrayDAGGroup dag_proxy); - entry void run(); - entry void receive_data(int node_id, int input_index, int name, - int ndims, int region_data[ndims*3], - int64_t size, -#ifdef USE_KOKKOS - nocopydevice -#endif - char data[size]); -#ifdef USE_KOKKOS - entry void receive_data(int node_id, int input_index, int name, - int ndims, int region_data[ndims*3], - int64_t size, char data[size], - CkDeviceBufferPost* devicePost); - entry void send_complete(int send_id); -#endif - entry void comm_done(int node_id); - entry void reduce_result(CkReductionMsg* msg); - }; - - array [2D] Partition2D - { - entry Partition2D(CProxy_ArrayDAGGroup dag_proxy); - entry void run(); - entry void receive_data(int node_id, int input_index, int name, - int ndims, int region_data[ndims*3], - int64_t size, -#ifdef USE_KOKKOS - nocopydevice -#endif - char data[size]); -#ifdef USE_KOKKOS - entry void receive_data(int node_id, int input_index, int name, - int ndims, int region_data[ndims*3], - int64_t size, char data[size], - CkDeviceBufferPost* devicePost); - entry void send_complete(int send_id); -#endif - entry void comm_done(int node_id); - entry void reduce_result(CkReductionMsg* msg); - }; - - array [3D] Partition3D - { - entry Partition3D(CProxy_ArrayDAGGroup dag_proxy); - entry void run(); - entry void receive_data(int node_id, int input_index, int name, - int ndims, int region_data[ndims*3], - int64_t size, -#ifdef USE_KOKKOS - nocopydevice -#endif - char data[size]); -#ifdef USE_KOKKOS - entry void receive_data(int node_id, int input_index, int name, - int ndims, int region_data[ndims*3], - int64_t size, char data[size], - CkDeviceBufferPost* devicePost); - entry void send_complete(int send_id); -#endif - entry void comm_done(int node_id); - entry void reduce_result(CkReductionMsg* msg); - }; -} diff --git a/src/backend.hpp b/src/backend.hpp deleted file mode 100644 index 4362a11..0000000 --- a/src/backend.hpp +++ /dev/null @@ -1,671 +0,0 @@ -#pragma once - -#include "array_region.hpp" -#include "opcodes.hpp" -#include -#include - -#include "charmnumeric.decl.h" - -#ifdef USE_KOKKOS -#define CT_MIN_TILE_1D 1048576 -#define CT_MIN_TILE_2D 1024 -#define CT_MIN_TILE_3D 128 -#else -#define CT_MIN_TILE_1D 262144 -#define CT_MIN_TILE_2D 512 -#define CT_MIN_TILE_3D 64 -#endif - - -// ---- Partition traits: map N -> concrete Charm++ chare/proxy/index types ---- - -template -struct PartitionTraits; - -template <> -struct PartitionTraits<1> { - using ProxyType = CProxy_Partition1D; - using CkIndexType = CkIndex_Partition1D; -}; -template <> -struct PartitionTraits<2> { - using ProxyType = CProxy_Partition2D; - using CkIndexType = CkIndex_Partition2D; -}; -template <> -struct PartitionTraits<3> { - using ProxyType = CProxy_Partition3D; - using CkIndexType = CkIndex_Partition3D; -}; - -/// Access an element of an N-D chare array proxy by ChareIndex. -template -inline auto proxy_at(typename PartitionTraits::ProxyType& proxy, const ChareIndex& ci); - -template <> -inline auto proxy_at<1>(CProxy_Partition1D& proxy, const ChareIndex<1>& ci) { - return proxy[ci.idx[0]]; -} -template <> -inline auto proxy_at<2>(CProxy_Partition2D& proxy, const ChareIndex<2>& ci) { - return proxy(ci.idx[0], ci.idx[1]); -} -template <> -inline auto proxy_at<3>(CProxy_Partition3D& proxy, const ChareIndex<3>& ci) { - return proxy(ci.idx[0], ci.idx[1], ci.idx[2]); -} - -#ifdef USE_KOKKOS -#include -using DeviceSpace = Kokkos::DefaultExecutionSpace::memory_space; -using HostSpace = Kokkos::HostSpace; -#ifdef USE_NVIDIA -#include "hapi.h" -#include -#endif -#endif - -class MLIRJitCompiler; // forward declaration — full definition only needed in backend.cpp - -/// Map C++ type to DType enum at compile time. -template -constexpr DType dtype_of(); -template <> -constexpr DType dtype_of() { - return DType::FLOAT32; -} -template <> -constexpr DType dtype_of() { - return DType::FLOAT64; -} -template <> -constexpr DType dtype_of() { - return DType::INT32; -} -template <> -constexpr DType dtype_of() { - return DType::INT64; -} - -/// Type-erased base class for N-D arrays. -template -class CTArrayBase { - public: - int name; - ArrayRegion region; - ArrayDecomp decomp; - std::array global_shape; - int global_size; // total elements across all dimensions - bool owner; - DType dtype; - - virtual ~CTArrayBase() = default; - int local_size() const { return region.size(); } - virtual void copyToHost() = 0; - virtual void copyToDevice() = 0; - virtual void* data_ptr() = 0; // host pointer - virtual void* device_data_ptr() = 0; // device pointer (or host if no GPU) - virtual int elem_size() const = 0; -}; - -template -class Array : public CTArrayBase { - public: -#ifdef USE_KOKKOS - Kokkos::View d_view; - typename Kokkos::View::HostMirror h_view; -#else - T* data; // host pointer (always allocated, used for gather/get) -#endif - - /// Construct from a local region + global shape + decomposition. - Array(ArrayRegion region_, std::array global_shape_, int name_, - ArrayDecomp decomp_) { - this->name = name_; - this->region = region_; - this->global_shape = global_shape_; - this->decomp = decomp_; - this->owner = true; - this->dtype = dtype_of(); - this->global_size = 1; - for (int d = 0; d < N; ++d) - this->global_size *= global_shape_[d]; - int n = this->region.size(); -#ifdef USE_KOKKOS - d_view = Kokkos::View("array_device", n); - h_view = Kokkos::create_mirror_view(HostSpace{}, d_view); - Kokkos::deep_copy(d_view, T(0)); - Kokkos::deep_copy(h_view, T(0)); -#else - data = new T[n]; - for (int i = 0; i < n; i++) - data[i] = T(0); -#endif - } - - /// Non-owning constructor (wraps existing buffer). - Array(T* data_, int name_, int size_) { - this->name = name_; - this->owner = false; - this->dtype = dtype_of(); - this->global_size = size_; -#ifdef USE_KOKKOS - h_view = Kokkos::View>(data_, size_); - // No device allocation for non-owning arrays -#else - data = data_; - for (int i = 0; i < size_; i++) - data[i] = T(0); -#endif - } - - ~Array() override { -#ifndef USE_KOKKOS - if (data != NULL && this->owner) - delete[] data; -#endif - // Kokkos Views are RAII — automatic cleanup - } - - void* data_ptr() override { -#ifdef USE_KOKKOS - return h_view.data(); -#else - return data; -#endif - } - - void* device_data_ptr() override { -#ifdef USE_KOKKOS - return d_view.data(); -#else - return data; // CPU-only: host pointer is the "device" pointer -#endif - } - - int elem_size() const override { return sizeof(T); } - - void copyToHost() override { -#ifdef USE_KOKKOS - Kokkos::deep_copy(h_view, d_view); -#endif - } - - void copyToDevice() override { -#ifdef USE_KOKKOS - Kokkos::deep_copy(d_view, h_view); -#endif - } -}; - -// ---- Communication Structures (templated on dimensionality only) ---- - -template -struct RemoteBuffer { - char* data; // host pointer (CPU) or device pointer (USE_KOKKOS) - ArrayRegion region; - int64_t byte_size; -}; - -template -struct PendingComm { - DAGNode* node; - int expected_msgs; - std::vector> my_inputs; - std::unordered_map>> remote_buffers; - bool incremental = false; // When true, compute as panels arrive - - void clear_remote_buffers() { - for (auto& [idx, bufs] : remote_buffers) - for (auto& rb : bufs) { -#ifdef USE_KOKKOS - Kokkos::kokkos_free(rb.data); -#else - delete[] rb.data; -#endif - } - remote_buffers.clear(); - } -}; - -// ---- Forward declarations ---- - -template -class PartitionImpl; - -// ---- ArrayDAGGroup ---- - -class ArrayDAGGroup : public CBase_ArrayDAGGroup { - public: - int num_compile; - - // Per-ndims partition state: current chare grid and proxy. - struct PartitionGrid { - int grid[3] = {0, 0, 0}; - int start_epoch = -1; // first epoch this partition was expanded - }; - std::unordered_map partition_grid; - CProxy_Partition1D partition_proxy_1; - CProxy_Partition2D partition_proxy_2; - CProxy_Partition3D partition_proxy_3; - std::unordered_map compile_cache; - std::unordered_map module_cache; - std::unordered_map gather_buffers; - std::unordered_map gather_counts; // bytes gathered so far - std::unordered_map gather_total; // total bytes expected - struct GatherFragment { - int64_t offset; - int64_t size; - char* data; // owned copy (bytes) - }; - std::unordered_map> gather_early; - - // Global metadata for arrays referenced by the current DAG. - // Maps array name -> metadata (ndims, global_shape, decomposition). - struct ArrayMetadata { - int ndims; - std::array global_shape; // padded to 3 dims - std::array offset; // decomposition offset (padded to 3 dims) - int tile; // tile size - bool decomp_final = false; // true once offset has been computed and must not change - - /// Construct an ArrayDecomp from the stored metadata. - template - ArrayDecomp decomp() const { - std::array shape{}, off{}; - for (int d = 0; d < N; ++d) { - shape[d] = global_shape[d]; - off[d] = offset[d]; - } - return ArrayDecomp::offset_decomp(shape, off, tile); - } - }; - std::unordered_map array_meta; - // Physical decomposition of arrays that have already been materialized. - // This persists across epochs so later DAGs do not reinterpret old data - // using freshly recomputed metadata unless the array is recreated. - std::unordered_map live_array_meta; - - MLIRJitCompiler* jit; - - ArrayDAGGroup(); - ArrayDAGGroup(CkMigrateMessage* m) {} - ~ArrayDAGGroup(); - - // Broadcast partition proxies from PE 0 to all PEs - void set_proxies(CkArrayID p1, CkArrayID p2, CkArrayID p3); - // Reduction target: all PEs have proxies, register CCS handlers - void proxies_ready(); - - // Unified DAG handling — dispatches to all relevant partition types - void receive_dag(int epoch, int size, char* serialized_dag); - void receive_get_request(int ndims, int epoch, int name, int size, int dtype = 0); - void gather(int epoch, int name, int64_t offset, int64_t size, char* data); - - // Templated execute_node — dispatched by dtype at call site - template - void execute_node_nd(DAGNode* node, PartitionImpl* partition, PendingComm* comm = nullptr); - - // Determine the decomposition for each array in the DAG. - // The runtime may later override entries with live_array_meta for arrays - // that already exist physically from an earlier epoch. - void compute_decompositions(DAG* dag); - - // Shared - void* compile_node(DAGNode* node); - void compile(DAG* dag); -}; - -// ---- Unified Executor (templated on N only) ---- - -template -class ArrayDAGExecutorND : public DAGExecutor { - private: - PartitionImpl* partition; - - public: - std::unordered_map> pending; - - ArrayDAGExecutorND(ArrayDAGGroup* group_, PartitionImpl* partition_) - : DAGExecutor(group_, N), partition(partition_) {} - virtual ~ArrayDAGExecutorND() = default; - - void execute_dag_node(DAGNode* node) override; - void delete_array(int name) override; - void execute_matmul_node(DAGNode* node); - void execute_matmatmul_node(DAGNode* node); - void execute_reduce_node(DAGNode* node); - void execute_cross_set_region_node(DAGNode* node); - void execute_diag_node(DAGNode* node); - void execute_tile_node(DAGNode* node); - void on_comm_done(int node_id); - void on_matmatmul_receive(int node_id, int input_index); - void on_matmul_partial(int node_id, RemoteBuffer& partial_buf); - void matmul_check_finalize(int node_id); - int ast_visitor(ASTNode* node, DType dtype); -}; - -// ---- PartitionImpl: all partition logic, independent of Charm++ chare base ---- - -template -class PartitionImpl { - friend class ArrayDAGExecutorND; - - public: - ArrayDAGExecutorND* executor; - typename PartitionTraits::ProxyType thisProxy; - CProxy_ArrayDAGGroup dag_proxy; - std::unordered_map*> arrays; // mixed-type array storage - std::array index; // N-D chare index (set by wrapper) - -#ifndef NDEBUG - int64_t comm_bytes_sent = 0; // cumulative communication volume (bytes) sent from this chare -#endif - - // Free list for array reuse: retired arrays keyed by exact buffer metadata. - // A buffer is only reusable when dtype, shape, and decomposition match. - static constexpr int FREE_LIST_MAX = 4; // max entries per key - struct FreeKey { - DType dtype; - std::array local_shape; - std::array global_shape; - std::array decomp_offset; - std::array decomp_global_shape; - int decomp_tile; - - bool operator==(const FreeKey& o) const { - return dtype == o.dtype && - local_shape == o.local_shape && - global_shape == o.global_shape && - decomp_offset == o.decomp_offset && - decomp_global_shape == o.decomp_global_shape && - decomp_tile == o.decomp_tile; - } - }; - struct FreeKeyHash { - static inline void hash_combine(std::size_t& seed, std::size_t value) { - seed ^= value + 0x9e3779b9 + (seed << 6) + (seed >> 2); - } - - static inline void hash_int_array(std::size_t& seed, const std::array& values) { - for (int value : values) - hash_combine(seed, std::hash()(value)); - } - - std::size_t operator()(const FreeKey& k) const { - std::size_t seed = std::hash()(static_cast(k.dtype)); - hash_int_array(seed, k.local_shape); - hash_int_array(seed, k.global_shape); - hash_int_array(seed, k.decomp_offset); - hash_int_array(seed, k.decomp_global_shape); - hash_combine(seed, std::hash()(k.decomp_tile)); - return seed; - } - }; - std::unordered_map*>, FreeKeyHash> free_arrays; - - static FreeKey make_free_key(DType dtype, - const ArrayRegion& region, - const std::array& global_shape, - const ArrayDecomp& decomp) { - std::array local_shape{}; - for (int d = 0; d < N; ++d) - local_shape[d] = region.size(d); - return FreeKey{ - dtype, - local_shape, - global_shape, - decomp.offset, - decomp.global_shape, - decomp.tile, - }; - } - - /// Move an array from `arrays` into the free list (or delete it if the - /// free list for its key is full). - void retire_array(int name) { - auto it = arrays.find(name); - if (it == arrays.end()) { - DBG_PRINT("[PE %d] Partition<%d> retire_array name=%d: no local buffer\n", - CkMyPe(), N, name); - return; - } - CTArrayBase* arr = it->second; - arrays.erase(it); - FreeKey key = make_free_key(arr->dtype, arr->region, arr->global_shape, arr->decomp); - auto& bucket = free_arrays[key]; - if (static_cast(bucket.size()) < FREE_LIST_MAX) { - bucket.push_back(arr); - DBG_PRINT("[PE %d] Partition<%d> retire_array name=%d: moved to freelist " - "(dtype=%d local_size=%d bucket=%d/%d)\n", - CkMyPe(), N, name, static_cast(arr->dtype), arr->local_size(), - static_cast(bucket.size()), FREE_LIST_MAX); - } else { - DBG_PRINT("[PE %d] Partition<%d> retire_array name=%d: deleting local buffer " - "(dtype=%d local_size=%d freelist_full=%d)\n", - CkMyPe(), N, name, static_cast(arr->dtype), arr->local_size(), - FREE_LIST_MAX); - delete arr; - } - } - - /// Try to pop a reusable buffer from the free list. - /// Returns nullptr if none available. - CTArrayBase* try_reuse(DType dtype, - const ArrayRegion& region, - const std::array& global_shape, - const ArrayDecomp& decomp) { - FreeKey key = make_free_key(dtype, region, global_shape, decomp); - auto it = free_arrays.find(key); - if (it == free_arrays.end() || it->second.empty()) - return nullptr; - CTArrayBase* arr = it->second.back(); - it->second.pop_back(); - if (it->second.empty()) - free_arrays.erase(it); - return arr; - } - - template - Array* allocate_or_reuse_typed(const ArrayRegion& region, - const std::array& global_shape, - int name, - const ArrayDecomp& decomp) { - CTArrayBase* reused = try_reuse(dtype_of(), region, global_shape, decomp); - if (reused == nullptr) - return new Array(region, global_shape, name, decomp); - - auto* arr = static_cast*>(reused); - arr->name = name; - arr->region = region; - arr->decomp = decomp; - arr->global_shape = global_shape; - arr->global_size = 1; - for (int d = 0; d < N; ++d) - arr->global_size *= global_shape[d]; - arr->owner = true; - arr->dtype = dtype_of(); - -#ifdef USE_KOKKOS - Kokkos::deep_copy(arr->d_view, T(0)); - Kokkos::deep_copy(arr->h_view, T(0)); -#else - T* data = static_cast(arr->data_ptr()); - for (int i = 0; i < arr->local_size(); ++i) - data[i] = T(0); -#endif - return arr; - } - - CTArrayBase* allocate_or_reuse(const ArrayRegion& region, - const std::array& global_shape, - int name, - DType dtype, - const ArrayDecomp& decomp) { - switch (dtype) { - case DType::FLOAT32: - return allocate_or_reuse_typed(region, global_shape, name, decomp); - case DType::FLOAT64: - return allocate_or_reuse_typed(region, global_shape, name, decomp); - case DType::INT32: - return allocate_or_reuse_typed(region, global_shape, name, decomp); - case DType::INT64: - return allocate_or_reuse_typed(region, global_shape, name, decomp); - } - return nullptr; - } - - template - Array* ensure_array_typed(const ArrayRegion& region, - const std::array& global_shape, - int name, - const ArrayDecomp& decomp) { - auto it = arrays.find(name); - if (it == arrays.end()) { - auto* arr = allocate_or_reuse_typed(region, global_shape, name, decomp); - arrays[name] = arr; - return arr; - } - return static_cast*>(it->second); - } - - CTArrayBase* ensure_array(const ArrayRegion& region, - const std::array& global_shape, - int name, - DType dtype, - const ArrayDecomp& decomp) { - auto it = arrays.find(name); - if (it == arrays.end()) { - auto* arr = allocate_or_reuse(region, global_shape, name, dtype, decomp); - arrays[name] = arr; - return arr; - } - return it->second; - } -#ifdef USE_KOKKOS - std::unordered_map pending_sends; // send_id → device ptr - int next_send_id = 0; -#ifdef USE_NVIDIA - cudaStream_t compute_stream_raw; - cudaStream_t comm_stream_raw; - Kokkos::Cuda compute_exec; - Kokkos::Cuda comm_exec; -#endif -#endif - - /// Return the N-D chare index (stored directly, no delinearization needed). - std::array nd_index() const { return index; } - - /// Callback to delegate contribute() to the owning chare. - std::function contribute_fn; - void contribute(int size, void* data, CkReduction::reducerType type, CkCallback cb) { - contribute_fn(size, data, type, cb); - } - - void init(typename PartitionTraits::ProxyType proxy, std::array idx, - CProxy_ArrayDAGGroup dag_proxy_); - ~PartitionImpl(); - - int create(ArrayRegion* region, int name, DType dtype, const ArrayDecomp& decomp); - void run(); - void receive_data(int node_id, int input_index, int name, int ndims, int* region_data, - int64_t size, char* data); -#ifdef USE_KOKKOS - void receive_data(int& node_id, int& input_index, int& name, int& ndims, int*& region_data, - int64_t& size, char*& data, CkDeviceBufferPost* devicePost); - void send_complete(int send_id); -#endif - void comm_done(int node_id); - void reduce_result(CkReductionMsg* msg); - void process_get(int epoch); -}; - -// ---- Thin wrapper chare classes ---- - -class Partition1D : public CBase_Partition1D { - public: - PartitionImpl<1> impl; - - Partition1D(CProxy_ArrayDAGGroup dag_proxy) { - impl.contribute_fn = [this](int sz, void* d, CkReduction::reducerType t, CkCallback cb) { - this->contribute(sz, d, t, cb); - }; - impl.init(thisProxy, {thisIndex}, dag_proxy); - } - Partition1D(CkMigrateMessage* m) {} - ~Partition1D() {} - - void run() { impl.run(); } - void receive_data(int node_id, int input_index, int name, int ndims, int* region_data, - int64_t size, char* data) { - impl.receive_data(node_id, input_index, name, ndims, region_data, size, data); - } -#ifdef USE_KOKKOS - void receive_data(int& node_id, int& input_index, int& name, int& ndims, int*& region_data, - int64_t& size, char*& data, CkDeviceBufferPost* devicePost) { - impl.receive_data(node_id, input_index, name, ndims, region_data, size, data, devicePost); - } - void send_complete(int send_id) { impl.send_complete(send_id); } -#endif - void comm_done(int node_id) { impl.comm_done(node_id); } - void reduce_result(CkReductionMsg* msg) { impl.reduce_result(msg); } -}; - -class Partition2D : public CBase_Partition2D { - public: - PartitionImpl<2> impl; - - Partition2D(CProxy_ArrayDAGGroup dag_proxy) { - impl.contribute_fn = [this](int sz, void* d, CkReduction::reducerType t, CkCallback cb) { - this->contribute(sz, d, t, cb); - }; - impl.init(thisProxy, {thisIndex.x, thisIndex.y}, dag_proxy); - } - Partition2D(CkMigrateMessage* m) {} - ~Partition2D() {} - - void run() { impl.run(); } - void receive_data(int node_id, int input_index, int name, int ndims, int* region_data, - int64_t size, char* data) { - impl.receive_data(node_id, input_index, name, ndims, region_data, size, data); - } -#ifdef USE_KOKKOS - void receive_data(int& node_id, int& input_index, int& name, int& ndims, int*& region_data, - int64_t& size, char*& data, CkDeviceBufferPost* devicePost) { - impl.receive_data(node_id, input_index, name, ndims, region_data, size, data, devicePost); - } - void send_complete(int send_id) { impl.send_complete(send_id); } -#endif - void comm_done(int node_id) { impl.comm_done(node_id); } - void reduce_result(CkReductionMsg* msg) { impl.reduce_result(msg); } -}; - -class Partition3D : public CBase_Partition3D { - public: - PartitionImpl<3> impl; - - Partition3D(CProxy_ArrayDAGGroup dag_proxy) { - impl.contribute_fn = [this](int sz, void* d, CkReduction::reducerType t, CkCallback cb) { - this->contribute(sz, d, t, cb); - }; - impl.init(thisProxy, {thisIndex.x, thisIndex.y, thisIndex.z}, dag_proxy); - } - Partition3D(CkMigrateMessage* m) {} - ~Partition3D() {} - - void run() { impl.run(); } - void receive_data(int node_id, int input_index, int name, int ndims, int* region_data, - int64_t size, char* data) { - impl.receive_data(node_id, input_index, name, ndims, region_data, size, data); - } -#ifdef USE_KOKKOS - void receive_data(int& node_id, int& input_index, int& name, int& ndims, int*& region_data, - int64_t& size, char*& data, CkDeviceBufferPost* devicePost) { - impl.receive_data(node_id, input_index, name, ndims, region_data, size, data, devicePost); - } - void send_complete(int send_id) { impl.send_complete(send_id); } -#endif - void comm_done(int node_id) { impl.comm_done(node_id); } - void reduce_result(CkReductionMsg* msg) { impl.reduce_result(msg); } -}; diff --git a/src/backend_internal.hpp b/src/backend_internal.hpp deleted file mode 100644 index 214c602..0000000 --- a/src/backend_internal.hpp +++ /dev/null @@ -1,490 +0,0 @@ -#pragma once - -#include "backend.hpp" -#include "dispatch.hpp" - -#ifdef USE_KOKKOS -#include -#else -#include -#endif - -/// Eigen-based gemv: res = mat[:, :actual_cols] * vec. -/// Matrix is row-major with `local_cols` stride. -#ifndef USE_KOKKOS -template -inline void eigen_gemv(T* mat_data, int local_rows, int local_cols, - T* vec_data, int actual_cols, T* res_data) { - Eigen::Map> - mat(mat_data, local_rows, local_cols); - Eigen::Map> vec(vec_data, actual_cols); - Eigen::Map> res(res_data, local_rows); - res.noalias() = mat.leftCols(actual_cols) * vec; -} - -/// Eigen-based sub-matrix gemv: res = mat[row_offset:row_offset+sub_rows, -/// col_offset:col_offset+sub_cols] * vec. -/// full_local_cols is the row stride of the full local matrix (row-major). -template -inline void eigen_gemv_sub(T* mat_data, int full_local_cols, - int row_offset, int col_offset, - int sub_rows, int sub_cols, - T* vec_data, T* res_data) { - typedef Eigen::Stride RowStride; - Eigen::Map, - 0, RowStride> - sub_mat(mat_data + row_offset * full_local_cols + col_offset, - sub_rows, sub_cols, RowStride(full_local_cols, 1)); - Eigen::Map> vec(vec_data, sub_cols); - Eigen::Map> res(res_data, sub_rows); - res.noalias() = sub_mat * vec; -} - -/// Eigen-based gemv on a 2D sub-matrix extracted from a 3D tile. -/// The 3D tile has local shape (local_d0, local_d1, local_d2) in row-major order. -/// dropped_dim identifies which dimension is the singleton (0, 1, or 2). -/// dd_local_offset is the local index in the dropped dimension. -/// row_offset/col_offset are local offsets in the remaining 2 dims. -template -inline void eigen_gemv_sub_3d(T* mat_data, - int local_d0, int local_d1, int local_d2, - int dropped_dim, int dd_local_offset, - int row_offset, int col_offset, - int sub_rows, int sub_cols, - T* vec_data, T* res_data) { - T* base; - int row_stride, col_stride; - - if (dropped_dim == 0) { - // Sub-matrix at (dd, row, col): contiguous 2D block - base = mat_data + dd_local_offset * local_d1 * local_d2 - + row_offset * local_d2 + col_offset; - row_stride = local_d2; - col_stride = 1; - } else if (dropped_dim == 1) { - // Sub-matrix at (row, dd, col) - base = mat_data + row_offset * local_d1 * local_d2 - + dd_local_offset * local_d2 + col_offset; - row_stride = local_d1 * local_d2; - col_stride = 1; - } else { - // Sub-matrix at (row, col, dd) - base = mat_data + row_offset * local_d1 * local_d2 - + col_offset * local_d2 + dd_local_offset; - row_stride = local_d1 * local_d2; - col_stride = local_d2; - } - - typedef Eigen::Stride GenStride; - Eigen::Map, - 0, GenStride> - sub_mat(base, sub_rows, sub_cols, GenStride(row_stride, col_stride)); - Eigen::Map> vec(vec_data, sub_cols); - Eigen::Map> res(res_data, sub_rows); - res.noalias() = sub_mat * vec; -} - -/// Eigen-based dot product: returns a.dot(b) for vectors of length n. -template -inline T eigen_dot(T* a_data, T* b_data, int n) { - Eigen::Map> a(a_data, n); - Eigen::Map> b(b_data, n); - return a.dot(b); -} - -/// Eigen-based gemm: C += A * B (accumulate). -/// A is (a_rows × a_cols), B is (a_cols × b_cols), C is (a_rows × b_cols). -/// All matrices are row-major, contiguous. -template -inline void eigen_gemm(T* a_data, int a_rows, int a_cols, - T* b_data, int b_cols, - T* c_data) { - using Mat = Eigen::Matrix; - Eigen::Map A(a_data, a_rows, a_cols); - Eigen::Map B(b_data, a_cols, b_cols); - Eigen::Map C(c_data, a_rows, b_cols); - C.noalias() += A * B; -} - -/// Eigen-based sub-matrix gemm: C += A_sub * B_sub (accumulate). -/// Extracts sub-blocks from row-major A and B by offset, accumulates into C. -/// A has row stride a_full_cols, B has row stride b_full_cols. -/// C is contiguous (sub_rows × sub_cols). -template -inline void eigen_gemm_sub(T* a_data, int a_full_cols, - int a_row_offset, int a_col_offset, - int sub_rows, int k_size, - T* b_data, int b_full_cols, - int b_row_offset, int b_col_offset, - int sub_cols, - T* c_data) { - typedef Eigen::Stride RowStride; - Eigen::Map, - 0, RowStride> - A_sub(a_data + a_row_offset * a_full_cols + a_col_offset, - sub_rows, k_size, RowStride(a_full_cols, 1)); - Eigen::Map, - 0, RowStride> - B_sub(b_data + b_row_offset * b_full_cols + b_col_offset, - k_size, sub_cols, RowStride(b_full_cols, 1)); - Eigen::Map> - C(c_data, sub_rows, sub_cols); - C.noalias() += A_sub * B_sub; -} -#endif - -/// Round up to the next power of 2 (returns v if already a power of 2). -inline int next_pow2(int v) { - if (v <= 1) - return 1; - v--; - v |= v >> 1; - v |= v >> 2; - v |= v >> 4; - v |= v >> 8; - v |= v >> 16; - return v + 1; -} - -/// Minimum tile size per dimension to keep tile volume >= ~1M elements. -/// Returns a power of 2. -inline int ct_min_tile(int ndims) { - switch (ndims) { - case 1: - return CT_MIN_TILE_1D; - case 2: - return CT_MIN_TILE_2D; - case 3: - return CT_MIN_TILE_3D; - default: - return 1; - } -} - -/// Get the tile size for array `name` from array_meta. -/// Falls back to ct_min_tile(ndims) if the array has no metadata entry or tile == 0. -/// Use this in all executor communication paths so that per-array tile sizes -/// assigned by compute_decompositions are respected. -inline int array_tile(const std::unordered_map& meta, - int name, int ndims) { - auto it = meta.find(name); - if (it != meta.end() && it->second.tile > 0) - return it->second.tile; - return ct_min_tile(ndims); -} - -/// Walk the AST to extract input and output regions for communication. -/// Returns true if region operations were found. -template -bool extract_regions_nd(DAGNode* node, std::unordered_map*>& arrays, - ArrayRegion& r_out, std::vector>& input_regions, - std::vector& input_source_names, std::array& global_shape_out); - -/// Determine the DType for a DAG node from its AST or the arrays map. -template -DType determine_dtype(DAGNode* node, std::unordered_map*>& arrays); - -/// Send matmul result from a 3D partition back to a 1D partition. -template -inline void cross_matmul_send_result_3d_to_1d(PartitionImpl<3>* partition, DAGNode* node, - int result_name, const ChareIndex<1>& target_ci, - T* result_data, int result_rows, int row_start, - CProxy_Partition1D& proxy_1d) { - int64_t byte_size = result_rows * sizeof(T); - - // Encode as 1D region: [row_start, row_start + result_rows) - int region_data[1 * 3]; - region_data[0] = row_start; - region_data[1] = row_start + result_rows; - region_data[2] = 1; - -#ifndef NDEBUG - partition->comm_bytes_sent += byte_size; -#endif - -#ifndef USE_KOKKOS - T* send_buf = new T[result_rows]; - memcpy(send_buf, result_data, byte_size); - proxy_at<1>(proxy_1d, target_ci) - .receive_data(node->id, /*input_index=*/0, result_name, 1, region_data, byte_size, - reinterpret_cast(send_buf)); - delete[] send_buf; -#else - T* send_buf = static_cast(Kokkos::kokkos_malloc(byte_size)); - Kokkos::parallel_for( - CT_COMM_POLICY(partition, result_rows), - KOKKOS_LAMBDA(int i) { send_buf[i] = result_data[i]; }); - device_pack_send<3, 1>(partition, proxy_1d, target_ci, node->id, 0, result_name, region_data, - byte_size, send_buf); -#endif -} - -/// Send matmul result from a 2D partition (column 0) back to a 1D partition. -template -inline void cross_matmul_send_result_2d_to_1d(PartitionImpl<2>* partition, DAGNode* node, - int result_name, const ChareIndex<1>& target_ci, - T* result_data, int result_rows, int row_start, - CProxy_Partition1D& proxy_1d) { - int64_t byte_size = result_rows * sizeof(T); - - // Encode as 1D region: [row_start, row_start + result_rows) - int region_data[1 * 3]; - region_data[0] = row_start; - region_data[1] = row_start + result_rows; - region_data[2] = 1; - -#ifndef NDEBUG - partition->comm_bytes_sent += byte_size; -#endif - -#ifndef USE_KOKKOS - T* send_buf = new T[result_rows]; - memcpy(send_buf, result_data, byte_size); - proxy_at<1>(proxy_1d, target_ci) - .receive_data(node->id, /*input_index=*/0, result_name, 1, region_data, byte_size, - reinterpret_cast(send_buf)); - delete[] send_buf; -#else - T* send_buf = static_cast(Kokkos::kokkos_malloc(byte_size)); - Kokkos::parallel_for( - CT_COMM_POLICY(partition, result_rows), - KOKKOS_LAMBDA(int i) { send_buf[i] = result_data[i]; }); - device_pack_send<2, 1>(partition, proxy_1d, target_ci, node->id, 0, result_name, region_data, - byte_size, send_buf); -#endif -} - - -/// --------------------------------------------------------------------------- -/// Custom Charm++ reduction for dot product (carries metadata + value) -/// --------------------------------------------------------------------------- - -/// Data contributed by each chare for a REDUCE (dot product) operation. -/// Layout: [node_id, result_name, dtype, pad, value[8]] = 24 bytes. -struct ReduceContrib { - int node_id; - int result_name; - int dtype_int; - int pad; - char value[8]; // large enough for float, double, int32_t, int64_t -}; -static_assert(sizeof(ReduceContrib) == 24, "ReduceContrib must be 24 bytes"); - -/// Custom reducer: sums the value field based on dtype, preserving metadata. -CkReductionMsg* reduce_dot_sum(int nMsg, CkReductionMsg** msgs); - -/// Global reducer type handle — registered once during init. -extern CkReduction::reducerType reduce_dot_sum_type; - -/// Call once to register the custom reducer. -void register_reduce_dot_sum(); - -/// --------------------------------------------------------------------------- -/// Cross-partition SET_REGION send helpers -/// --------------------------------------------------------------------------- - -/// Compute the dimension mapping between source (N_src) and target (N_tgt). -/// For higher→lower: finds non-singleton source dims. -/// For lower→higher: finds non-singleton target region dims. -/// dim_map[tgt_dim] = src_dim (maps each target dimension to a source dimension) -/// Singleton target dims get mapped to the corresponding singleton source dim. -template -inline void compute_dim_map(const int* src_global_shape, Region* tgt_region_base, int* dim_map) { - if constexpr (N_src > N_tgt) { - // Higher→Lower: non-singleton source dims map to target dims in order - int tgt_d = 0; - for (int sd = 0; sd < N_src && tgt_d < N_tgt; ++sd) { - if (src_global_shape[sd] > 1) - dim_map[tgt_d++] = sd; - } - // If all source dims are size 1, just map in order - if (tgt_d == 0) - for (int d = 0; d < N_tgt; ++d) - dim_map[d] = d; - } else { - // Lower→Higher: non-singleton target region dims receive source dims in order - auto* tgt_region = static_cast*>(tgt_region_base); - int src_d = 0; - for (int td = 0; td < N_tgt && src_d < N_src; ++td) { - if (tgt_region->size(td) > 1) - dim_map[td] = src_d++; - else - dim_map[td] = -1; // singleton — use target region's fixed value - } - } -} - -/// Send local source data from PartitionImpl to target chares on Partition. -/// The data is packed and sent with region encoded in N_tgt-dimensional target coordinates. -template -inline void cross_set_region_send(PartitionImpl* partition, DAGNode* node, int source_name, - const int* dim_map, Region* tgt_region_base, - typename PartitionTraits::ProxyType& target_proxy, - const ArrayDecomp& tgt_decomp) { - auto src_it = partition->arrays.find(source_name); - if (src_it == partition->arrays.end() || src_it->second->local_size() == 0) - return; - - auto* src = static_cast*>(src_it->second); - auto* tgt_region = static_cast*>(tgt_region_base); - - auto src_nd = partition->nd_index(); - auto& src_decomp = src->decomp; - auto src_chare_global = src_decomp.chare_region_global(src_nd); - - // Compute source local region in global coordinates - std::array src_cs; - std::array src_local_start, src_local_stop; - for (int d = 0; d < N_src; ++d) { - src_cs[d] = src_chare_global.start[d]; - src_local_start[d] = src_chare_global.start[d]; - src_local_stop[d] = src_chare_global.start[d] + src->region.size(d); - } - - // Map to target coordinates: build the target region this chare's data covers - std::array tgt_start, tgt_stop, tgt_step; - for (int td = 0; td < N_tgt; ++td) { - tgt_step[td] = 1; - if (dim_map[td] >= 0) { - int sd = dim_map[td]; - tgt_start[td] = src_local_start[sd]; - tgt_stop[td] = src_local_stop[sd]; - } else { - // Singleton target dim: use the target region's fixed value - tgt_start[td] = tgt_region->start[td]; - tgt_stop[td] = tgt_region->stop[td]; - } - } - ArrayRegion mapped_region(tgt_start, tgt_stop, tgt_step); - - // Decompose into target chares using the target array's decomposition - auto chare_map = decompose(mapped_region, tgt_decomp); - - // Compute source strides - std::array src_strides; - src_strides[N_src - 1] = 1; - for (int d = N_src - 2; d >= 0; --d) - src_strides[d] = src_strides[d + 1] * src->region.size(d + 1); - - for (auto& [ci, cr] : chare_map) { - // Compute the size of this fragment - int64_t total_size = cr.size(); - int64_t byte_size = total_size * sizeof(T); - - // Encode region in N_tgt coordinates - int region_data[N_tgt * 3]; - for (int d = 0; d < N_tgt; ++d) { - region_data[d * 3 + 0] = cr.start[d]; - region_data[d * 3 + 1] = cr.stop[d]; - region_data[d * 3 + 2] = 1; - } - -#ifndef NDEBUG - partition->comm_bytes_sent += byte_size; -#endif - -#ifndef USE_KOKKOS - src->copyToHost(); - T* send_buf = new T[total_size]; - T* src_data = static_cast(src->data_ptr()); - - // Pack data: iterate over the target fragment, map back to source coords - std::array idx; - for (int d = 0; d < N_tgt; ++d) - idx[d] = cr.start[d]; - for (int64_t flat = 0; flat < total_size; ++flat) { - int src_flat = 0; - for (int sd = 0; sd < N_src; ++sd) { - // Find which target dim maps to this source dim - int coord = 0; - for (int td = 0; td < N_tgt; ++td) { - if (dim_map[td] == sd) { - coord = idx[td] - src_cs[sd]; - break; - } - } - src_flat += coord * src_strides[sd]; - } - send_buf[flat] = src_data[src_flat]; - - // Advance odometer - for (int d = N_tgt - 1; d >= 0; --d) { - if (++idx[d] < cr.stop[d]) - break; - idx[d] = cr.start[d]; - } - } - - proxy_at(target_proxy, ci) - .receive_data(node->id, /*input_index=*/0, source_name, N_tgt, region_data, byte_size, - reinterpret_cast(send_buf)); - delete[] send_buf; -#else - T* send_buf = static_cast(Kokkos::kokkos_malloc(byte_size)); - T* src_device = static_cast(src->device_data_ptr()); - - // Capture for lambda - int dm[N_tgt], sc[N_src], ss[N_src], cr_start[N_tgt], cr_sizes[N_tgt]; - for (int d = 0; d < N_tgt; ++d) { - dm[d] = dim_map[d]; - cr_start[d] = cr.start[d]; - cr_sizes[d] = cr.size(d); - } - for (int d = 0; d < N_src; ++d) { - sc[d] = src_cs[d]; - ss[d] = src_strides[d]; - } - - Kokkos::parallel_for( - CT_COMM_POLICY(partition, total_size), KOKKOS_LAMBDA(int flat_idx) { - int remaining = flat_idx; - int tgt_coords[N_tgt]; - for (int d = N_tgt - 1; d >= 0; --d) { - tgt_coords[d] = cr_start[d] + remaining % cr_sizes[d]; - remaining /= cr_sizes[d]; - } - int src_flat = 0; - for (int sd = 0; sd < N_src; ++sd) { - int coord = 0; - for (int td = 0; td < N_tgt; ++td) { - if (dm[td] == sd) { - coord = tgt_coords[td] - sc[sd]; - break; - } - } - src_flat += coord * ss[sd]; - } - send_buf[flat_idx] = src_device[src_flat]; - }); - device_pack_send(partition, target_proxy, ci, node->id, 0, source_name, - region_data, byte_size, send_buf); -#endif - } -} - -/// Dispatch cross_set_region_send by DType. -template -inline void dispatch_cross_set_region_send(DType dt, PartitionImpl* partition, DAGNode* node, - int source_name, const int* dim_map, - Region* tgt_region_base, - typename PartitionTraits::ProxyType& target_proxy, - const ArrayDecomp& tgt_decomp) { - switch (dt) { - case DType::FLOAT32: - cross_set_region_send(partition, node, source_name, dim_map, - tgt_region_base, target_proxy, tgt_decomp); - break; - case DType::FLOAT64: - cross_set_region_send(partition, node, source_name, dim_map, - tgt_region_base, target_proxy, tgt_decomp); - break; - case DType::INT32: - cross_set_region_send(partition, node, source_name, dim_map, - tgt_region_base, target_proxy, tgt_decomp); - break; - case DType::INT64: - cross_set_region_send(partition, node, source_name, dim_map, - tgt_region_base, target_proxy, tgt_decomp); - break; - } -} diff --git a/src/dag_group_compile.cpp b/src/dag_group_compile.cpp deleted file mode 100644 index ae4807e..0000000 --- a/src/dag_group_compile.cpp +++ /dev/null @@ -1,56 +0,0 @@ -#include "backend.hpp" -#include "jit.hpp" - -#include - -void* ArrayDAGGroup::compile_node(DAGNode* node) { - MLIRJitCompiler node_jit; - node_jit.buildFromAST(node->ast); - node_jit.optimizeAndFuse(); - void* moduleHandle = nullptr; - void* funcPtr = nullptr; - -#if defined(USE_NVIDIA) - std::string ptx = node_jit.generateNVIDIA(); - if (ptx.empty()) - return nullptr; - if (!node_jit.loadNVIDIA(ptx, "fused_kernel", &moduleHandle, &funcPtr)) - return nullptr; -#elif defined(USE_AMD) - std::string gcn = node_jit.generateAMD(); - if (gcn.empty()) - return nullptr; - if (!node_jit.loadAMD(gcn, "fused_kernel", &moduleHandle, &funcPtr)) - return nullptr; -#elif defined(USE_INTEL) - std::string spirv = node_jit.generateIntel(); - if (spirv.empty()) - return nullptr; - if (!node_jit.loadIntel(spirv, "fused_kernel", &moduleHandle, &funcPtr)) - return nullptr; -#else - auto engine = node_jit.generateCPU(); - if (!engine) - return nullptr; - if (!node_jit.loadCPU(engine, "fused_kernel", &moduleHandle, &funcPtr)) - return nullptr; -#endif - - module_cache[node->identifier] = moduleHandle; - return funcPtr; -} - -void ArrayDAGGroup::compile(DAG* dag) { - int newly_compiled = 0; - for (auto& [id, node] : dag->nodes) { - auto it = compile_cache.find(node->identifier); - if (it == compile_cache.end() && node->fusible) { - num_compile++; - newly_compiled++; - void* compiled_fn = compile_node(node); - compile_cache[node->identifier] = compiled_fn; - } - } - DBG_PRINT("[PE %d] compile: %d new kernels, %d cached kernels total\n", - CkMyPe(), newly_compiled, (int)compile_cache.size()); -} diff --git a/src/dag_group_internal.hpp b/src/dag_group_internal.hpp deleted file mode 100644 index c18efc1..0000000 --- a/src/dag_group_internal.hpp +++ /dev/null @@ -1,39 +0,0 @@ -#pragma once - -#include "backend.hpp" - -#include -#include - -template -inline void dag_group_for_each_node_topo(DAG* dag, Fn&& fn) { - std::queue topo_queue; - std::unordered_map remaining_parents; - remaining_parents.reserve(dag->nodes.size()); - - for (auto& [id, node] : dag->nodes) { - remaining_parents[node] = node->num_parents; - if (node->num_parents == 0) - topo_queue.push(node); - } - - int processed = 0; - while (!topo_queue.empty()) { - DAGNode* dag_node = topo_queue.front(); - topo_queue.pop(); - processed++; - fn(dag_node); - - for (DAGNode* child : dag_node->children) { - auto it = remaining_parents.find(child); - if (it == remaining_parents.end()) - continue; - if (--it->second == 0) - topo_queue.push(child); - } - } - - if (processed != static_cast(dag->nodes.size())) - CkAbort("dag_group_for_each_node_topo processed %d/%d nodes", processed, - static_cast(dag->nodes.size())); -} diff --git a/src/dag_group_receive.cpp b/src/dag_group_receive.cpp deleted file mode 100644 index ad58707..0000000 --- a/src/dag_group_receive.cpp +++ /dev/null @@ -1,575 +0,0 @@ -#include "backend.hpp" -#include "dag_group_internal.hpp" - -#include -#include -#include -#include - -namespace { - -template -Region* deserialize_array_region_impl(char*& msg) { - std::array s, e, st; - for (int d = 0; d < N; ++d) { - s[d] = extract(msg); - e[d] = extract(msg); - st[d] = extract(msg); - } - return new ArrayRegion(s, e, st); -} - -Region* deserialize_array_region(int ndims, char*& msg) { - switch (ndims) { - case 1: - return deserialize_array_region_impl<1>(msg); - case 2: - return deserialize_array_region_impl<2>(msg); - case 3: - return deserialize_array_region_impl<3>(msg); - default: - CkAbort("Unsupported ndims=%d in deserialize_array_region", ndims); - return nullptr; - } -} - -bool region_to_shape(Region* region, int ndims, std::array& shape_out) { - shape_out = {0, 0, 0}; - if (region == nullptr || region->is_global) - return false; - - switch (ndims) { - case 1: { - auto* r = static_cast*>(region); - shape_out[0] = r->size(0); - return true; - } - case 2: { - auto* r = static_cast*>(region); - shape_out[0] = r->size(0); - shape_out[1] = r->size(1); - return true; - } - case 3: { - auto* r = static_cast*>(region); - shape_out[0] = r->size(0); - shape_out[1] = r->size(1); - shape_out[2] = r->size(2); - return true; - } - default: - return false; - } -} - -bool infer_result_shape_from_operands( - ASTNode* ast_node, const std::unordered_map& array_meta, - std::array& shape_out) { - shape_out = {0, 0, 0}; - bool found = false; - - for (int op_idx = 0; op_idx < (int)ast_node->operands.size(); ++op_idx) { - ASTNode* operand = ast_node->operands[op_idx]; - if (operand == nullptr || operand->is_scalar) - continue; - - std::array candidate = {0, 0, 0}; - bool have_candidate = false; - - Region* operand_region = ast_node->get_operand_region(op_idx); - if (operand_region && !operand_region->is_global) - have_candidate = region_to_shape(operand_region, operand->ndims, candidate); - - if (!have_candidate) { - auto meta_it = array_meta.find(operand->result_name); - if (meta_it != array_meta.end()) { - candidate = meta_it->second.global_shape; - have_candidate = true; - } - } - - if (!have_candidate) - continue; - - shape_out = candidate; - found = true; - - // Prefer a non-broadcast operand when available; size-1 broadcast - // operands are only a fallback when every array-like operand is scalar-shaped. - if (!operand->is_broadcast) - return true; - } - - return found; -} - -template -void insert_chare(typename PartitionTraits::ProxyType& part_proxy, CProxy_ArrayDAGGroup proxy, - const std::array& idx) { - ChareIndex ci; - for (int d = 0; d < N; ++d) - ci.idx[d] = idx[d]; - proxy_at(part_proxy, ci).insert(proxy); -} - -template -void expand_partition_nd(CProxy_ArrayDAGGroup proxy, typename PartitionTraits::ProxyType& part_proxy, - int* current_grid, Region* region_base, int tile) { - auto* region = static_cast*>(region_base); - int needed[N]; - - bool needs_expansion = false; - for (int d = 0; d < N; ++d) { - needed[d] = (region->size(d) + tile - 1) / tile; - if (needed[d] > current_grid[d]) - needs_expansion = true; - } - - if (!needs_expansion) - return; - - int expanded[N]; - for (int d = 0; d < N; ++d) - expanded[d] = std::max(current_grid[d], needed[d]); - - if (CkMyPe() == 0) { - part_proxy.beginInserting(); - - // Insert new chares by iterating dimension-by-dimension over the - // "L-shaped" expansion region. For each dimension d where the grid - // grew, iterate over the slab [old..expanded) in that dimension - // with [0..expanded) in dimensions > d and [0..old) in dimensions < d - // (the latter are already covered by earlier slabs). - int inserted = 0; - for (int dim = 0; dim < N; ++dim) { - if (expanded[dim] <= current_grid[dim]) - continue; - - int slab_total = 1; - int slab_extent[N]; - int slab_start[N]; - for (int d = 0; d < N; ++d) { - if (d < dim) { - slab_start[d] = 0; - slab_extent[d] = current_grid[d]; - } else if (d == dim) { - slab_start[d] = current_grid[d]; - slab_extent[d] = expanded[d] - current_grid[d]; - } else { - slab_start[d] = 0; - slab_extent[d] = expanded[d]; - } - slab_total *= slab_extent[d]; - } - - std::array idx; - for (int d = 0; d < N; ++d) - idx[d] = slab_start[d]; - - for (int i = 0; i < slab_total; ++i) { - insert_chare(part_proxy, proxy, idx); - inserted++; - - for (int d = N - 1; d >= 0; --d) { - if (++idx[d] < slab_start[d] + slab_extent[d]) - break; - idx[d] = slab_start[d]; - } - } - } - - part_proxy.doneInserting(); - DBG_PRINT("Partition<%d>: expanded grid, inserted %d chares\n", N, inserted); - } - - for (int d = 0; d < N; ++d) - current_grid[d] = expanded[d]; -} - -} // namespace - -void ArrayDAGGroup::receive_dag(int epoch, int size, char* serialized_dag) { - DAG* dag = DAG::deserialize(serialized_dag, deserialize_array_region); - - // Collect which ndims are used in this DAG - std::set active_ndims; - for (auto& [id, dag_node] : dag->nodes) { - for (ASTNode* ast_node : dag_node->ast->roots) { - if (ast_node->ndims >= 1 && ast_node->ndims <= 3) - active_ndims.insert(ast_node->ndims); - for (ASTNode* operand : ast_node->operands) - if (!operand->is_scalar && !operand->is_broadcast && operand->ndims >= 1 && - operand->ndims <= 3) - active_ndims.insert(operand->ndims); - if (static_cast(ast_node->opcode) == Opcode::DIAG) { - active_ndims.insert(1); - active_ndims.insert(2); - } - if (static_cast(ast_node->opcode) == Opcode::TILE) { - int input_nd = ast_node->operands[0]->ndims; - int result_nd = ast_node->ndims; - active_ndims.insert(input_nd); - active_ndims.insert(result_nd); - } - } - } - DBG_PRINT("[PE %d] receive_dag epoch=%d: active_ndims={", CkMyPe(), epoch); - for (int nd : active_ndims) - DBG_PRINT(" %d", nd); - DBG_PRINT(" }, num_nodes=%d\n", (int)dag->nodes.size()); - - // Pass 1: populate array_meta shapes in topological order. - dag_group_for_each_node_topo(dag, [&](DAGNode* dag_node) { - for (ASTNode* ast_node : dag_node->ast->roots) { - Opcode op = static_cast(ast_node->opcode); - - if (op == Opcode::CREATE && ast_node->region) { - int nd = ast_node->ndims; - switch (nd) { - case 1: { - auto* r = static_cast*>(ast_node->region); - array_meta[ast_node->result_name] = {nd, {r->size(0), 0, 0}, {}, 0}; - break; - } - case 2: { - auto* r = static_cast*>(ast_node->region); - array_meta[ast_node->result_name] = {nd, {r->size(0), r->size(1), 0}, {}, 0}; - break; - } - case 3: { - auto* r = static_cast*>(ast_node->region); - array_meta[ast_node->result_name] = - {nd, {r->size(0), r->size(1), r->size(2)}, {}, 0}; - break; - } - default: - CkAbort("Unsupported ndims=%d for partition creation", nd); - } - } else if (op == Opcode::REDUCE) { - array_meta[ast_node->result_name] = {1, {1, 0, 0}, {}, 0}; - } else if (op == Opcode::MATMUL) { - int mat_op_name = ast_node->operands[0]->result_name; - auto mat_it = array_meta.find(mat_op_name); - if (mat_it != array_meta.end()) { - int mat_ndims = mat_it->second.ndims; - int mat_rows = mat_it->second.global_shape[0]; - if (ast_node->operand_regions.size() >= 1) { - if (mat_ndims == 3) { - auto* mr = - static_cast*>(ast_node->operand_regions[0]); - int dropped = -1; - for (int d = 0; d < 3; d++) { - if (mr->stop[d] - mr->start[d] == 1) { - dropped = d; - break; - } - } - int row_dim = (dropped == 0) ? 1 : 0; - mat_rows = mr->stop[row_dim] - mr->start[row_dim]; - } else { - auto* mr = - static_cast*>(ast_node->operand_regions[0]); - mat_rows = mr->stop[0] - mr->start[0]; - } - } - array_meta[ast_node->result_name] = {1, {mat_rows, 0, 0}, {}, 0}; - } - } else if (op == Opcode::MATMATMUL) { - int lhs_name = ast_node->operands[0]->result_name; - int rhs_name = ast_node->operands[1]->result_name; - auto lhs_it = array_meta.find(lhs_name); - auto rhs_it = array_meta.find(rhs_name); - if (lhs_it != array_meta.end() && rhs_it != array_meta.end()) { - int M, N_cols; - if (ast_node->operand_regions.size() >= 2) { - int lhs_ndims = lhs_it->second.ndims; - int rhs_ndims = rhs_it->second.ndims; - if (lhs_ndims == 3) { - auto* lr = - static_cast*>(ast_node->operand_regions[0]); - int dd = -1; - for (int d = 0; d < 3; d++) - if (lr->stop[d] - lr->start[d] == 1) { - dd = d; - break; - } - int rd = (dd == 0) ? 1 : 0; - M = lr->stop[rd] - lr->start[rd]; - } else { - auto* lr = - static_cast*>(ast_node->operand_regions[0]); - M = lr->stop[0] - lr->start[0]; - } - if (rhs_ndims == 3) { - auto* rr = - static_cast*>(ast_node->operand_regions[1]); - int dd = -1; - for (int d = 0; d < 3; d++) - if (rr->stop[d] - rr->start[d] == 1) { - dd = d; - break; - } - int cd = (dd <= 1) ? 2 : 1; - N_cols = rr->stop[cd] - rr->start[cd]; - } else { - auto* rr = - static_cast*>(ast_node->operand_regions[1]); - N_cols = rr->stop[1] - rr->start[1]; - } - } else { - M = lhs_it->second.global_shape[0]; - N_cols = rhs_it->second.global_shape[1]; - } - array_meta[ast_node->result_name] = {2, {M, N_cols, 0}, {}, 0}; - } - } else if (op == Opcode::DIAG) { - int input_name = ast_node->operands[0]->result_name; - auto input_it = array_meta.find(input_name); - if (input_it != array_meta.end()) { - int input_ndims = input_it->second.ndims; - int k_offset = 0; - if (ast_node->operands.size() >= 2 && ast_node->operands[1]->is_scalar) - k_offset = (int)ast_node->operands[1]->scalar; - if (input_ndims == 1) { - int vec_len = input_it->second.global_shape[0]; - int n = vec_len + std::abs(k_offset); - array_meta[ast_node->result_name] = {2, {n, n, 0}, {}, 0}; - } else if (input_ndims == 2) { - int M = input_it->second.global_shape[0]; - int N_cols = input_it->second.global_shape[1]; - int diag_len; - if (k_offset >= 0) - diag_len = std::max(0, std::min(M, N_cols - k_offset)); - else - diag_len = std::max(0, std::min(M + k_offset, N_cols)); - array_meta[ast_node->result_name] = {1, {diag_len, 0, 0}, {}, 0}; - } - } - } else if (op == Opcode::TILE) { - int input_name = ast_node->operands[0]->result_name; - auto input_it = array_meta.find(input_name); - if (input_it != array_meta.end()) { - int input_ndims = input_it->second.ndims; - int out_ndims = ast_node->ndims; - std::array reps = {1, 1, 1}; - int num_reps = (int)ast_node->operands.size() - 1; - for (int d = 0; d < num_reps && d < 3; d++) { - if (ast_node->operands[d + 1]->is_scalar) - reps[d] = (int)ast_node->operands[d + 1]->scalar; - } - int delta = out_ndims - input_ndims; - std::array out_shape = {0, 0, 0}; - for (int d = 0; d < out_ndims; d++) { - int in_dim = - (d < delta) ? 1 : input_it->second.global_shape[d - delta]; - out_shape[d] = in_dim * reps[d]; - } - array_meta[ast_node->result_name] = {out_ndims, out_shape, {}, 0}; - } - } else if (op == Opcode::SET_REGION) { - // SET_REGION writes into the target array's partition. - } else if (is_elementwise(op) || op == Opcode::COPY) { - if (ast_node->is_temp) - continue; - std::array inferred_shape = {0, 0, 0}; - if (infer_result_shape_from_operands(ast_node, array_meta, inferred_shape)) { - array_meta[ast_node->result_name] = {ast_node->ndims, inferred_shape, {}, 0}; - } else { - array_meta[ast_node->result_name] = {1, {1, 0, 0}, {}, 0}; - } - } else if (op != Opcode::NOOP) { - bool found = false; - for (ASTNode* operand : ast_node->operands) { - if (operand->is_scalar || operand->is_broadcast) - continue; - auto it = array_meta.find(operand->result_name); - if (it != array_meta.end()) { - array_meta[ast_node->result_name] = it->second; - found = true; - break; - } - } - if (!found) - array_meta[ast_node->result_name] = {1, {1, 0, 0}, {}, 0}; - } - } - }); - DBG_PRINT("[PE %d] receive_dag epoch=%d: metadata pass complete\n", CkMyPe(), epoch); - - compute_decompositions(dag); - for (auto& [name, live_meta] : live_array_meta) { - auto meta_it = array_meta.find(name); - if (meta_it == array_meta.end()) { - array_meta[name] = live_meta; - continue; - } - meta_it->second.ndims = live_meta.ndims; - meta_it->second.global_shape = live_meta.global_shape; - meta_it->second.tile = live_meta.tile; - meta_it->second.offset = live_meta.offset; - } - - // Pass 2: expand partitions now that tile sizes have been assigned. - dag_group_for_each_node_topo(dag, [&](DAGNode* dag_node) { - for (ASTNode* ast_node : dag_node->ast->roots) { - Opcode op = static_cast(ast_node->opcode); - - if (op == Opcode::CREATE && ast_node->region) { - int nd = ast_node->ndims; - int tile = array_meta[ast_node->result_name].tile; - if (partition_grid[nd].start_epoch < 0) - partition_grid[nd].start_epoch = epoch; - switch (nd) { - case 1: - expand_partition_nd<1>(thisProxy, partition_proxy_1, partition_grid[nd].grid, - ast_node->region, tile); - break; - case 2: - expand_partition_nd<2>(thisProxy, partition_proxy_2, partition_grid[nd].grid, - ast_node->region, tile); - break; - case 3: - expand_partition_nd<3>(thisProxy, partition_proxy_3, partition_grid[nd].grid, - ast_node->region, tile); - break; - default: - CkAbort("Unsupported ndims=%d in expansion pass", nd); - } - } else if (op == Opcode::MATMUL) { - auto res_it = array_meta.find(ast_node->result_name); - if (res_it != array_meta.end()) { - int mat_rows = res_it->second.global_shape[0]; - int tile_1d = res_it->second.tile; - std::array s = {0}, e = {mat_rows}, st = {1}; - ArrayRegion<1> result_region(s, e, st); - expand_partition_nd<1>(thisProxy, partition_proxy_1, partition_grid[1].grid, - &result_region, tile_1d); - } - } else if (op == Opcode::MATMATMUL) { - auto res_it = array_meta.find(ast_node->result_name); - if (res_it != array_meta.end()) { - int M = res_it->second.global_shape[0]; - int N_cols = res_it->second.global_shape[1]; - int tile_2d = res_it->second.tile; - std::array s = {0, 0}, e = {M, N_cols}, st = {1, 1}; - ArrayRegion<2> result_region(s, e, st); - expand_partition_nd<2>(thisProxy, partition_proxy_2, partition_grid[2].grid, - &result_region, tile_2d); - } - } else if (op == Opcode::DIAG) { - auto res_it = array_meta.find(ast_node->result_name); - if (res_it != array_meta.end()) { - int res_nd = res_it->second.ndims; - int tile = res_it->second.tile; - if (res_nd == 2) { - int n = res_it->second.global_shape[0]; - std::array s = {0, 0}, e = {n, n}, st = {1, 1}; - ArrayRegion<2> result_region(s, e, st); - if (partition_grid[2].start_epoch < 0) - partition_grid[2].start_epoch = epoch; - expand_partition_nd<2>(thisProxy, partition_proxy_2, partition_grid[2].grid, - &result_region, tile); - } else { - int diag_len = res_it->second.global_shape[0]; - std::array s = {0}, e = {diag_len}, st = {1}; - ArrayRegion<1> result_region(s, e, st); - if (partition_grid[1].start_epoch < 0) - partition_grid[1].start_epoch = epoch; - expand_partition_nd<1>(thisProxy, partition_proxy_1, partition_grid[1].grid, - &result_region, tile); - } - } - } else if (op == Opcode::TILE) { - auto res_it = array_meta.find(ast_node->result_name); - if (res_it != array_meta.end()) { - int res_nd = res_it->second.ndims; - int tile = res_it->second.tile; - auto& sh = res_it->second.global_shape; - if (partition_grid[res_nd].start_epoch < 0) - partition_grid[res_nd].start_epoch = epoch; - switch (res_nd) { - case 1: { - std::array s = {0}, e = {sh[0]}, st = {1}; - ArrayRegion<1> result_region(s, e, st); - expand_partition_nd<1>(thisProxy, partition_proxy_1, - partition_grid[res_nd].grid, &result_region, tile); - break; - } - case 2: { - std::array s = {0, 0}, e = {sh[0], sh[1]}, st = {1, 1}; - ArrayRegion<2> result_region(s, e, st); - expand_partition_nd<2>(thisProxy, partition_proxy_2, - partition_grid[res_nd].grid, &result_region, tile); - break; - } - case 3: { - std::array s = {0, 0, 0}, e = {sh[0], sh[1], sh[2]}, st = {1, 1, 1}; - ArrayRegion<3> result_region(s, e, st); - expand_partition_nd<3>(thisProxy, partition_proxy_3, - partition_grid[res_nd].grid, &result_region, tile); - break; - } - } - } - } else if (is_elementwise(op) || op == Opcode::COPY) { - auto res_it = array_meta.find(ast_node->result_name); - if (res_it != array_meta.end() && res_it->second.global_shape[0] > 0) { - int nd = res_it->second.ndims; - int tile = res_it->second.tile; - auto& sh = res_it->second.global_shape; - if (partition_grid[nd].start_epoch < 0) - partition_grid[nd].start_epoch = epoch; - switch (nd) { - case 1: { - std::array s = {0}, e = {sh[0]}, st = {1}; - ArrayRegion<1> region(s, e, st); - expand_partition_nd<1>(thisProxy, partition_proxy_1, partition_grid[nd].grid, - ®ion, tile); - break; - } - case 2: { - std::array s = {0, 0}, e = {sh[0], sh[1]}, st = {1, 1}; - ArrayRegion<2> region(s, e, st); - expand_partition_nd<2>(thisProxy, partition_proxy_2, partition_grid[nd].grid, - ®ion, tile); - break; - } - case 3: { - std::array s = {0, 0, 0}, e = {sh[0], sh[1], sh[2]}, st = {1, 1, 1}; - ArrayRegion<3> region(s, e, st); - expand_partition_nd<3>(thisProxy, partition_proxy_3, partition_grid[nd].grid, - ®ion, tile); - break; - } - } - } - } - } - }); - DBG_PRINT("[PE %d] receive_dag epoch=%d: partition expansion pass complete\n", CkMyPe(), - epoch); - - compile(dag); - - bool first = true; - for (int nd : active_ndims) { - if (first) { - add_dag(nd, epoch, dag); - first = false; - } else { - add_dag(nd, epoch, dag->copy()); - } - } - - for (auto& [nd, pg] : partition_grid) { - if (pg.grid[0] > 0 && active_ndims.find(nd) == active_ndims.end()) { - add_dag(nd, epoch, new DAG()); - DBG_PRINT("[PE %d] receive_dag epoch=%d: stored empty DAG for ndims=%d\n", CkMyPe(), - epoch, nd); - } - } - - partition_proxy_1.run(); - partition_proxy_2.run(); - partition_proxy_3.run(); -} diff --git a/src/dag_group_runtime.cpp b/src/dag_group_runtime.cpp deleted file mode 100644 index a728e69..0000000 --- a/src/dag_group_runtime.cpp +++ /dev/null @@ -1,118 +0,0 @@ -#include "backend_internal.hpp" -#include "decomposition_solver.hpp" -#include "jit.hpp" - -#include - -ArrayDAGGroup::ArrayDAGGroup() : num_compile(0) { -#ifdef USE_KOKKOS - if (!Kokkos::is_initialized()) - Kokkos::initialize(); -#endif - register_reduce_dot_sum(); - jit = new MLIRJitCompiler(); - if (CkMyPe() == 0) { - partition_proxy_1 = CProxy_Partition1D::ckNew(); - partition_proxy_2 = CProxy_Partition2D::ckNew(); - partition_proxy_3 = CProxy_Partition3D::ckNew(); - thisProxy.set_proxies(partition_proxy_1, partition_proxy_2, partition_proxy_3); - } -} - -void ArrayDAGGroup::set_proxies(CkArrayID p1, CkArrayID p2, CkArrayID p3) { - partition_proxy_1 = CProxy_Partition1D(p1); - partition_proxy_2 = CProxy_Partition2D(p2); - partition_proxy_3 = CProxy_Partition3D(p3); - CkCallback cb(CkReductionTarget(ArrayDAGGroup, proxies_ready), thisProxy[0]); - contribute(cb); -} - -void ArrayDAGGroup::proxies_ready() { - Server::initialize(thisProxy); -} - -ArrayDAGGroup::~ArrayDAGGroup() { - delete jit; -#ifdef USE_KOKKOS - if (Kokkos::is_initialized()) - Kokkos::finalize(); -#endif -} - -void ArrayDAGGroup::compute_decompositions(DAG* dag) { - decomposition_solver::compute_decompositions(array_meta, dag, odf, CkNumPes()); - - DBG_PRINT("[PE %d] compute_decompositions: assigned tiles and offsets for %d arrays\n", - CkMyPe(), (int)array_meta.size()); - for (auto& [name, meta] : array_meta) { - DBG_PRINT(" array %d: ndims=%d shape=(%d,%d,%d) tile=%d offset=(%d,%d,%d)\n", name, - meta.ndims, meta.global_shape[0], meta.global_shape[1], - meta.global_shape[2], meta.tile, meta.offset[0], meta.offset[1], - meta.offset[2]); - } -} - -void ArrayDAGGroup::receive_get_request(int ndims, int epoch, int name, int size, int dtype) { - DAGGroup::receive_get_request(ndims, epoch, name, size, dtype); - - // Store empty DAGs for ndims that have active partitions but aren't involved in this GET. - for (auto& [nd, pg] : partition_grid) { - if (pg.grid[0] > 0 && nd != ndims) { - add_dag(nd, epoch, new DAG()); - } - } - - // Wake up all partition chares so they can check for new work. - partition_proxy_1.run(); - partition_proxy_2.run(); - partition_proxy_3.run(); - - // Only PE 0 allocates the gather buffer for assembly - // size is element count; convert to bytes using the wire dtype - int elem_size = dtype_size(static_cast(dtype)); - int64_t byte_size = (int64_t)size * elem_size; - if (CkMyPe() == 0) { - gather_buffers[epoch] = new char[byte_size]; - gather_total[epoch] = byte_size; - if (gather_counts.find(epoch) == gather_counts.end()) - gather_counts[epoch] = 0; - - // Flush any gather fragments that arrived before this allocation - auto early_it = gather_early.find(epoch); - if (early_it != gather_early.end()) { - for (auto& frag : early_it->second) { - memcpy(gather_buffers[epoch] + frag.offset, frag.data, frag.size); - delete[] frag.data; - } - gather_early.erase(early_it); - if (gather_counts[epoch] >= gather_total[epoch]) { - Server::send_reply(epoch, byte_size, gather_buffers[epoch]); - gather_counts.erase(epoch); - gather_total.erase(epoch); - delete[] gather_buffers[epoch]; - gather_buffers.erase(epoch); - } - } - } -} - -void ArrayDAGGroup::gather(int epoch, int name, int64_t offset, int64_t size, char* data) { - // offset and size are in bytes - if (gather_buffers.find(epoch) == gather_buffers.end()) { - char* buf = new char[size]; - memcpy(buf, data, size); - gather_early[epoch].push_back({offset, size, buf}); - gather_counts[epoch] += size; - return; - } - - memcpy(gather_buffers[epoch] + offset, data, size); - gather_counts[epoch] += size; - if (gather_counts[epoch] >= gather_total[epoch]) { - Server::send_reply(epoch, gather_total[epoch], gather_buffers[epoch]); - gather_counts.erase(epoch); - gather_total.erase(epoch); - delete[] gather_buffers[epoch]; - gather_buffers.erase(epoch); - } -} diff --git a/src/decomposition_solver.hpp b/src/decomposition_solver.hpp deleted file mode 100644 index dd98aa5..0000000 --- a/src/decomposition_solver.hpp +++ /dev/null @@ -1,1716 +0,0 @@ -#pragma once - -#include "array_region.hpp" -#include "opcodes.hpp" - -#include -#include -#include -#include -#include -#include -#include -#include -#include -#include - -#ifndef CT_MIN_TILE_1D -#ifdef USE_KOKKOS -#define CT_MIN_TILE_1D 1048576 -#define CT_MIN_TILE_2D 1024 -#define CT_MIN_TILE_3D 128 -#else -#define CT_MIN_TILE_1D 262144 -#define CT_MIN_TILE_2D 512 -#define CT_MIN_TILE_3D 64 -#endif -#endif - -namespace decomposition_solver { - -namespace detail { - -inline int next_pow2(int v) { - if (v <= 1) - return 1; - v--; - v |= v >> 1; - v |= v >> 2; - v |= v >> 4; - v |= v >> 8; - v |= v >> 16; - return v + 1; -} - -inline int ct_min_tile(int ndims) { - switch (ndims) { - case 1: - return CT_MIN_TILE_1D; - case 2: - return CT_MIN_TILE_2D; - case 3: - return CT_MIN_TILE_3D; - default: - return 1; - } -} - -template -inline void for_each_node_topo(DAG* dag, Fn&& fn) { - if (dag == nullptr) - return; - - std::queue topo_queue; - std::unordered_map remaining_parents; - remaining_parents.reserve(dag->nodes.size()); - - for (auto& [id, node] : dag->nodes) { - remaining_parents[node] = node->num_parents; - if (node->num_parents == 0) - topo_queue.push(node); - } - - int processed = 0; - while (!topo_queue.empty()) { - DAGNode* dag_node = topo_queue.front(); - topo_queue.pop(); - processed++; - fn(dag_node); - - for (DAGNode* child : dag_node->children) { - auto it = remaining_parents.find(child); - if (it == remaining_parents.end()) - continue; - if (--it->second == 0) - topo_queue.push(child); - } - } - - assert(processed == static_cast(dag->nodes.size())); -} - -} // namespace detail - -template -void compute_decompositions(ArrayMetaMap& array_meta, DAG* dag, int odf, int num_pes) { - using ArrayMetadata = typename ArrayMetaMap::mapped_type; - - auto zero_offset = std::array{0, 0, 0}; - - auto compute_tile_count = [](const ArrayMetadata& m) -> int { - int count = 1; - for (int d = 0; d < m.ndims; d++) - count *= (m.global_shape[d] + m.tile - 1) / m.tile; - return count; - }; - - auto clamp_tile_for_max_count = [&](ArrayMetadata& m) { - int max_tiles = std::max(1, odf) * std::max(1, num_pes); - while (compute_tile_count(m) > max_tiles) - m.tile *= 2; - }; - - for (auto& [name, meta] : array_meta) { - if (meta.tile == 0) - meta.tile = detail::ct_min_tile(meta.ndims); - - if (!meta.decomp_final) { - meta.tile = std::max(detail::next_pow2(meta.tile), detail::ct_min_tile(meta.ndims)); - clamp_tile_for_max_count(meta); - meta.offset = zero_offset; - } - } - - if (dag != nullptr) { - std::unordered_set matmul_arrays; - - auto positive_mod = [](int value, int mod) { - int r = value % mod; - return r < 0 ? r + mod : r; - }; - - auto gcd_positive = [](int a, int b) { - if (a < 0) - a = -a; - if (b < 0) - b = -b; - while (b != 0) { - int t = a % b; - a = b; - b = t; - } - return a; - }; - - auto region_stride_factor = [&](Region* region_base, int ndims) { - if (region_base == nullptr || region_base->is_global) - return 1; - - int factor = 0; - auto fold_step = [&](int step) { - if (step < 0) - step = -step; - if (step <= 1) - return; - factor = (factor == 0) ? step : gcd_positive(factor, step); - }; - - switch (ndims) { - case 1: { - auto* r = static_cast*>(region_base); - fold_step(r->step[0]); - break; - } - case 2: { - auto* r = static_cast*>(region_base); - fold_step(r->step[0]); - fold_step(r->step[1]); - break; - } - case 3: { - auto* r = static_cast*>(region_base); - fold_step(r->step[0]); - fold_step(r->step[1]); - fold_step(r->step[2]); - break; - } - default: - break; - } - - return factor > 1 ? factor : 1; - }; - - auto region_is_dense = [](Region* region_base, int ndims) { - if (region_base == nullptr || region_base->is_global) - return true; - - switch (ndims) { - case 1: { - auto* r = static_cast*>(region_base); - return r->step[0] == 1; - } - case 2: { - auto* r = static_cast*>(region_base); - return r->step[0] == 1 && r->step[1] == 1; - } - case 3: { - auto* r = static_cast*>(region_base); - return r->step[0] == 1 && r->step[1] == 1 && r->step[2] == 1; - } - default: - return false; - } - }; - - auto tile_fits_grid = [](const ArrayMetadata& meta, int tile) { - (void)meta; - return tile > 0; - }; - - auto try_assign_tile = [&](int array_name, int candidate_tile) { - auto array_it = array_meta.find(array_name); - if (array_it == array_meta.end()) - return false; - - auto& meta = array_it->second; - if (meta.decomp_final) - return false; - if (candidate_tile <= 0) - return false; - - candidate_tile = - std::max(detail::next_pow2(candidate_tile), detail::ct_min_tile(meta.ndims)); - - if (meta.tile > 0 && candidate_tile >= meta.tile) - return false; - if (!tile_fits_grid(meta, candidate_tile)) - return false; - - meta.tile = candidate_tile; - return true; - }; - - auto propose_tile_from_access = [&](int array_name, int ndims, int source_name, - Region* access_region) { - if (source_name == array_name) - return false; - - auto src_it = array_meta.find(source_name); - if (src_it == array_meta.end() || src_it->second.ndims != ndims) - return false; - - int candidate_tile = src_it->second.tile; - int stride_factor = region_stride_factor(access_region, ndims); - if (stride_factor > 1) - candidate_tile = std::max(1, candidate_tile / stride_factor); - - return try_assign_tile(array_name, candidate_tile); - }; - - auto collect_tile_candidates = - [&](auto&& self, int array_name, int ndims, ASTNode* node, Region* carried_region, - const std::unordered_map& temp_defs) -> bool { - if (node == nullptr || node->is_scalar || node->is_broadcast) - return false; - - auto meta_it = array_meta.find(node->result_name); - if (meta_it != array_meta.end()) { - if (meta_it->second.ndims == ndims) - return propose_tile_from_access(array_name, ndims, node->result_name, - carried_region); - return false; - } - - if (node->operands.empty()) { - auto def_it = temp_defs.find(node->result_name); - if (def_it != temp_defs.end()) { - ASTNode* def_root = def_it->second; - bool changed = false; - for (int i = 0; i < (int)def_root->operands.size(); ++i) { - changed = self(self, array_name, ndims, def_root->operands[i], - def_root->get_operand_region(i), temp_defs) || - changed; - } - return changed; - } - return false; - } - - bool changed = false; - for (int child_idx = 0; child_idx < (int)node->operands.size(); ++child_idx) { - changed = self(self, array_name, ndims, node->operands[child_idx], - node->get_operand_region(child_idx), temp_defs) || - changed; - } - return changed; - }; - - auto is_ast_leaf = [](ASTNode* node) -> bool { - for (ASTNode* child : node->operands) - if (child != nullptr && !child->is_scalar && !child->is_broadcast) - return false; - return true; - }; - - bool tiles_changed = true; - int max_tile_passes = std::max(64, static_cast(array_meta.size())); - for (int iter = 0; iter < max_tile_passes && tiles_changed; ++iter) { - tiles_changed = false; - - detail::for_each_node_topo(dag, [&](DAGNode* dag_node) { - std::unordered_map temp_defs; - for (ASTNode* root : dag_node->ast->roots) { - if (root->is_temp) - temp_defs[root->result_name] = root; - } - - for (ASTNode* ast_node : dag_node->ast->roots) { - Opcode op = static_cast(ast_node->opcode); - - if (op == Opcode::SET_REGION) { - if (ast_node->operands.size() < 2) - continue; - - int dst_name = ast_node->operands[0]->result_name; - auto dst_it = array_meta.find(dst_name); - if (dst_it == array_meta.end()) - continue; - - if (region_is_dense(ast_node->region, dst_it->second.ndims)) { - for (int op_idx = 1; op_idx < (int)ast_node->operands.size(); - ++op_idx) { - tiles_changed = - collect_tile_candidates(collect_tile_candidates, dst_name, - dst_it->second.ndims, - ast_node->operands[op_idx], - ast_node->get_operand_region(op_idx), - temp_defs) || - tiles_changed; - } - } - - for (int op_idx = 1; op_idx < (int)ast_node->operands.size(); ++op_idx) { - ASTNode* src_op = ast_node->operands[op_idx]; - if (!is_ast_leaf(src_op)) - continue; - auto src_it_cd = array_meta.find(src_op->result_name); - if (src_it_cd == array_meta.end()) - continue; - if (src_it_cd->second.ndims == dst_it->second.ndims) - continue; - int higher_tile; - int lower_name_cd; - if (src_it_cd->second.ndims > dst_it->second.ndims) { - higher_tile = src_it_cd->second.tile; - lower_name_cd = dst_name; - } else { - higher_tile = dst_it->second.tile; - lower_name_cd = src_op->result_name; - } - tiles_changed = - try_assign_tile(lower_name_cd, higher_tile) || tiles_changed; - } - continue; - } - - if (is_elementwise(op) || op == Opcode::COPY) { - auto out_it = array_meta.find(ast_node->result_name); - if (out_it == array_meta.end()) - continue; - - for (int op_idx = 0; op_idx < (int)ast_node->operands.size(); ++op_idx) { - tiles_changed = - collect_tile_candidates(collect_tile_candidates, - ast_node->result_name, - out_it->second.ndims, - ast_node->operands[op_idx], - ast_node->get_operand_region(op_idx), - temp_defs) || - tiles_changed; - } - } - } - }); - } - - for (auto& [name, meta] : array_meta) { - if (!meta.decomp_final) - clamp_tile_for_max_count(meta); - } - - auto extract_region_start = [](Region* region_base, int ndims) { - std::array region_start = {0, 0, 0}; - if (region_base == nullptr || region_base->is_global) - return region_start; - switch (ndims) { - case 1: { - auto* r = static_cast*>(region_base); - region_start[0] = r->start[0]; - break; - } - case 2: { - auto* r = static_cast*>(region_base); - region_start[0] = r->start[0]; - region_start[1] = r->start[1]; - break; - } - case 3: { - auto* r = static_cast*>(region_base); - region_start[0] = r->start[0]; - region_start[1] = r->start[1]; - region_start[2] = r->start[2]; - break; - } - default: - break; - } - return region_start; - }; - - auto extract_region_step = [](Region* region_base, int ndims) { - std::array region_step = {1, 1, 1}; - if (region_base == nullptr || region_base->is_global) - return region_step; - switch (ndims) { - case 1: { - auto* r = static_cast*>(region_base); - region_step[0] = r->step[0]; - break; - } - case 2: { - auto* r = static_cast*>(region_base); - region_step[0] = r->step[0]; - region_step[1] = r->step[1]; - break; - } - case 3: { - auto* r = static_cast*>(region_base); - region_step[0] = r->step[0]; - region_step[1] = r->step[1]; - region_step[2] = r->step[2]; - break; - } - default: - break; - } - return region_step; - }; - - auto extract_region_stop = [](Region* region_base, int ndims, - const std::array& global_shape) { - std::array region_stop = {0, 0, 0}; - if (region_base == nullptr || region_base->is_global) { - for (int d = 0; d < ndims; ++d) - region_stop[d] = global_shape[d]; - return region_stop; - } - switch (ndims) { - case 1: { - auto* r = static_cast*>(region_base); - region_stop[0] = r->stop[0]; - break; - } - case 2: { - auto* r = static_cast*>(region_base); - region_stop[0] = r->stop[0]; - region_stop[1] = r->stop[1]; - break; - } - case 3: { - auto* r = static_cast*>(region_base); - region_stop[0] = r->stop[0]; - region_stop[1] = r->stop[1]; - region_stop[2] = r->stop[2]; - break; - } - default: - break; - } - return region_stop; - }; - - auto circ_dist = [](int a, int b, int mod) -> int { - int d = ((a - b) % mod + mod) % mod; - return std::min(d, mod - d); - }; - - std::unordered_map uf_parent, uf_rank; - - auto uf_make = [&](int x) { - if (uf_parent.find(x) == uf_parent.end()) { - uf_parent[x] = x; - uf_rank[x] = 0; - } - }; - - std::function uf_find = [&](int x) -> int { - if (uf_parent[x] != x) - uf_parent[x] = uf_find(uf_parent[x]); - return uf_parent[x]; - }; - - auto uf_unite = [&](int x, int y) { - int rx = uf_find(x), ry = uf_find(y); - if (rx == ry) - return; - if (uf_rank[rx] < uf_rank[ry]) - std::swap(rx, ry); - uf_parent[ry] = rx; - if (uf_rank[rx] == uf_rank[ry]) - uf_rank[rx]++; - }; - - for (auto& [name, meta] : array_meta) - uf_make(name); - - std::unordered_set source_arrays; - - detail::for_each_node_topo(dag, [&](DAGNode* dag_node) { - for (ASTNode* ast_node : dag_node->ast->roots) - if (static_cast(ast_node->opcode) == Opcode::CREATE) - source_arrays.insert(ast_node->result_name); - }); - - auto collect_unionable_nonfixed_leaves = [&](auto&& self, ASTNode* node, int out_name, - int nd, int out_tile, - Region* carried_region, - std::vector& leaves) -> void { - if (node == nullptr || node->is_scalar || node->is_broadcast) - return; - - if (is_ast_leaf(node)) { - int src_name = node->result_name; - if (src_name == out_name) - return; - - auto src_it = array_meta.find(src_name); - if (src_it == array_meta.end() || src_it->second.ndims != nd || - src_it->second.tile != out_tile) - return; - - bool is_fixed_source = src_it->second.decomp_final || source_arrays.count(src_name); - if (is_fixed_source) - return; - - if (carried_region != nullptr && !carried_region->is_global) { - auto rstart = extract_region_start(carried_region, nd); - auto rstep = extract_region_step(carried_region, nd); - for (int d = 0; d < nd; ++d) { - if (rstart[d] != 0 || rstep[d] != 1) - return; - } - } - - leaves.push_back(src_name); - return; - } - - for (int i = 0; i < (int)node->operands.size(); ++i) { - ASTNode* child = node->operands[i]; - if (child == nullptr || child->is_scalar || child->is_broadcast) - continue; - Region* child_region = node->get_operand_region(i); - self(self, child, out_name, nd, out_tile, - child_region ? child_region : carried_region, leaves); - } - }; - - detail::for_each_node_topo(dag, [&](DAGNode* dag_node) { - for (ASTNode* ast_node : dag_node->ast->roots) { - Opcode op = static_cast(ast_node->opcode); - - if (op == Opcode::CREATE) { - continue; - } - - if (is_elementwise(op) || op == Opcode::COPY) { - int out_name = ast_node->result_name; - auto out_it = array_meta.find(out_name); - if (out_it == array_meta.end()) - continue; - int nd = out_it->second.ndims; - int out_tile = out_it->second.tile; - - std::vector union_leaves; - for (int op_idx = 0; op_idx < (int)ast_node->operands.size(); ++op_idx) { - ASTNode* operand = ast_node->operands[op_idx]; - if (operand == nullptr || operand->is_scalar || operand->is_broadcast) - continue; - Region* op_region = ast_node->get_operand_region(op_idx); - collect_unionable_nonfixed_leaves(collect_unionable_nonfixed_leaves, - operand, out_name, nd, out_tile, - op_region, union_leaves); - } - for (int src_name : union_leaves) - uf_unite(out_name, src_name); - continue; - } - - if (op == Opcode::MATMUL || op == Opcode::MATMATMUL) { - for (ASTNode* operand : ast_node->operands) - if (operand != nullptr && !operand->is_scalar) - matmul_arrays.insert(operand->result_name); - matmul_arrays.insert(ast_node->result_name); - } - } - }); - - struct ClassInfo { - int tile = 0; - int ndims = 0; - bool has_fixed_offset = false; - std::array fixed_offset = {0, 0, 0}; - }; - std::unordered_map class_info; - - for (auto& [name, meta] : array_meta) { - int rep = uf_find(name); - auto it = class_info.find(rep); - bool is_fixed = meta.decomp_final || source_arrays.count(name); - if (it == class_info.end()) { - class_info[rep] = {meta.tile, meta.ndims, is_fixed, meta.offset}; - } else if (is_fixed) { - it->second.has_fixed_offset = true; - it->second.fixed_offset = meta.offset; - } - } - - struct ShiftEdge { - int from_class; - int to_class; - std::array shift; - std::array stride; - int64_t weight; - bool forward; - std::array from_to_to = {0, 1, 2}; - std::array to_to_from = {0, 1, 2}; - }; - std::vector shift_edges; - - auto add_shift_edge_if_needed = [&](int out_name, int nd, ASTNode* leaf, Region* region) { - int src_name = leaf->result_name; - if (src_name == out_name) - return; - auto src_it = array_meta.find(src_name); - if (src_it == array_meta.end() || src_it->second.ndims != nd) - return; - - int class_src = uf_find(src_name); - int class_out = uf_find(out_name); - if (class_src == class_out) - return; - - std::array rstart = {0, 0, 0}; - std::array rstep = {1, 1, 1}; - if (region != nullptr && !region->is_global) { - rstart = extract_region_start(region, nd); - rstep = extract_region_step(region, nd); - } - - int64_t weight = 1; - for (int d = 0; d < nd; ++d) - weight *= static_cast(src_it->second.global_shape[d]); - - shift_edges.push_back({class_src, class_out, rstart, rstep, weight, true}); - }; - - auto visit_leaves_for_edges = [&](auto&& self, ASTNode* node, int out_name, int nd, - Region* carried_region) -> void { - if (node == nullptr || node->is_scalar || node->is_broadcast) - return; - if (is_ast_leaf(node)) { - add_shift_edge_if_needed(out_name, nd, node, carried_region); - return; - } - for (int i = 0; i < (int)node->operands.size(); ++i) { - ASTNode* child = node->operands[i]; - if (child == nullptr || child->is_scalar || child->is_broadcast) - continue; - Region* child_region = node->get_operand_region(i); - if (is_ast_leaf(child)) { - add_shift_edge_if_needed(out_name, nd, child, - child_region ? child_region : carried_region); - } else { - self(self, child, out_name, nd, child_region); - } - } - }; - - auto visit_set_region_leaves = [&](auto&& self, ASTNode* node, int target_name, int nd, - const std::array& region_start) -> void { - if (node == nullptr || node->is_scalar || node->is_broadcast) - return; - if (is_ast_leaf(node)) { - int src_name = node->result_name; - if (src_name == target_name) - return; - auto src_it = array_meta.find(src_name); - if (src_it == array_meta.end() || src_it->second.ndims != nd) - return; - - int class_src = uf_find(src_name); - int class_tgt = uf_find(target_name); - if (class_src == class_tgt) - return; - - int64_t weight = 1; - for (int d = 0; d < nd; ++d) - weight *= static_cast(src_it->second.global_shape[d]); - - shift_edges.push_back( - {class_tgt, class_src, region_start, {1, 1, 1}, weight, true}); - return; - } - for (ASTNode* child : node->operands) - self(self, child, target_name, nd, region_start); - }; - - detail::for_each_node_topo(dag, [&](DAGNode* dag_node) { - for (ASTNode* ast_node : dag_node->ast->roots) { - Opcode op = static_cast(ast_node->opcode); - - if (op == Opcode::SET_REGION) { - if (ast_node->operands.size() < 2 || ast_node->region == nullptr) - continue; - int target_name = ast_node->operands[0]->result_name; - auto target_it = array_meta.find(target_name); - if (target_it == array_meta.end()) - continue; - int nd = target_it->second.ndims; - auto region_start = extract_region_start(ast_node->region, nd); - - for (int op_idx = 1; op_idx < (int)ast_node->operands.size(); ++op_idx) - visit_set_region_leaves(visit_set_region_leaves, - ast_node->operands[op_idx], target_name, nd, - region_start); - - for (int op_idx = 1; op_idx < (int)ast_node->operands.size(); ++op_idx) { - ASTNode* src_op = ast_node->operands[op_idx]; - if (!is_ast_leaf(src_op)) - continue; - int src_name_cd = src_op->result_name; - if (src_name_cd == target_name) - continue; - auto src_it_cd = array_meta.find(src_name_cd); - if (src_it_cd == array_meta.end()) - continue; - int nd_src = src_it_cd->second.ndims; - if (nd_src == nd) - continue; - - std::array from_to_to_cd = {-1, -1, -1}; - std::array to_to_from_cd = {-1, -1, -1}; - - if (nd < nd_src) { - int td = 0; - for (int sd = 0; sd < nd_src && td < nd; ++sd) { - if (src_it_cd->second.global_shape[sd] > 1) { - from_to_to_cd[td] = sd; - to_to_from_cd[sd] = td; - ++td; - } - } - if (td != nd) - continue; - } else { - auto region_stop = extract_region_stop( - ast_node->region, nd, target_it->second.global_shape); - int sd = 0; - for (int td = 0; td < nd && sd < nd_src; ++td) { - int size = region_stop[td] - region_start[td]; - if (size > 1) { - from_to_to_cd[td] = sd; - to_to_from_cd[sd] = td; - ++sd; - } - } - if (sd != nd_src) - continue; - } - - int class_tgt_cd = uf_find(target_name); - int class_src_cd = uf_find(src_name_cd); - if (class_tgt_cd == class_src_cd) - continue; - - int64_t weight_cd = 1; - int nd_lo = std::min(nd, nd_src); - auto& lo_meta = (nd <= nd_src) ? target_it->second : src_it_cd->second; - for (int d = 0; d < nd_lo; ++d) - weight_cd *= static_cast(lo_meta.global_shape[d]); - - shift_edges.push_back({class_tgt_cd, class_src_cd, region_start, - {1, 1, 1}, weight_cd, true, from_to_to_cd, - to_to_from_cd}); - } - continue; - } - - if (is_elementwise(op) || op == Opcode::COPY) { - int out_name = ast_node->result_name; - auto out_it = array_meta.find(out_name); - if (out_it == array_meta.end()) - continue; - int nd = out_it->second.ndims; - - for (int op_idx = 0; op_idx < (int)ast_node->operands.size(); ++op_idx) { - ASTNode* operand = ast_node->operands[op_idx]; - if (operand == nullptr || operand->is_scalar || operand->is_broadcast) - continue; - Region* op_region = ast_node->get_operand_region(op_idx); - if (is_ast_leaf(operand)) { - add_shift_edge_if_needed(out_name, nd, operand, op_region); - } else { - visit_leaves_for_edges(visit_leaves_for_edges, operand, out_name, nd, - op_region); - } - } - } - } - }); - - struct AdjEdge { - int neighbor; - std::array shift; - std::array stride; - std::array neighbor_dim = {0, 1, 2}; - int64_t weight; - bool forward; - int from_class; - int to_class; - }; - std::unordered_map> adj; - - for (auto& e : shift_edges) { - if (e.from_class == e.to_class) - continue; - adj[e.from_class].push_back( - {e.to_class, e.shift, e.stride, e.from_to_to, e.weight, true, e.from_class, - e.to_class}); - adj[e.to_class].push_back( - {e.from_class, e.shift, e.stride, e.to_to_from, e.weight, false, e.from_class, - e.to_class}); - } - - std::unordered_map comp_id; - std::unordered_map tree_parent; - std::unordered_map> tree_children; - std::unordered_set cyclic_components; - int num_components = 0; - - std::unordered_set all_classes; - for (auto& [name, _] : array_meta) - all_classes.insert(uf_find(name)); - - std::unordered_map comp_root; - - for (int cls : all_classes) { - if (comp_id.count(cls)) - continue; - int cid = num_components++; - comp_root[cid] = cls; - - std::queue bfs; - bfs.push(cls); - comp_id[cls] = cid; - tree_parent[cls] = -1; - - while (!bfs.empty()) { - int v = bfs.front(); - bfs.pop(); - std::unordered_set seen_neighbors; - for (auto& e : adj[v]) { - if (seen_neighbors.count(e.neighbor)) - continue; - seen_neighbors.insert(e.neighbor); - if (!comp_id.count(e.neighbor)) { - comp_id[e.neighbor] = cid; - tree_parent[e.neighbor] = v; - tree_children[v].push_back(e.neighbor); - bfs.push(e.neighbor); - } else if (e.neighbor != tree_parent[v]) { - cyclic_components.insert(cid); - } - } - auto ci = class_info.find(v); - if (ci != class_info.end() && ci->second.has_fixed_offset) - comp_root[cid] = v; - } - } - - tree_parent.clear(); - tree_children.clear(); - - for (auto& [cid, root] : comp_root) { - std::queue bfs; - bfs.push(root); - tree_parent[root] = -1; - - std::unordered_set visited; - visited.insert(root); - - while (!bfs.empty()) { - int v = bfs.front(); - bfs.pop(); - for (auto& e : adj[v]) { - if (!visited.count(e.neighbor) && comp_id[e.neighbor] == cid) { - visited.insert(e.neighbor); - tree_parent[e.neighbor] = v; - tree_children[v].push_back(e.neighbor); - bfs.push(e.neighbor); - } - } - } - } - - constexpr int64_t INF = std::numeric_limits::max() / 2; - - std::unordered_map, 3>> cands; - - auto add_cand = [](std::vector& v, int o) { - auto it = std::lower_bound(v.begin(), v.end(), o); - if (it == v.end() || *it != o) - v.insert(it, o); - }; - - auto compute_desired = [&](int o_from, int shift_d, int stride_d, int tile_to) -> int { - int step = stride_d > 1 ? stride_d : 1; - return positive_mod((o_from + shift_d) / step, tile_to); - }; - - auto reverse_desired = [&](int o_target, int shift_d, int stride_d, int tile_from, - int tile_to) -> std::vector { - int step = stride_d > 1 ? stride_d : 1; - std::vector result; - int base = o_target * step - shift_d; - int period = tile_to * step; - if (period <= 0) - period = 1; - int num_periods = (tile_from + period - 1) / period; - for (int k = 0; k <= num_periods; ++k) { - for (int delta = 0; delta < step; ++delta) { - int raw = base + delta + k * period; - int cand = ((raw % tile_from) + tile_from) % tile_from; - if (cand >= 0 && cand < tile_from && - compute_desired(cand, shift_d, stride_d, tile_to) == o_target) { - result.push_back(cand); - } - } - } - std::sort(result.begin(), result.end()); - result.erase(std::unique(result.begin(), result.end()), result.end()); - return result; - }; - - std::vector post_order; - for (auto& [cid, root] : comp_root) { - if (cyclic_components.count(cid)) - continue; - - std::stack> stk; - stk.push({root, false}); - while (!stk.empty()) { - auto [v, processed] = stk.top(); - stk.pop(); - if (processed) { - post_order.push_back(v); - continue; - } - stk.push({v, true}); - for (int c : tree_children[v]) - stk.push({c, false}); - } - } - - for (int v : post_order) { - auto ci = class_info.find(v); - if (ci == class_info.end()) - continue; - int nd = ci->second.ndims; - int tile_v = ci->second.tile; - if (tile_v <= 0) - tile_v = 1; - - auto& cands_v = cands[v]; - - if (ci->second.has_fixed_offset) { - for (int d = 0; d < nd; ++d) - add_cand(cands_v[d], positive_mod(ci->second.fixed_offset[d], tile_v)); - } - - for (int c : tree_children[v]) { - auto cc = class_info.find(c); - if (cc == class_info.end()) - continue; - int tile_c = cc->second.tile; - if (tile_c <= 0) - tile_c = 1; - - for (auto& e : adj[v]) { - if (e.neighbor != c) - continue; - for (int d = 0; d < nd; ++d) { - int d_c = e.neighbor_dim[d]; - if (d_c < 0) - continue; - int d_from = e.forward ? d : d_c; - for (int oc : cands[c][d_c]) { - if (e.forward) { - for (int rv : reverse_desired(oc, e.shift[d_from], - e.stride[d_from], tile_v, tile_c)) - add_cand(cands_v[d], rv); - } else { - add_cand(cands_v[d], compute_desired(oc, e.shift[d_from], - e.stride[d_from], tile_v)); - } - } - } - } - } - - for (int d = 0; d < nd; ++d) { - if (cands_v[d].empty()) - add_cand(cands_v[d], 0); - } - } - - { - std::queue work; - for (auto& [cid, root] : comp_root) { - if (!cyclic_components.count(cid)) - work.push(root); - } - - while (!work.empty()) { - int v = work.front(); - work.pop(); - auto ci = class_info.find(v); - if (ci == class_info.end()) - continue; - int nd = ci->second.ndims; - int tile_v = ci->second.tile; - if (tile_v <= 0) - tile_v = 1; - - for (int c : tree_children[v]) { - auto cc = class_info.find(c); - if (cc == class_info.end()) - continue; - int tile_c = cc->second.tile; - if (tile_c <= 0) - tile_c = 1; - - for (auto& e : adj[v]) { - if (e.neighbor != c) - continue; - for (int d = 0; d < nd; ++d) { - int d_c = e.neighbor_dim[d]; - if (d_c < 0) - continue; - int d_from = e.forward ? d : d_c; - for (int ov : cands[v][d]) { - if (e.forward) { - add_cand(cands[c][d_c], - compute_desired(ov, e.shift[d_from], - e.stride[d_from], tile_c)); - } else { - for (int rc : reverse_desired(ov, e.shift[d_from], - e.stride[d_from], tile_c, tile_v)) - add_cand(cands[c][d_c], rc); - } - } - } - } - work.push(c); - } - } - } - - std::unordered_map, 3>> dp; - - for (int v : post_order) { - auto ci = class_info.find(v); - if (ci == class_info.end()) - continue; - int nd = ci->second.ndims; - int tile_v = ci->second.tile; - if (tile_v <= 0) - tile_v = 1; - - auto& dp_v = dp[v]; - auto& cands_v = cands[v]; - - for (int d = 0; d < 3; ++d) { - int nc = (int)cands_v[d].size(); - dp_v[d].assign(nc, 0); - if (d < nd && ci->second.has_fixed_offset) { - int fixed = positive_mod(ci->second.fixed_offset[d], tile_v); - for (int i = 0; i < nc; ++i) - dp_v[d][i] = (cands_v[d][i] == fixed) ? 0 : INF; - } - } - - for (int c : tree_children[v]) { - auto cc = class_info.find(c); - if (cc == class_info.end()) - continue; - int tile_c = cc->second.tile; - if (tile_c <= 0) - tile_c = 1; - - std::vector edges_vc; - for (auto& e : adj[v]) - if (e.neighbor == c) - edges_vc.push_back(&e); - - auto& cands_c = cands[c]; - auto& dp_c = dp[c]; - - for (int d = 0; d < nd; ++d) { - int d_c = d; - for (auto* ep : edges_vc) { - if (ep->neighbor_dim[d] != d) { - d_c = ep->neighbor_dim[d]; - break; - } - } - if (d_c < 0) - continue; - - int nc_v = (int)cands_v[d].size(); - int nc_cd = (int)cands_c[d_c].size(); - - std::vector h(nc_v, INF); - for (int iv = 0; iv < nc_v; ++iv) { - int ov = cands_v[d][iv]; - for (int ic = 0; ic < nc_cd; ++ic) { - int oc = cands_c[d_c][ic]; - int64_t cost = dp_c[d_c][ic]; - if (cost >= INF) - continue; - for (auto* ep : edges_vc) { - int ep_dc = ep->neighbor_dim[d]; - if (ep_dc < 0) - continue; - int d_from = ep->forward ? d : ep_dc; - int o_from, tile_to, o_to; - if (ep->forward) { - o_from = ov; - tile_to = tile_c; - o_to = oc; - } else { - o_from = oc; - tile_to = tile_v; - o_to = ov; - } - int step = ep->stride[d_from] > 1 ? ep->stride[d_from] : 1; - int desired = positive_mod( - (o_from + ep->shift[d_from]) / step, tile_to); - cost += ep->weight * circ_dist(desired, o_to, tile_to); - if (cost >= INF) - break; - } - if (cost < h[iv]) - h[iv] = cost; - } - } - for (int iv = 0; iv < nc_v; ++iv) { - if (dp_v[d][iv] < INF && h[iv] < INF) - dp_v[d][iv] += h[iv]; - else - dp_v[d][iv] = INF; - } - } - } - } - - std::unordered_map> class_phase; - - for (auto& [cid, root] : comp_root) { - if (cyclic_components.count(cid)) - continue; - auto ci = class_info.find(root); - if (ci == class_info.end()) - continue; - int nd = ci->second.ndims; - - std::array best = {0, 0, 0}; - for (int d = 0; d < nd; ++d) { - int64_t best_cost = INF; - auto& cv = cands[root][d]; - auto& dv = dp[root][d]; - for (int i = 0; i < (int)cv.size(); ++i) { - if (dv[i] < best_cost) { - best_cost = dv[i]; - best[d] = cv[i]; - } - } - } - class_phase[root] = best; - } - - { - std::queue work; - for (auto& [cid, root] : comp_root) { - if (!cyclic_components.count(cid)) - work.push(root); - } - - while (!work.empty()) { - int v = work.front(); - work.pop(); - auto ci_v = class_info.find(v); - if (ci_v == class_info.end()) - continue; - int nd = ci_v->second.ndims; - int tile_v = ci_v->second.tile; - if (tile_v <= 0) - tile_v = 1; - - for (int c : tree_children[v]) { - auto ci_c = class_info.find(c); - if (ci_c == class_info.end()) - continue; - int tile_c = ci_c->second.tile; - if (tile_c <= 0) - tile_c = 1; - - std::vector edges_vc; - for (auto& e : adj[v]) - if (e.neighbor == c) - edges_vc.push_back(&e); - - auto& cands_c = cands[c]; - auto& dp_c = dp[c]; - int nd_c = ci_c->second.ndims; - - std::array best_c = {0, 0, 0}; - for (int d_c = 0; d_c < nd_c; ++d_c) { - int ov = -1; - int d_parent = -1; - for (int dp_d = 0; dp_d < nd; ++dp_d) { - for (auto* ep : edges_vc) { - if (ep->neighbor_dim[dp_d] == d_c) { - d_parent = dp_d; - ov = class_phase[v][dp_d]; - break; - } - } - if (d_parent >= 0) - break; - } - - if (ov < 0) { - int64_t best_cost = INF; - for (int ic = 0; ic < (int)cands_c[d_c].size(); ++ic) { - if (dp_c[d_c][ic] < best_cost) { - best_cost = dp_c[d_c][ic]; - best_c[d_c] = cands_c[d_c][ic]; - } - } - continue; - } - - int64_t best_cost = INF; - int nc_cd = (int)cands_c[d_c].size(); - for (int ic = 0; ic < nc_cd; ++ic) { - int oc = cands_c[d_c][ic]; - int64_t cost = dp_c[d_c][ic]; - if (cost >= INF) - continue; - for (auto* ep : edges_vc) { - int ep_dc = ep->neighbor_dim[d_parent]; - if (ep_dc < 0 || ep_dc != d_c) - continue; - int d_from = ep->forward ? d_parent : d_c; - int o_from, tile_to, o_to; - if (ep->forward) { - o_from = ov; - tile_to = tile_c; - o_to = oc; - } else { - o_from = oc; - tile_to = tile_v; - o_to = ov; - } - int step = ep->stride[d_from] > 1 ? ep->stride[d_from] : 1; - int desired = positive_mod( - (o_from + ep->shift[d_from]) / step, tile_to); - cost += ep->weight * circ_dist(desired, o_to, tile_to); - if (cost >= INF) - break; - } - if (cost < best_cost) { - best_cost = cost; - best_c[d_c] = oc; - } - } - } - class_phase[c] = best_c; - work.push(c); - } - } - } - - auto positive_mod64 = [](int64_t value, int mod) -> int { - if (mod <= 0) - return static_cast(value); - int64_t r = value % mod; - if (r < 0) - r += mod; - return static_cast(r); - }; - - auto abs_dist64 = [](int64_t a, int64_t b) -> int64_t { - return (a >= b) ? (a - b) : (b - a); - }; - - auto nearest_congruent = [&](int64_t anchor, int phase, int tile) -> int { - if (tile <= 0) - return static_cast(anchor); - - int target = positive_mod(phase, tile); - int64_t lower = target; - if (anchor >= target) { - lower += ((anchor - target) / tile) * static_cast(tile); - } else { - lower -= ((target - anchor + tile - 1) / tile) * static_cast(tile); - } - int64_t upper = lower + tile; - return (abs_dist64(upper, anchor) < abs_dist64(lower, anchor)) - ? static_cast(upper) - : static_cast(lower); - }; - - auto reverse_absolute = [&](int64_t target_abs, int shift_d, int stride_d, int phase_from, - int tile_from) -> int { - int step = stride_d > 1 ? stride_d : 1; - int64_t low = target_abs * step - shift_d; - int64_t center2 = 2 * low + (step - 1); - bool found = false; - int64_t best = low; - int64_t best_score = 0; - int phase = positive_mod(phase_from, tile_from); - - for (int delta = 0; delta < step; ++delta) { - int64_t cand = low + delta; - if (positive_mod64(cand, tile_from) != phase) - continue; - if ((cand + shift_d) / step != target_abs) - continue; - int64_t score = abs_dist64(2 * cand, center2); - if (!found || score < best_score || (score == best_score && cand < best)) { - found = true; - best = cand; - best_score = score; - } - } - - if (found) - return static_cast(best); - - return nearest_congruent(low, phase_from, tile_from); - }; - - struct AbsoluteProposal { - int value = 0; - int64_t weight = 0; - }; - - auto choose_absolute = [&](const std::vector& proposals, int phase, - int tile) -> int { - std::vector candidates; - candidates.push_back(nearest_congruent(phase, phase, tile)); - for (auto const& p : proposals) { - auto it = std::lower_bound(candidates.begin(), candidates.end(), p.value); - if (it == candidates.end() || *it != p.value) - candidates.insert(it, p.value); - } - - int best = candidates.front(); - int64_t best_cost = std::numeric_limits::max(); - for (int cand : candidates) { - int64_t cost = 0; - for (auto const& p : proposals) - cost += p.weight * abs_dist64(cand, p.value); - if (cost < best_cost || (cost == best_cost && cand < best)) { - best_cost = cost; - best = cand; - } - } - return best; - }; - - std::unordered_map> class_abs_offset; - - for (auto& [cid, root] : comp_root) { - if (cyclic_components.count(cid)) - continue; - - auto ci = class_info.find(root); - auto phase_it = class_phase.find(root); - if (ci == class_info.end() || phase_it == class_phase.end()) - continue; - - std::array root_abs = {0, 0, 0}; - int tile_root = ci->second.tile; - if (tile_root <= 0) - tile_root = 1; - - for (int d = 0; d < ci->second.ndims; ++d) { - if (ci->second.has_fixed_offset) - root_abs[d] = ci->second.fixed_offset[d]; - else - root_abs[d] = - nearest_congruent(phase_it->second[d], phase_it->second[d], tile_root); - } - class_abs_offset[root] = root_abs; - } - - { - std::queue work; - for (auto& [cid, root] : comp_root) { - if (!cyclic_components.count(cid) && class_abs_offset.count(root)) - work.push(root); - } - - while (!work.empty()) { - int v = work.front(); - work.pop(); - - auto ci_v = class_info.find(v); - auto abs_v_it = class_abs_offset.find(v); - if (ci_v == class_info.end() || abs_v_it == class_abs_offset.end()) - continue; - int nd_v = ci_v->second.ndims; - - for (int c : tree_children[v]) { - auto ci_c = class_info.find(c); - auto phase_c_it = class_phase.find(c); - if (ci_c == class_info.end() || phase_c_it == class_phase.end()) - continue; - - int tile_c = ci_c->second.tile; - if (tile_c <= 0) - tile_c = 1; - - std::vector edges_vc; - for (auto& e : adj[v]) - if (e.neighbor == c) - edges_vc.push_back(&e); - - std::array abs_c = {0, 0, 0}; - for (int d_c = 0; d_c < ci_c->second.ndims; ++d_c) { - if (ci_c->second.has_fixed_offset) { - abs_c[d_c] = ci_c->second.fixed_offset[d_c]; - continue; - } - - std::vector proposals; - for (auto* ep : edges_vc) { - for (int d_parent = 0; d_parent < nd_v; ++d_parent) { - if (ep->neighbor_dim[d_parent] != d_c) - continue; - - int value = 0; - if (ep->forward) { - int d_from = d_parent; - int step = ep->stride[d_from] > 1 ? ep->stride[d_from] : 1; - int64_t anchor = - (static_cast(abs_v_it->second[d_parent]) + - ep->shift[d_from]) / - step; - value = - nearest_congruent(anchor, phase_c_it->second[d_c], tile_c); - } else { - int d_from = d_c; - value = reverse_absolute(abs_v_it->second[d_parent], - ep->shift[d_from], - ep->stride[d_from], - phase_c_it->second[d_c], tile_c); - } - proposals.push_back({value, ep->weight}); - } - } - - abs_c[d_c] = - choose_absolute(proposals, phase_c_it->second[d_c], tile_c); - } - - class_abs_offset[c] = abs_c; - work.push(c); - } - } - } - - for (auto& [name, meta] : array_meta) { - if (meta.decomp_final) - continue; - if (matmul_arrays.count(name)) - continue; - int rep = uf_find(name); - auto comp_it = comp_id.find(rep); - if (comp_it != comp_id.end() && cyclic_components.count(comp_it->second)) - continue; - auto it = class_abs_offset.find(rep); - if (it != class_abs_offset.end()) - meta.offset = it->second; - } - - if (!cyclic_components.empty()) { - struct DesiredOffset { - std::array desired_phase; - std::array desired_anchor; - int64_t weight; - }; - std::unordered_map> all_desired; - - auto collect_fallback_producer = - [&](auto&& self, ASTNode* node, int out_name, int nd, int out_tile) -> void { - if (node == nullptr || node->is_scalar || node->is_broadcast) - return; - if (is_ast_leaf(node)) { - int src_name = node->result_name; - if (src_name == out_name) - return; - auto src_it = array_meta.find(src_name); - if (src_it == array_meta.end() || src_it->second.ndims != nd) - return; - if (out_tile <= 0) - return; - std::array desired_phase = {0, 0, 0}; - std::array desired_anchor = {0, 0, 0}; - for (int d = 0; d < nd; ++d) { - desired_anchor[d] = src_it->second.offset[d]; - desired_phase[d] = positive_mod(src_it->second.offset[d], out_tile); - } - int64_t weight = 1; - for (int d = 0; d < nd; ++d) - weight *= static_cast(src_it->second.global_shape[d]); - all_desired[out_name].push_back({desired_phase, desired_anchor, weight}); - return; - } - for (int i = 0; i < (int)node->operands.size(); ++i) { - ASTNode* child = node->operands[i]; - if (child == nullptr || child->is_scalar || child->is_broadcast) - continue; - Region* child_region = node->get_operand_region(i); - if (is_ast_leaf(child) && child_region != nullptr && !child_region->is_global) { - int src_name = child->result_name; - if (src_name == out_name) - continue; - auto src_it = array_meta.find(src_name); - if (src_it == array_meta.end() || src_it->second.ndims != nd) - continue; - if (out_tile <= 0) - continue; - auto rstart = extract_region_start(child_region, nd); - auto rstep = extract_region_step(child_region, nd); - std::array desired_phase = {0, 0, 0}; - std::array desired_anchor = {0, 0, 0}; - for (int d = 0; d < nd; ++d) { - int step = (rstep[d] > 1) ? rstep[d] : 1; - int raw = src_it->second.offset[d] + rstart[d]; - desired_anchor[d] = (step > 1) ? (raw / step) : raw; - desired_phase[d] = (step > 1) - ? positive_mod(raw / step, out_tile) - : positive_mod(raw, out_tile); - } - int64_t weight = 1; - for (int d = 0; d < nd; ++d) - weight *= static_cast(src_it->second.global_shape[d]); - all_desired[out_name].push_back( - {desired_phase, desired_anchor, weight}); - } else { - self(self, child, out_name, nd, out_tile); - } - } - }; - - auto collect_fallback_consumer = - [&](auto&& self, ASTNode* node, int target_name, int nd, - const ArrayMetadata& target_meta, const std::array& region_start) - -> void { - if (node == nullptr || node->is_scalar || node->is_broadcast) - return; - if (is_ast_leaf(node)) { - int src_name = node->result_name; - if (src_name == target_name) - return; - auto src_it = array_meta.find(src_name); - if (src_it == array_meta.end() || src_it->second.ndims != nd) - return; - int tile = src_it->second.tile; - if (tile <= 0) - return; - std::array desired_phase = {0, 0, 0}; - std::array desired_anchor = {0, 0, 0}; - for (int d = 0; d < nd; ++d) { - desired_anchor[d] = target_meta.offset[d] + region_start[d]; - desired_phase[d] = - positive_mod(target_meta.offset[d] + region_start[d], tile); - } - int64_t weight = 1; - for (int d = 0; d < nd; ++d) - weight *= static_cast(src_it->second.global_shape[d]); - all_desired[src_name].push_back({desired_phase, desired_anchor, weight}); - return; - } - for (ASTNode* child : node->operands) - self(self, child, target_name, nd, target_meta, region_start); - }; - - auto is_cyclic_array = [&](int name) { - int rep = uf_find(name); - auto cit = comp_id.find(rep); - return cit != comp_id.end() && cyclic_components.count(cit->second); - }; - - detail::for_each_node_topo(dag, [&](DAGNode* dag_node) { - for (ASTNode* ast_node : dag_node->ast->roots) { - Opcode op = static_cast(ast_node->opcode); - - if (op == Opcode::SET_REGION) { - if (ast_node->operands.size() < 2 || !ast_node->region) - continue; - int target_name = ast_node->operands[0]->result_name; - auto target_it = array_meta.find(target_name); - if (target_it == array_meta.end()) - continue; - int nd = target_it->second.ndims; - auto region_start = extract_region_start(ast_node->region, nd); - for (int i = 1; i < (int)ast_node->operands.size(); ++i) - collect_fallback_consumer(collect_fallback_consumer, - ast_node->operands[i], target_name, nd, - target_it->second, region_start); - continue; - } - - if (is_elementwise(op) || op == Opcode::COPY) { - int out_name = ast_node->result_name; - if (!is_cyclic_array(out_name)) - continue; - auto out_it = array_meta.find(out_name); - if (out_it == array_meta.end()) - continue; - int nd = out_it->second.ndims; - int out_tile = out_it->second.tile; - for (int i = 0; i < (int)ast_node->operands.size(); ++i) { - ASTNode* operand = ast_node->operands[i]; - if (!operand || operand->is_scalar || operand->is_broadcast) - continue; - collect_fallback_producer(collect_fallback_producer, operand, out_name, - nd, out_tile); - } - } - } - }); - - for (auto& [name, constraints] : all_desired) { - if (constraints.empty()) - continue; - if (!is_cyclic_array(name)) - continue; - auto meta_it = array_meta.find(name); - if (meta_it == array_meta.end() || meta_it->second.decomp_final) - continue; - if (matmul_arrays.count(name)) - continue; - - int nd = meta_it->second.ndims; - int tile = meta_it->second.tile; - if (tile <= 0) - continue; - - std::array best_phase = {0, 0, 0}; - for (int d = 0; d < nd; ++d) { - std::vector fb_cands; - fb_cands.push_back(0); - for (auto& c : constraints) { - int o = positive_mod(c.desired_phase[d], tile); - auto it2 = std::lower_bound(fb_cands.begin(), fb_cands.end(), o); - if (it2 == fb_cands.end() || *it2 != o) - fb_cands.insert(it2, o); - } - - int64_t best_cost = std::numeric_limits::max(); - int best_o = 0; - for (int o : fb_cands) { - int64_t cost = 0; - for (auto& c : constraints) - cost += c.weight * circ_dist(o, c.desired_phase[d], tile); - if (cost < best_cost) { - best_cost = cost; - best_o = o; - } - } - best_phase[d] = best_o; - } - - std::array best_offset = {0, 0, 0}; - for (int d = 0; d < nd; ++d) { - std::vector proposals; - for (auto& c : constraints) { - proposals.push_back( - {nearest_congruent(c.desired_anchor[d], best_phase[d], tile), - c.weight}); - } - best_offset[d] = choose_absolute(proposals, best_phase[d], tile); - } - meta_it->second.offset = best_offset; - } - } - - for (int name : matmul_arrays) { - auto meta_it = array_meta.find(name); - if (meta_it != array_meta.end() && !meta_it->second.decomp_final) - meta_it->second.offset = zero_offset; - } - } - - for (auto& [name, meta] : array_meta) - meta.decomp_final = true; -} - -} // namespace decomposition_solver diff --git a/src/dispatch.hpp b/src/dispatch.hpp deleted file mode 100644 index 13ae000..0000000 --- a/src/dispatch.hpp +++ /dev/null @@ -1,265 +0,0 @@ -#pragma once - -#include "backend.hpp" - -// ---- Stream policy helpers (USE_KOKKOS) ---- -// Under USE_NVIDIA: Kokkos::Cuda exec space instances bound to per-chare streams. -// Under USE_KOKKOS without USE_NVIDIA: default execution space + Kokkos::fence(). - -#ifdef USE_KOKKOS -#ifdef USE_NVIDIA -#define CT_COMPUTE_POLICY(partition, n) \ - Kokkos::RangePolicy((partition)->compute_exec, 0, (n)) -#define CT_COMM_POLICY(partition, n) \ - Kokkos::RangePolicy((partition)->comm_exec, 0, (n)) -#else -#define CT_COMPUTE_POLICY(partition, n) Kokkos::RangePolicy<>(0, (n)) -#define CT_COMM_POLICY(partition, n) Kokkos::RangePolicy<>(0, (n)) -#endif -#endif - -// Helper: fence Kokkos on non-NVIDIA GPU builds (no-op on CPU-only and NVIDIA). -#if defined(USE_KOKKOS) && !defined(USE_NVIDIA) -#define CT_KOKKOS_FENCE() Kokkos::fence() -#else -#define CT_KOKKOS_FENCE() ((void)0) -#endif - -// ---- HAPI callback structs and functions (USE_NVIDIA only) ---- - -#ifdef USE_NVIDIA -/// Param struct for compute-done callbacks (hapiAddCallback). -/// Heap-allocated, freed by the callback function. -template -struct ComputeDoneParam { - PartitionImpl* partition; - int node_id; - bool has_comm; -}; - -/// C callback invoked by HAPI when all prior compute_stream work completes. -template -static void compute_done_cb(void* param, void* msg) { - auto* p = static_cast*>(param); - if (p->has_comm) { - auto it = p->partition->executor->pending.find(p->node_id); - if (it != p->partition->executor->pending.end()) { - it->second.clear_remote_buffers(); - p->partition->executor->pending.erase(it); - } - } - p->partition->executor->node_finished(p->node_id); - delete p; -} - -/// Param struct for deferred-send callbacks (hapiAddCallback). -/// N_sender is the sender partition dimension, N_target is the target partition dimension. -/// For same-dimension sends, N_sender == N_target (the default). -template -struct DeferredSendParam { - PartitionImpl* partition; - typename PartitionTraits::ProxyType target_proxy; - int node_id; - int input_index; - int inp_name; - ChareIndex target_ci; - int region_data[N_target * 3]; - int64_t byte_size; - void* send_buf; - int send_id; - - /// Same-dimension send constructor: target_proxy defaults to partition's own proxy. - DeferredSendParam(PartitionImpl* p) - : partition(p), target_proxy(p->thisProxy) {} - - /// Cross-dimensional send constructor: target_proxy is explicitly provided. - DeferredSendParam(PartitionImpl* p, typename PartitionTraits::ProxyType tp) - : partition(p), target_proxy(tp) {} -}; - -/// C callback invoked by HAPI when comm_stream packing completes. -/// Performs the actual CkDeviceBuffer send. -template -static void deferred_send_cb(void* param, void* msg) { - auto* p = static_cast*>(param); - p->partition->pending_sends[p->send_id] = p->send_buf; - ChareIndex self_ci; - for (int d = 0; d < N_sender; ++d) - self_ci.idx[d] = p->partition->index[d]; - CkCallback cb(PartitionTraits::CkIndexType::send_complete(nullptr), - proxy_at(p->partition->thisProxy, self_ci)); - proxy_at(p->target_proxy, p->target_ci).receive_data( - p->node_id, p->input_index, p->inp_name, N_target, p->region_data, p->byte_size, - CkDeviceBuffer(reinterpret_cast(p->send_buf), cb)); - delete p; -} - -#endif - -#ifdef USE_KOKKOS -/// Helper: complete a device-side pack-and-send after a parallel_for on comm_stream. -/// On NVIDIA: defers the send via hapiAddCallback on comm_stream. -/// On non-NVIDIA Kokkos: fences then sends synchronously. -template -static void device_pack_send(PartitionImpl* partition, - typename PartitionTraits::ProxyType& target_proxy, - const ChareIndex& target_ci, int node_id, int input_index, - int name, int* region_data, int64_t byte_size, void* send_buf) { - auto* ds = new DeferredSendParam(partition, target_proxy); - ds->node_id = node_id; - ds->input_index = input_index; - ds->inp_name = name; - ds->target_ci = target_ci; - memcpy(ds->region_data, region_data, sizeof(int) * N_target * 3); - ds->byte_size = byte_size; - ds->send_buf = send_buf; - ds->send_id = partition->next_send_id++; -#ifdef USE_NVIDIA - CkCallback hapi_cb(deferred_send_cb, ds); - hapiAddCallback(partition->comm_stream_raw, &hapi_cb); -#else - Kokkos::fence(); - deferred_send_cb(ds, nullptr); -#endif -} -#endif - -/// Dispatch a JIT-compiled kernel with the given packed memref descriptors -/// and optional broadcast scalar arguments. -/// -/// Function signature layout (matching JIT buildFromAST): -/// [n_input_memrefs memref args] [n_scalars scalar args] [n_output_memrefs memref args] -/// -/// @param n_input_memrefs Number of input memrefs (descs[0..n_input_memrefs-1]) -/// @param scalar_args Pointers to scalar values (broadcast leaves), inserted -/// between input and output memref args -/// @param descs All memref descriptors: inputs first, then outputs -template -static void dispatch_kernel(void* func_ptr, int n_memrefs, std::vector>& descs, - int n_input_memrefs = -1, - const std::vector& scalar_args = {} -#ifdef USE_NVIDIA - , - cudaStream_t stream = nullptr -#endif -) { - if (n_input_memrefs < 0) - n_input_memrefs = n_memrefs; // backward compat: all are inputs (no outputs separate) - - constexpr int FIELDS_PER_MEMREF = MemRef::fields_per_memref(); - int n_scalar_args = (int)scalar_args.size(); - - // Pack args: input memrefs, then scalar args, then output memrefs - std::vector args; - args.reserve(n_memrefs * FIELDS_PER_MEMREF + n_scalar_args); - - // Input memrefs - auto pack_memref = [&](int i) { - args.push_back(&descs[i].allocated); - args.push_back(&descs[i].aligned); - args.push_back(&descs[i].offset); - for (int d = 0; d < N; ++d) - args.push_back(&descs[i].sizes[d]); - for (int d = 0; d < N; ++d) - args.push_back(&descs[i].strides[d]); - }; - - for (int i = 0; i < n_input_memrefs; ++i) - pack_memref(i); - - // Broadcast scalar arguments - for (int i = 0; i < n_scalar_args; ++i) - args.push_back(scalar_args[i]); - - // Output memrefs - for (int i = n_input_memrefs; i < n_memrefs; ++i) - pack_memref(i); - - auto total_elements = [&]() -> int64_t { - if (n_memrefs == 0) - return 0; - int64_t n = 1; - for (int d = 0; d < N; ++d) - n *= descs[0].sizes[d]; - return n; - }; - -#if defined(USE_NVIDIA) - { - CUfunction cuFunc = reinterpret_cast(func_ptr); - int64_t n_elements = total_elements(); - int blockSize = 256; - int gridSize = ((int)n_elements + blockSize - 1) / blockSize; - cuLaunchKernel(cuFunc, gridSize, 1, 1, blockSize, 1, 1, 0, stream, args.data(), nullptr); - } -#elif defined(USE_AMD) - { - hipFunction_t hipFunc = reinterpret_cast(func_ptr); - int64_t n_elements = total_elements(); - int blockSize = 256; - int gridSize = ((int)n_elements + blockSize - 1) / blockSize; - hipModuleLaunchKernel(hipFunc, gridSize, 1, 1, blockSize, 1, 1, 0, nullptr, args.data(), - nullptr); - hipDeviceSynchronize(); - } -#elif defined(USE_INTEL) - { - ze_kernel_handle_t zeKernel = reinterpret_cast(func_ptr); - int64_t n_elements = total_elements(); - uint32_t groupSizeX = 256; - zeKernelSetGroupSize(zeKernel, groupSizeX, 1, 1); - - uint32_t argIdx = 0; - // Input memrefs - for (int i = 0; i < n_input_memrefs; ++i) { - zeKernelSetArgumentValue(zeKernel, argIdx++, sizeof(T*), &descs[i].allocated); - zeKernelSetArgumentValue(zeKernel, argIdx++, sizeof(T*), &descs[i].aligned); - zeKernelSetArgumentValue(zeKernel, argIdx++, sizeof(int64_t), &descs[i].offset); - for (int d = 0; d < N; ++d) - zeKernelSetArgumentValue(zeKernel, argIdx++, sizeof(int64_t), &descs[i].sizes[d]); - for (int d = 0; d < N; ++d) - zeKernelSetArgumentValue(zeKernel, argIdx++, sizeof(int64_t), &descs[i].strides[d]); - } - // Broadcast scalars - for (int i = 0; i < n_scalar_args; ++i) - zeKernelSetArgumentValue(zeKernel, argIdx++, sizeof(T), scalar_args[i]); - // Output memrefs - for (int i = n_input_memrefs; i < n_memrefs; ++i) { - zeKernelSetArgumentValue(zeKernel, argIdx++, sizeof(T*), &descs[i].allocated); - zeKernelSetArgumentValue(zeKernel, argIdx++, sizeof(T*), &descs[i].aligned); - zeKernelSetArgumentValue(zeKernel, argIdx++, sizeof(int64_t), &descs[i].offset); - for (int d = 0; d < N; ++d) - zeKernelSetArgumentValue(zeKernel, argIdx++, sizeof(int64_t), &descs[i].sizes[d]); - for (int d = 0; d < N; ++d) - zeKernelSetArgumentValue(zeKernel, argIdx++, sizeof(int64_t), &descs[i].strides[d]); - } - - ze_group_count_t dispatchArgs = {(uint32_t)((n_elements + groupSizeX - 1) / groupSizeX), 1, - 1}; - - ze_result_t res = zeInit(ZE_INIT_FLAG_GPU_ONLY); - uint32_t driverCount = 1; - ze_driver_handle_t driver; - zeDriverGet(&driverCount, &driver); - uint32_t deviceCount = 1; - ze_device_handle_t device; - zeDeviceGet(driver, &deviceCount, &device); - ze_context_desc_t ctxDesc = {ZE_STRUCTURE_TYPE_CONTEXT_DESC, nullptr, 0}; - ze_context_handle_t zeContext; - zeContextCreate(driver, &ctxDesc, &zeContext); - - ze_command_queue_desc_t cmdQueueDesc = {}; - cmdQueueDesc.stype = ZE_STRUCTURE_TYPE_COMMAND_QUEUE_DESC; - cmdQueueDesc.mode = ZE_COMMAND_QUEUE_MODE_SYNCHRONOUS; - ze_command_list_handle_t cmdList; - zeCommandListCreateImmediate(zeContext, device, &cmdQueueDesc, &cmdList); - - zeCommandListAppendLaunchKernel(cmdList, zeKernel, &dispatchArgs, nullptr, 0, nullptr); - zeCommandListDestroy(cmdList); - zeContextDestroy(zeContext); - } -#else - using PackedFunc = void (*)(void**); - reinterpret_cast(func_ptr)(args.data()); -#endif -} diff --git a/src/execute_node.cpp b/src/execute_node.cpp deleted file mode 100644 index 2507bc1..0000000 --- a/src/execute_node.cpp +++ /dev/null @@ -1,1599 +0,0 @@ -#include "backend_internal.hpp" -#include "dispatch.hpp" - -#include - -#ifdef USE_KOKKOS -#include -#endif - -/// Collect memrefLeaves, broadcastLeaves, and outputRoots from a DAG node's AST. -static void collect_leaves_and_outputs(DAGNode* node, std::vector& memrefLeaves, - std::vector& outputRoots, - std::vector& broadcastLeaves) { - std::unordered_set internalResults; - for (ASTNode* root : node->ast->roots) - internalResults.insert(root->result_name); - - std::unordered_set seenBroadcasts; - - for (ASTNode* root : node->ast->roots) { - auto root_opc = static_cast(root->opcode); - - for (int op_idx = 0; op_idx < (int)root->operands.size(); ++op_idx) { - ASTNode* operand = root->operands[op_idx]; - if (!operand) - continue; - if (internalResults.count(operand->result_name)) - continue; - if (root_opc == Opcode::SET_REGION && op_idx == 0) - continue; - auto opc = static_cast(operand->opcode); - if ((opc == Opcode::NOOP || opc == Opcode::CREATE) && !operand->is_scalar) { - if (operand->is_broadcast) { - if (seenBroadcasts.insert(operand->result_name).second) - broadcastLeaves.push_back(operand); - } else { - memrefLeaves.push_back(operand); - } - } - } - } - - for (ASTNode* root : node->ast->roots) { - auto opc = static_cast(root->opcode); - if (opc != Opcode::NOOP && opc != Opcode::CREATE && !root->is_temp) - outputRoots.push_back(root); - } -} - -// ---- Unified execute_node_nd ---- - -/// Build a MemRef descriptor for an input array. -template -static MemRef make_input_memref(Array* arr) { - MemRef desc; - T* ptr = static_cast(arr->device_data_ptr()); - desc.allocated = ptr; - desc.aligned = ptr; - // Compute row-major strides and offset - desc.strides[N - 1] = 1LL; - for (int d = N - 2; d >= 0; --d) - desc.strides[d] = desc.strides[d + 1] * (int64_t)arr->region.size(d + 1); - desc.offset = 0; - for (int d = 0; d < N; ++d) { - desc.sizes[d] = (int64_t)arr->region.size(d); - desc.offset += (int64_t)arr->region.start[d] * desc.strides[d]; - } - return desc; -} - -/// Build an output MemRef descriptor. -template -static MemRef make_output_memref(Array* arr) { - MemRef desc; - T* ptr = static_cast(arr->device_data_ptr()); - desc.allocated = ptr; - desc.aligned = ptr; - desc.offset = 0LL; - desc.strides[N - 1] = 1LL; - for (int d = N - 2; d >= 0; --d) - desc.strides[d] = desc.strides[d + 1] * (int64_t)arr->region.size(d + 1); - for (int d = 0; d < N; ++d) - desc.sizes[d] = (int64_t)arr->region.size(d); - return desc; -} - -template -void ArrayDAGGroup::execute_node_nd(DAGNode* node, PartitionImpl* partition, PendingComm* comm) { - auto nd_idx = partition->nd_index(); - int chare_idx = nd_idx[0]; - if constexpr (N >= 2) - chare_idx = chare_idx * partition_grid[N].grid[1] + nd_idx[1]; - if constexpr (N >= 3) - chare_idx = chare_idx * partition_grid[N].grid[2] + nd_idx[2]; - auto& ameta = this->array_meta; - - - // DBG_PRINT("[Chare %d] execute_node_nd<%d>: node_id=%d identifier=%lld fusible=%d comm=%s\n", - // chare_idx, N, node->id, (long long)node->identifier, (int)node->fusible, - // comm ? "yes" : "no"); - - auto it = compile_cache.find(node->identifier); - - std::vector memrefLeaves_check; - std::vector outputRoots_check; - std::vector broadcastLeaves_check; - collect_leaves_and_outputs(node, memrefLeaves_check, outputRoots_check, broadcastLeaves_check); - bool force_interpreter = memrefLeaves_check.empty() && broadcastLeaves_check.empty(); - - if (force_interpreter || it == compile_cache.end() || !it->second) { - // Interpreter fallback - // DBG_PRINT("[Chare %d] -> interpreter fallback\n", chare_idx); - bool handled = false; - if (comm) { - for (auto ast_node : node->ast->roots) { - if (static_cast(ast_node->opcode) == Opcode::MATMUL) { - if constexpr (N == 2) { - // --- 2D MATMUL interpreter with communication data --- - int mat_name = ast_node->operands[0]->result_name; - int vec_name = ast_node->operands[1]->result_name; - int result_name = ast_node->result_name; - - // Get local matrix - auto mat_it = partition->arrays.find(mat_name); - if (mat_it == partition->arrays.end()) { - handled = true; - continue; - } - Array* mat = static_cast*>(mat_it->second); - int local_rows = mat->region.size(0); - int local_cols = mat->region.size(1); - - // Extract slice regions if present - int mat_rs = 0, mat_re = mat->global_shape[0]; - int mat_cs = 0, mat_ce = mat->global_shape[1]; - if (ast_node->operand_regions.size() >= 2) { - auto* mr = static_cast*>(ast_node->operand_regions[0]); - mat_rs = mr->start[0]; mat_re = mr->stop[0]; - mat_cs = mr->start[1]; mat_ce = mr->stop[1]; - } - - // Translate AST regions to global space - auto mat_decomp = mat_it->second->decomp; - mat_rs += mat_decomp.offset[0]; mat_re += mat_decomp.offset[0]; - mat_cs += mat_decomp.offset[1]; mat_ce += mat_decomp.offset[1]; - - // Compute row/col overlap with local tile - auto mat_chare = mat_decomp.chare_region_global(nd_idx); - int row_lo = std::max(mat_chare.start[0], mat_rs); - int row_hi = std::min(mat_chare.stop[0], mat_re); - int col_lo = std::max(mat_chare.start[1], mat_cs); - int col_hi = std::min(mat_chare.stop[1], mat_ce); - int sub_rows = row_hi - row_lo; - int sub_cols = col_hi - col_lo; - int row_offset = row_lo - mat_chare.start[0]; - int col_offset = col_lo - mat_chare.start[1]; - - // Reassemble vector from remote buffers (may be multiple) - T* vec_data = nullptr; - int vec_len = 0; - T* assembled_vec = nullptr; - { - auto rb_it = comm->remote_buffers.find(0); - if (rb_it != comm->remote_buffers.end() && !rb_it->second.empty()) { - auto& buffers = rb_it->second; - if (buffers.size() == 1) { - vec_data = reinterpret_cast(buffers[0].data); - vec_len = buffers[0].byte_size / sizeof(T); - } else { - std::sort(buffers.begin(), buffers.end(), - [](auto& a, auto& b) { - return a.region.start[0] < b.region.start[0]; - }); - for (auto& rb : buffers) - vec_len += rb.byte_size / sizeof(T); - assembled_vec = new T[vec_len]; - int offset = 0; - for (auto& rb : buffers) { - int len = rb.byte_size / sizeof(T); - memcpy(assembled_vec + offset, - reinterpret_cast(rb.data), - rb.byte_size); - offset += len; - } - vec_data = assembled_vec; - } - } - } - - if (!vec_data || sub_rows <= 0 || sub_cols <= 0) { - delete[] assembled_vec; - handled = true; - continue; - } - - int slice_rows = mat_re - mat_rs; - std::array out_start = {}, out_stop, out_step, out_gs; - out_stop[0] = sub_rows; out_step[0] = 1; out_gs[0] = slice_rows; - out_stop[1] = 1; out_step[1] = 1; out_gs[1] = 1; - ArrayRegion out_region(out_start, out_stop, out_step); - ArrayDecomp out_decomp = array_meta[result_name].decomp(); - Array* result = partition->template ensure_array_typed( - out_region, out_gs, result_name, out_decomp); - - int actual_cols = std::min(sub_cols, vec_len); -#ifdef USE_KOKKOS - { - Kokkos::View> - d_mat(static_cast(mat->device_data_ptr()), - local_rows, local_cols); - auto d_sub_mat = Kokkos::subview(d_mat, - Kokkos::make_pair(row_offset, row_offset + sub_rows), - Kokkos::make_pair(col_offset, col_offset + actual_cols)); - Kokkos::View> - d_vec(vec_data, actual_cols); - Kokkos::View> - d_res(static_cast(result->device_data_ptr()), - sub_rows); - KokkosBlas::gemv("N", T(1), d_sub_mat, d_vec, T(0), d_res); - } - CT_KOKKOS_FENCE(); -#else - eigen_gemv_sub(static_cast(mat->data_ptr()), local_cols, - row_offset, col_offset, - sub_rows, actual_cols, - vec_data, static_cast(result->data_ptr())); -#endif - delete[] assembled_vec; - - // Send partial result to 1D partition(s) - { - auto* dag_group = static_cast( - partition->dag_proxy.ckLocalBranch()); - int tile_1d = array_tile(ameta, result_name, 1); - - - // result_lo/hi in result-local space (0-based) - int result_lo = row_lo - mat_rs; - int result_hi = row_hi - mat_rs; - // Translate to global for chare index computation - auto res_meta = dag_group->array_meta.find(result_name); - int res_offset = (res_meta != dag_group->array_meta.end()) - ? res_meta->second.offset[0] - : 0; - int result_lo_g = result_lo + res_offset; - int result_hi_g = result_hi + res_offset; - int first_result_chare = result_lo_g / tile_1d; - int last_result_chare = (result_hi_g - 1) / tile_1d; - - T* result_data = static_cast(result->device_data_ptr()); - for (int rk = first_result_chare; rk <= last_result_chare; rk++) { - // Tile boundaries in result-local space - int send_start = std::max(result_lo, rk * tile_1d - res_offset); - int send_end = std::min(result_hi, (rk + 1) * tile_1d - res_offset); - int send_len = send_end - send_start; - int data_offset = send_start - result_lo; - - ChareIndex<1> target_ci; - target_ci.idx[0] = rk; - cross_matmul_send_result_2d_to_1d( - partition, node, result_name, target_ci, - result_data + data_offset, send_len, send_start, - dag_group->partition_proxy_1); - } - - partition->retire_array(result_name); - } - - handled = true; - continue; - - } else if constexpr (N == 3) { - // --- 3D MATMUL interpreter: dimension-dropped matvec --- - int mat_name = ast_node->operands[0]->result_name; - int vec_name = ast_node->operands[1]->result_name; - int result_name = ast_node->result_name; - - auto mat_it = partition->arrays.find(mat_name); - if (mat_it == partition->arrays.end()) { - handled = true; - continue; - } - Array* mat = static_cast*>(mat_it->second); - int local_d0 = mat->region.size(0); - int local_d1 = mat->region.size(1); - int local_d2 = mat->region.size(2); - - // Extract 3D region: identify dropped/row/col dims - auto* mr = static_cast*>(ast_node->operand_regions[0]); - int dd = -1; - for (int d = 0; d < 3; d++) { - if (mr->stop[d] - mr->start[d] == 1) { - dd = d; break; - } - } - int rd = (dd == 0) ? 1 : 0; - int cd = (dd <= 1) ? 2 : 1; - - int dd_s = mr->start[dd], dd_e = mr->stop[dd]; - int rd_s = mr->start[rd], rd_e = mr->stop[rd]; - int cd_s = mr->start[cd], cd_e = mr->stop[cd]; - - // Translate AST regions to global space - auto mat_decomp = mat_it->second->decomp; - dd_s += mat_decomp.offset[dd]; dd_e += mat_decomp.offset[dd]; - rd_s += mat_decomp.offset[rd]; rd_e += mat_decomp.offset[rd]; - cd_s += mat_decomp.offset[cd]; cd_e += mat_decomp.offset[cd]; - - // Local tile sizes for each dim - int local_dims[3] = {local_d0, local_d1, local_d2}; - - // Compute overlap in each dim using decomp - auto mat_chare = mat_decomp.chare_region_global(nd_idx); - int dd_lo = std::max(mat_chare.start[dd], dd_s); - int dd_hi = std::min(mat_chare.stop[dd], dd_e); - int row_lo = std::max(mat_chare.start[rd], rd_s); - int row_hi = std::min(mat_chare.stop[rd], rd_e); - int col_lo = std::max(mat_chare.start[cd], cd_s); - int col_hi = std::min(mat_chare.stop[cd], cd_e); - - int sub_rows = row_hi - row_lo; - int sub_cols = col_hi - col_lo; - int dd_local_offset = dd_lo - mat_chare.start[dd]; - int row_offset = row_lo - mat_chare.start[rd]; - int col_offset = col_lo - mat_chare.start[cd]; - - // Reassemble vector from remote buffers - T* vec_data = nullptr; - int vec_len = 0; - T* assembled_vec = nullptr; - { - auto rb_it = comm->remote_buffers.find(0); - if (rb_it != comm->remote_buffers.end() && !rb_it->second.empty()) { - auto& buffers = rb_it->second; - if (buffers.size() == 1) { - vec_data = reinterpret_cast(buffers[0].data); - vec_len = buffers[0].byte_size / sizeof(T); - } else { - std::sort(buffers.begin(), buffers.end(), - [](auto& a, auto& b) { - return a.region.start[0] < b.region.start[0]; - }); - for (auto& rb : buffers) - vec_len += rb.byte_size / sizeof(T); - assembled_vec = new T[vec_len]; - int offset = 0; - for (auto& rb : buffers) { - int len = rb.byte_size / sizeof(T); - memcpy(assembled_vec + offset, - reinterpret_cast(rb.data), - rb.byte_size); - offset += len; - } - vec_data = assembled_vec; - } - } - } - - if (!vec_data || sub_rows <= 0 || sub_cols <= 0) { - delete[] assembled_vec; - handled = true; - continue; - } - - int slice_rows = rd_e - rd_s; - - // Create temporary result array for sub_rows - std::array out_start = {}, out_stop = {}, out_step = {}, out_gs = {}; - out_stop[0] = sub_rows; out_step[0] = 1; out_gs[0] = slice_rows; - out_stop[1] = 1; out_step[1] = 1; out_gs[1] = 1; - out_stop[2] = 1; out_step[2] = 1; out_gs[2] = 1; - ArrayRegion out_region(out_start, out_stop, out_step); - ArrayDecomp out_decomp = array_meta[result_name].decomp(); - Array* result = partition->template ensure_array_typed( - out_region, out_gs, result_name, out_decomp); - - int actual_cols = std::min(sub_cols, vec_len); - - // Perform 3D sub-matrix GEMV (zero-copy from 3D tile) -#ifndef USE_KOKKOS - eigen_gemv_sub_3d(static_cast(mat->data_ptr()), - local_d0, local_d1, local_d2, - dd, dd_local_offset, - row_offset, col_offset, - sub_rows, actual_cols, - vec_data, static_cast(result->data_ptr())); -#else - // Kokkos: use 3D subview then reshape to 2D for gemv - { - Kokkos::View> - d_mat(static_cast(mat->device_data_ptr()), - local_d0, local_d1, local_d2); - - // Create subview for the 2D slice - int r0_lo, r0_hi, r1_lo, r1_hi, r2_lo, r2_hi; - if (dd == 0) { - r0_lo = dd_local_offset; r0_hi = dd_local_offset + 1; - r1_lo = row_offset; r1_hi = row_offset + sub_rows; - r2_lo = col_offset; r2_hi = col_offset + actual_cols; - } else if (dd == 1) { - r0_lo = row_offset; r0_hi = row_offset + sub_rows; - r1_lo = dd_local_offset; r1_hi = dd_local_offset + 1; - r2_lo = col_offset; r2_hi = col_offset + actual_cols; - } else { - r0_lo = row_offset; r0_hi = row_offset + sub_rows; - r1_lo = col_offset; r1_hi = col_offset + actual_cols; - r2_lo = dd_local_offset; r2_hi = dd_local_offset + 1; - } - auto d_sub = Kokkos::subview(d_mat, - Kokkos::make_pair(r0_lo, r0_hi), - Kokkos::make_pair(r1_lo, r1_hi), - Kokkos::make_pair(r2_lo, r2_hi)); - - // Manual gemv on the subview - Kokkos::View> - d_vec(vec_data, actual_cols); - Kokkos::View> - d_res(static_cast(result->device_data_ptr()), - sub_rows); - // Flatten the 3D subview to 2D for KokkosBlas - // For dd==0: subview is (1, sub_rows, actual_cols) -> reshape to (sub_rows, actual_cols) - // For other cases, strides may differ; use manual parallel_for - Kokkos::parallel_for( - Kokkos::RangePolicy<>(0, sub_rows), - KOKKOS_LAMBDA(int i) { - T sum = T(0); - for (int j = 0; j < actual_cols; j++) { - if (dd == 0) - sum += d_sub(0, i, j) * d_vec(j); - else if (dd == 1) - sum += d_sub(i, 0, j) * d_vec(j); - else - sum += d_sub(i, j, 0) * d_vec(j); - } - d_res(i) = sum; - }); - } - CT_KOKKOS_FENCE(); -#endif - delete[] assembled_vec; - - // Send partial result to 1D partition(s) - { - auto* dag_group = static_cast( - partition->dag_proxy.ckLocalBranch()); - int tile_1d = array_tile(ameta, result_name, 1); - - - // result_lo/hi in result-local space (0-based) - int result_lo = row_lo - rd_s; - int result_hi = row_hi - rd_s; - // Translate to global for chare index computation - auto res_meta = dag_group->array_meta.find(result_name); - int res_offset = (res_meta != dag_group->array_meta.end()) - ? res_meta->second.offset[0] - : 0; - int result_lo_g = result_lo + res_offset; - int result_hi_g = result_hi + res_offset; - int first_result_chare = result_lo_g / tile_1d; - int last_result_chare = (result_hi_g - 1) / tile_1d; - - T* result_data = static_cast(result->device_data_ptr()); - for (int rk = first_result_chare; rk <= last_result_chare; rk++) { - // Tile boundaries in result-local space - int send_start = std::max(result_lo, rk * tile_1d - res_offset); - int send_end = std::min(result_hi, (rk + 1) * tile_1d - res_offset); - int send_len = send_end - send_start; - int data_offset = send_start - result_lo; - - ChareIndex<1> target_ci; - target_ci.idx[0] = rk; - cross_matmul_send_result_3d_to_1d( - partition, node, result_name, target_ci, - result_data + data_offset, send_len, send_start, - dag_group->partition_proxy_1); - } - - partition->retire_array(result_name); - } - - handled = true; - continue; - } else if constexpr (N == 1) { - // --- 1D partition: accumulate partial results from 2D/3D chares --- - int result_name = ast_node->result_name; - int mat_name = ast_node->operands[0]->result_name; - auto* dag_group = - static_cast(partition->dag_proxy.ckLocalBranch()); - - // Extract slice region for result size (works for 2D or 3D regions) - int slice_rows = 0; - if (ast_node->operand_regions.size() >= 1) { - auto meta_it = dag_group->array_meta.find(mat_name); - int mat_ndims = (meta_it != dag_group->array_meta.end()) ? meta_it->second.ndims : 2; - if (mat_ndims == 3) { - auto* mr = static_cast*>(ast_node->operand_regions[0]); - // Find dropped dim and extract row dim size - int dd = -1; - for (int d = 0; d < 3; d++) { - if (mr->stop[d] - mr->start[d] == 1) { dd = d; break; } - } - int rd = (dd == 0) ? 1 : 0; - slice_rows = mr->stop[rd] - mr->start[rd]; - } else { - auto* mr = static_cast*>(ast_node->operand_regions[0]); - slice_rows = mr->stop[0] - mr->start[0]; - } - } else { - // Fallback: use matrix global rows - auto meta_it = dag_group->array_meta.find(mat_name); - if (meta_it != dag_group->array_meta.end()) - slice_rows = meta_it->second.global_shape[0]; - } - - int tile_1d = array_tile(ameta, result_name, 1); - - int k = partition->nd_index()[0]; - // Decomp-aware result chare start (in result-local space) - auto res_meta_1d = dag_group->array_meta.find(result_name); - ArrayDecomp<1> res_decomp_1d; - if (res_meta_1d != dag_group->array_meta.end()) - res_decomp_1d = res_meta_1d->second.decomp<1>(); - std::array nd_1d = {k}; - auto res_chare_1d = res_decomp_1d.chare_region_local(nd_1d); - int my_result_start = res_chare_1d.start[0]; - int local_result_size = - std::min(res_chare_1d.stop[0] - my_result_start, - slice_rows - my_result_start); - if (local_result_size <= 0) { - handled = true; - continue; - } - - // Create result array - std::array out_start = {0}; - std::array out_stop = {local_result_size}; - std::array out_step = {1}; - std::array out_gs = {slice_rows}; - ArrayRegion<1> out_region(out_start, out_stop, out_step); - ArrayDecomp<1> out_decomp = array_meta[result_name].decomp<1>(); - partition->template ensure_array_typed( - out_region, out_gs, result_name, out_decomp); - - auto rb_it = comm->remote_buffers.find(0); - if (rb_it != comm->remote_buffers.end() && !rb_it->second.empty()) { - auto& buffers = rb_it->second; - T* res_data = static_cast(partition->arrays[result_name]->data_ptr()); - - // Accumulate partials with offset awareness. - // Each partial has a 1D region [start, stop) in output space. - // Result array is zero-initialized, so always accumulate with +=. - for (auto& rb : buffers) { - T* received = reinterpret_cast(rb.data); - int partial_start = rb.region.start[0]; - int partial_size = rb.byte_size / sizeof(T); - int local_offset = partial_start - my_result_start; -#ifdef USE_KOKKOS - Kokkos::View> - d_partial(received, partial_size); - Kokkos::View> - d_res(static_cast( - partition->arrays[result_name]->device_data_ptr()) + local_offset, - partial_size); - KokkosBlas::axpy(T(1), d_partial, d_res); -#else - Eigen::Map> - eigen_res(res_data + local_offset, partial_size); - Eigen::Map> - eigen_partial(received, partial_size); - eigen_res += eigen_partial; -#endif - } - DBG_PRINT("[1D Chare %d] MATMUL: accumulated %d partials (%d elements)\n", - chare_idx, (int)buffers.size(), local_result_size); - } - handled = true; - continue; - } else { - CkAbort("MATMUL only supported for N >= 2"); - } - } - if (static_cast(ast_node->opcode) == Opcode::MATMATMUL) { - if constexpr (N == 2) { - // Incremental mode: all computation was done in on_matmatmul_receive - // as panels arrived. C array was allocated in execute_matmatmul_node. - // Nothing to do here — just mark as handled. - DBG_PRINT("[2D Chare %d] MATMATMUL: incremental done\n", chare_idx); - handled = true; - continue; - } else { - CkAbort("MATMATMUL only supported for N == 2"); - } - } - if (static_cast(ast_node->opcode) == Opcode::DIAG) { - int input_name = ast_node->operands[0]->result_name; - int result_name = ast_node->result_name; - int k_offset = 0; - if (ast_node->operands.size() >= 2 && ast_node->operands[1]->is_scalar) - k_offset = (int)ast_node->operands[1]->scalar; - - auto* dag_group = static_cast( - partition->dag_proxy.ckLocalBranch()); - auto result_meta_it = dag_group->array_meta.find(result_name); - - if constexpr (N == 2) { - // 1D → 2D: construct diagonal matrix from received vector elements - int out_n = result_meta_it->second.global_shape[0]; - auto result_decomp = result_meta_it->second.decomp<2>(); - auto result_chare = result_decomp.chare_region_global(nd_idx); - int local_rows = result_chare.stop[0] - result_chare.start[0]; - int local_cols = result_chare.stop[1] - result_chare.start[1]; - - // Create zero-initialized output array - std::array out_start = {0, 0}; - std::array out_stop = {local_rows, local_cols}; - std::array out_step = {1, 1}; - std::array out_gs = {out_n, out_n}; - ArrayRegion<2> out_region(out_start, out_stop, out_step); - partition->template ensure_array_typed( - out_region, out_gs, result_name, result_decomp); - int row_start = result_chare.start[0]; - int col_start = result_chare.start[1]; - - // Place diagonal elements from remote buffers - // Array is already zero-initialized by constructor -#ifdef USE_KOKKOS - // Device-side placement: rb.data is a device pointer - if (comm) { - auto rb_it = comm->remote_buffers.find(0); - if (rb_it != comm->remote_buffers.end()) { - for (auto& rb : rb_it->second) { - T* received = reinterpret_cast(rb.data); - T* out_device = static_cast( - partition->arrays[result_name]->device_data_ptr()); - int rb_row_start = rb.region.start[0]; - int rb_col_start = rb.region.start[1]; - int rb_count = rb.byte_size / sizeof(T); - int rs_cap = row_start, cs_cap = col_start; - int lr_cap = local_rows, lc_cap = local_cols; - Kokkos::parallel_for( - CT_COMPUTE_POLICY(partition, rb_count), - KOKKOS_LAMBDA(int j) { - int row = rb_row_start + j; - int col = rb_col_start + j; - int local_r = row - rs_cap; - int local_c = col - cs_cap; - if (local_r >= 0 && local_r < lr_cap && - local_c >= 0 && local_c < lc_cap) { - out_device[local_r * lc_cap + local_c] = - received[j]; - } - }); - } - } - } -#else - T* out_data = static_cast(partition->arrays[result_name]->data_ptr()); - std::memset(out_data, 0, local_rows * local_cols * sizeof(T)); - - if (comm) { - auto rb_it = comm->remote_buffers.find(0); - if (rb_it != comm->remote_buffers.end()) { - for (auto& rb : rb_it->second) { - T* received = reinterpret_cast(rb.data); - int rb_row_start = rb.region.start[0]; - int rb_col_start = rb.region.start[1]; - int rb_count = rb.byte_size / sizeof(T); - for (int j = 0; j < rb_count; ++j) { - int row = rb_row_start + j; - int col = rb_col_start + j; - int local_r = row - row_start; - int local_c = col - col_start; - if (local_r >= 0 && local_r < local_rows && - local_c >= 0 && local_c < local_cols) { - out_data[local_r * local_cols + local_c] = received[j]; - } - } - } - } - } -#endif - DBG_PRINT("[2D Chare %d] DIAG: 1D->2D construct done\n", chare_idx); - handled = true; - continue; - } else if constexpr (N == 1) { - // 2D → 1D: place received diagonal elements into output - int diag_len = result_meta_it->second.global_shape[0]; - auto result_decomp_1d = result_meta_it->second.decomp<1>(); - auto result_chare_1d = result_decomp_1d.chare_region_global(nd_idx); - int my_start = result_chare_1d.start[0]; - int local_size = result_chare_1d.stop[0] - my_start; - - std::array out_start = {0}; - std::array out_stop = {local_size}; - std::array out_step = {1}; - std::array out_gs = {diag_len}; - ArrayRegion<1> out_region(out_start, out_stop, out_step); - partition->template ensure_array_typed( - out_region, out_gs, result_name, result_decomp_1d); - // Array is already zero-initialized by constructor -#ifdef USE_KOKKOS - // Device-side placement: rb.data is a device pointer - if (comm) { - auto rb_it = comm->remote_buffers.find(0); - if (rb_it != comm->remote_buffers.end()) { - for (auto& rb : rb_it->second) { - T* received = reinterpret_cast(rb.data); - T* out_device = static_cast( - partition->arrays[result_name]->device_data_ptr()); - int rb_start = rb.region.start[0]; - int rb_count = rb.byte_size / sizeof(T); - int ms_cap = my_start, ls_cap = local_size; - Kokkos::parallel_for( - CT_COMPUTE_POLICY(partition, rb_count), - KOKKOS_LAMBDA(int j) { - int local_idx = (rb_start + j) - ms_cap; - if (local_idx >= 0 && local_idx < ls_cap) - out_device[local_idx] = received[j]; - }); - } - } - } -#else - T* out_data = static_cast(partition->arrays[result_name]->data_ptr()); - std::memset(out_data, 0, local_size * sizeof(T)); - - if (comm) { - auto rb_it = comm->remote_buffers.find(0); - if (rb_it != comm->remote_buffers.end()) { - for (auto& rb : rb_it->second) { - T* received = reinterpret_cast(rb.data); - int rb_start = rb.region.start[0]; - int rb_count = rb.byte_size / sizeof(T); - for (int j = 0; j < rb_count; ++j) { - int local_idx = (rb_start + j) - my_start; - if (local_idx >= 0 && local_idx < local_size) { - out_data[local_idx] = received[j]; - } - } - } - } - } -#endif - DBG_PRINT("[1D Chare %d] DIAG: 2D->1D extract done\n", chare_idx); - handled = true; - continue; - } - } - if (static_cast(ast_node->opcode) == Opcode::TILE) { - int input_name = ast_node->operands[0]->result_name; - int result_name = ast_node->result_name; - int out_ndims = ast_node->ndims; - - // Extract reps - std::array reps = {1, 1, 1}; - int num_reps = (int)ast_node->operands.size() - 1; - for (int d = 0; d < num_reps && d < 3; d++) { - if (ast_node->operands[d + 1]->is_scalar) - reps[d] = (int)ast_node->operands[d + 1]->scalar; - } - - auto* dag_group = static_cast( - partition->dag_proxy.ckLocalBranch()); - auto result_meta_it = dag_group->array_meta.find(result_name); - auto input_meta_it = dag_group->array_meta.find(input_name); - - if (result_meta_it != dag_group->array_meta.end() && - input_meta_it != dag_group->array_meta.end() && - N == out_ndims) { - auto result_decomp = result_meta_it->second.template decomp(); - auto result_chare = result_decomp.chare_region_global(nd_idx); - int input_ndims = input_meta_it->second.ndims; - int delta = out_ndims - input_ndims; - - // Compute local sizes - std::array local_sizes; - std::array out_gs; - for (int d = 0; d < N; ++d) { - local_sizes[d] = result_chare.stop[d] - result_chare.start[d]; - out_gs[d] = result_meta_it->second.global_shape[d]; - } - - // Create output array if not yet allocated - std::array out_start, out_stop, out_step; - for (int d = 0; d < N; ++d) { - out_start[d] = 0; - out_stop[d] = local_sizes[d]; - out_step[d] = 1; - } - ArrayRegion out_region(out_start, out_stop, out_step); - partition->template ensure_array_typed( - out_region, out_gs, result_name, result_decomp); - - // Get input shape (padded to out_ndims) - std::array input_shape = input_meta_it->second.global_shape; - - // Place received data using modular index mapping - // Each remote buffer contains source data from one source chare. - // The buffer's region encodes the source's global coords (with delta padding). -#ifndef USE_KOKKOS - T* out_data = static_cast(partition->arrays[result_name]->data_ptr()); - - if (comm) { - auto rb_it = comm->remote_buffers.find(0); - if (rb_it != comm->remote_buffers.end()) { - for (auto& rb : rb_it->second) { - T* received = reinterpret_cast(rb.data); - // rb.region encodes source global coords in N_tgt dims - // For dims < delta: [0, 1) (prepended dim) - // For dims >= delta: source chare's global range - std::array src_start, src_size; - int src_total_size = 1; - for (int d = 0; d < N; ++d) { - src_start[d] = rb.region.start[d]; - src_size[d] = rb.region.stop[d] - rb.region.start[d]; - src_total_size *= src_size[d]; - } - - // For each output position in this chare's local region, - // check if it maps (via mod) to this source's range - int total_out = 1; - for (int d = 0; d < N; ++d) - total_out *= local_sizes[d]; - - for (int flat = 0; flat < total_out; ++flat) { - // Decompose flat index into N-dimensional local coords - std::array local_idx; - int rem = flat; - for (int d = N - 1; d >= 0; --d) { - local_idx[d] = rem % local_sizes[d]; - rem /= local_sizes[d]; - } - - // Convert to global output coords - std::array global_out; - for (int d = 0; d < N; ++d) - global_out[d] = result_chare.start[d] + local_idx[d]; - - // Map to input coords via mod - bool from_this_src = true; - std::array src_local; - for (int d = 0; d < N; ++d) { - int in_size; - if (d < delta) - in_size = 1; - else - in_size = input_shape[d - delta]; - int inp_pos = global_out[d] % in_size; - if (inp_pos < src_start[d] || - inp_pos >= src_start[d] + src_size[d]) { - from_this_src = false; - break; - } - src_local[d] = inp_pos - src_start[d]; - } - if (!from_this_src) - continue; - - // Compute flat index into received buffer - int src_flat = 0; - for (int d = 0; d < N; ++d) { - src_flat = src_flat * src_size[d] + src_local[d]; - } - - out_data[flat] = received[src_flat]; - } - } - } - } -#else - if (comm) { - auto rb_it = comm->remote_buffers.find(0); - if (rb_it != comm->remote_buffers.end()) { - for (auto& rb : rb_it->second) { - T* received = reinterpret_cast(rb.data); - T* out_device = static_cast( - partition->arrays[result_name]->device_data_ptr()); - - std::array src_start_arr, src_size_arr; - for (int d = 0; d < N; ++d) { - src_start_arr[d] = rb.region.start[d]; - src_size_arr[d] = rb.region.stop[d] - rb.region.start[d]; - } - - int total_out = 1; - for (int d = 0; d < N; ++d) - total_out *= local_sizes[d]; - - // Capture arrays for lambda - int lsz[N], ssz[N], sst[N], rcs[N], ishp[3]; - for (int d = 0; d < N; ++d) { - lsz[d] = local_sizes[d]; - ssz[d] = src_size_arr[d]; - sst[d] = src_start_arr[d]; - rcs[d] = result_chare.start[d]; - } - for (int d = 0; d < 3; ++d) - ishp[d] = input_shape[d]; - int delta_cap = delta; - - Kokkos::parallel_for( - CT_COMPUTE_POLICY(partition, total_out), - KOKKOS_LAMBDA(int flat) { - int local_idx[N]; - int r = flat; - for (int d = N - 1; d >= 0; --d) { - local_idx[d] = r % lsz[d]; - r /= lsz[d]; - } - - bool from_this_src = true; - int src_local[N]; - for (int d = 0; d < N; ++d) { - int global_out = rcs[d] + local_idx[d]; - int in_size = (d < delta_cap) ? 1 : ishp[d - delta_cap]; - int inp_pos = global_out % in_size; - if (inp_pos < sst[d] || inp_pos >= sst[d] + ssz[d]) { - from_this_src = false; - break; - } - src_local[d] = inp_pos - sst[d]; - } - if (!from_this_src) - return; - - int src_flat = 0; - for (int d = 0; d < N; ++d) - src_flat = src_flat * ssz[d] + src_local[d]; - - out_device[flat] = received[src_flat]; - }); - } - } - } -#endif - DBG_PRINT("[%dD Chare %d] TILE: assembly done\n", N, chare_idx); - handled = true; - continue; - } - } - if (static_cast(ast_node->opcode) == Opcode::SET_REGION) { - int target_name = ast_node->operands[0]->result_name; - int source_name = ast_node->operands[1]->result_name; - auto* region = static_cast*>(ast_node->region); - - auto tgt_it = partition->arrays.find(target_name); - if (tgt_it == partition->arrays.end()) { - // Target array not on this chare — nothing to write - handled = true; - continue; - } - Array* target = - static_cast*>(tgt_it->second); - auto r_chare_g = target->decomp.chare_region_global(nd_idx); - auto dst_region_global = target->decomp.to_global(*region); - - // Use the target array's decomp-global coordinates before - // intersecting with the owning chare. Without this shift, - // any nonzero target offset makes same-epoch SET_REGION - // writes land on the wrong cells even though the source - // communication and fragment kernels are otherwise correct. - auto [dst_global, has_dst_overlap] = - intersect(dst_region_global, r_chare_g); - if (!has_dst_overlap) { - handled = true; - continue; - } - - if (!comm->my_inputs.empty()) { - auto& my_inp = comm->my_inputs[0]; - - // Compute row-major strides for target - std::array tgt_strides; - tgt_strides[N - 1] = 1; - for (int d = N - 2; d >= 0; --d) - tgt_strides[d] = tgt_strides[d + 1] * target->region.size(d + 1); - - // Chare starts for target and source arrays - std::array cs_tgt, cs_src; - for (int d = 0; d < N; ++d) - cs_tgt[d] = r_chare_g.start[d]; - - // Source array's chare region for local intersection - auto src_it_tmp = partition->arrays.find(source_name); - ArrayRegion r_chare_src; - if (src_it_tmp != partition->arrays.end()) { - r_chare_src = src_it_tmp->second->decomp.chare_region_global(nd_idx); - for (int d = 0; d < N; ++d) - cs_src[d] = src_it_tmp->second->decomp.chare_start_global(d, nd_idx[d]); - } else { - r_chare_src = r_chare_g; - cs_src = cs_tgt; - } - - // Phase 1: copy remote buffers - for (auto& [ridx, rbufs] : comm->remote_buffers) { - for (auto& rb : rbufs) { - // Map the received source fragment back into this chare's - // destination slice using the same region algebra as the - // communication planner. The previous hand-rolled arithmetic - // was fragile for shifted decompositions. - ArrayRegion rb_in_dst = map(rb.region, my_inp, dst_global); - auto [overlap, has_overlap] = intersect(rb_in_dst, dst_global); - if (!has_overlap) - continue; - - // Compute row-major strides for remote buffer (densely packed) - std::array rb_strides; - rb_strides[N - 1] = 1; - for (int d = N - 2; d >= 0; --d) - rb_strides[d] = rb_strides[d + 1] * rb.region.size(d + 1); - -#ifdef USE_KOKKOS - // Device-side copy: remote buffer and target are both on device - T* tgt_device = static_cast(target->device_data_ptr()); - T* rb_device = reinterpret_cast(rb.data); - int64_t overlap_total = overlap.size(); - - int cs_arr[N], tgt_s[N], rb_s[N]; - int ol_start[N], ol_sizes[N], ol_step[N], rb_in_dst_start[N], - rb_in_dst_step[N]; - for (int d = 0; d < N; d++) { - cs_arr[d] = cs_tgt[d]; - tgt_s[d] = tgt_strides[d]; - rb_s[d] = rb_strides[d]; - ol_start[d] = overlap.start[d]; - ol_sizes[d] = overlap.size(d); - ol_step[d] = overlap.step[d]; - rb_in_dst_start[d] = rb_in_dst.start[d]; - rb_in_dst_step[d] = rb_in_dst.step[d]; - } - - Kokkos::parallel_for( - CT_COMPUTE_POLICY(partition, overlap_total), - KOKKOS_LAMBDA(int flat_idx) { - int remaining = flat_idx; - int dst_flat = 0, rb_flat = 0; - for (int d = N - 1; d >= 0; --d) { - int coord_d = remaining % ol_sizes[d]; - remaining /= ol_sizes[d]; - int global_d = ol_start[d] + coord_d * ol_step[d]; - dst_flat += (global_d - cs_arr[d]) * tgt_s[d]; - int rb_logical = - (global_d - rb_in_dst_start[d]) / rb_in_dst_step[d]; - rb_flat += rb_logical * rb_s[d]; - } - tgt_device[dst_flat] = rb_device[rb_flat]; - }); -#else - // Host-side odometer over overlap (step-aware) - std::array idx; - for (int d = 0; d < N; ++d) - idx[d] = overlap.start[d]; - while (true) { - int dst_flat = 0, rb_flat = 0; - for (int d = 0; d < N; ++d) { - dst_flat += (idx[d] - cs_tgt[d]) * tgt_strides[d]; - int rb_logical = - (idx[d] - rb_in_dst.start[d]) / rb_in_dst.step[d]; - rb_flat += rb_logical * rb_strides[d]; - } - static_cast(target->data_ptr())[dst_flat] = - reinterpret_cast(rb.data)[rb_flat]; - - int d = N - 1; - while (d >= 0) { - idx[d] += overlap.step[d]; - if (idx[d] < overlap.stop[d]) - break; - idx[d] = overlap.start[d]; - --d; - } - if (d < 0) - break; - } -#endif - } - } - - // Phase 2: copy local source - auto [local_inp, has_local] = intersect(my_inp, r_chare_src); - if (has_local) { - ArrayRegion local_in_dst = map(local_inp, my_inp, dst_global); - auto [loc_overlap, has_loc_overlap] = - intersect(local_in_dst, dst_global); - if (has_loc_overlap) { - auto src_it = partition->arrays.find(source_name); - if (src_it != partition->arrays.end()) { - // Compute row-major strides for source - std::array src_strides; - src_strides[N - 1] = 1; - for (int d = N - 2; d >= 0; --d) - src_strides[d] = - src_strides[d + 1] * src_it->second->region.size(d + 1); - -#ifdef USE_KOKKOS - // Device-side copy: local source and target both on device - T* tgt_device = static_cast(target->device_data_ptr()); - T* src_device = - static_cast(src_it->second->device_data_ptr()); - int64_t loc_total = loc_overlap.size(); - - int cs_tgt_arr[N], cs_src_arr[N], tgt_s[N], src_s[N]; - int lo_start[N], lo_sizes[N], lo_step[N]; - int local_inp_start[N], local_inp_step[N]; - int local_in_dst_start[N], local_in_dst_step[N]; - for (int d = 0; d < N; d++) { - cs_tgt_arr[d] = cs_tgt[d]; - cs_src_arr[d] = cs_src[d]; - tgt_s[d] = tgt_strides[d]; - src_s[d] = src_strides[d]; - lo_start[d] = loc_overlap.start[d]; - lo_sizes[d] = loc_overlap.size(d); - lo_step[d] = loc_overlap.step[d]; - local_inp_start[d] = local_inp.start[d]; - local_inp_step[d] = local_inp.step[d]; - local_in_dst_start[d] = local_in_dst.start[d]; - local_in_dst_step[d] = local_in_dst.step[d]; - } - - Kokkos::parallel_for( - CT_COMPUTE_POLICY(partition, loc_total), - KOKKOS_LAMBDA(int flat_idx) { - int remaining = flat_idx; - int dst_flat = 0, src_flat = 0; - for (int d = N - 1; d >= 0; --d) { - int coord_d = remaining % lo_sizes[d]; - remaining /= lo_sizes[d]; - int global_d = lo_start[d] + coord_d * lo_step[d]; - dst_flat += (global_d - cs_tgt_arr[d]) * tgt_s[d]; - int logical_idx = - (global_d - local_in_dst_start[d]) / - local_in_dst_step[d]; - int local_src_d = local_inp_start[d] + - logical_idx * local_inp_step[d] - - cs_src_arr[d]; - src_flat += local_src_d * src_s[d]; - } - tgt_device[dst_flat] = src_device[src_flat]; - }); -#else - // Host-side odometer over loc_overlap (step-aware) - std::array idx; - for (int d = 0; d < N; ++d) - idx[d] = loc_overlap.start[d]; - while (true) { - int dst_flat = 0, src_flat = 0; - for (int d = 0; d < N; ++d) { - dst_flat += (idx[d] - cs_tgt[d]) * tgt_strides[d]; - int logical_idx = - (idx[d] - local_in_dst.start[d]) / - local_in_dst.step[d]; - int local_src_d = local_inp.start[d] + - logical_idx * local_inp.step[d] - - cs_src[d]; - src_flat += local_src_d * src_strides[d]; - } - static_cast(target->data_ptr())[dst_flat] = - static_cast(src_it->second->data_ptr())[src_flat]; - - int d = N - 1; - while (d >= 0) { - idx[d] += loc_overlap.step[d]; - if (idx[d] < loc_overlap.stop[d]) - break; - idx[d] = loc_overlap.start[d]; - --d; - } - if (d < 0) - break; - } -#endif - } - } - } - handled = true; - } - } else { - partition->executor->ast_visitor(ast_node, dtype_of()); - } - } - } - if (!handled) - for (auto ast_node : node->ast->roots) - partition->executor->ast_visitor(ast_node, dtype_of()); - return; - } - - // ---- JIT path ---- - DBG_PRINT("[Chare %d] -> JIT path\n", chare_idx); - void* func_ptr = it->second; - - std::vector memrefLeaves; - std::vector outputRoots; - std::vector broadcastLeaves; - collect_leaves_and_outputs(node, memrefLeaves, outputRoots, broadcastLeaves); - - int n_inputs = (int)memrefLeaves.size(); - int n_outputs = (int)outputRoots.size(); - int n_broadcast = (int)broadcastLeaves.size(); - bool has_comm = comm && !comm->my_inputs.empty(); - - // Load broadcast scalar values from local arrays or received buffers. - // Each broadcast leaf is a size-1 array; we load its single element as a T scalar. - std::vector broadcast_values(n_broadcast); - std::vector broadcast_ptrs(n_broadcast); - for (int i = 0; i < n_broadcast; ++i) { - int bname = broadcastLeaves[i]->result_name; - auto arr_it = partition->arrays.find(bname); - if (arr_it != partition->arrays.end()) { -#ifdef USE_KOKKOS - Kokkos::deep_copy( - Kokkos::View>( - &broadcast_values[i]), - Kokkos::View>( - static_cast(arr_it->second->device_data_ptr()))); -#else - arr_it->second->copyToHost(); - broadcast_values[i] = *static_cast(arr_it->second->data_ptr()); -#endif - } else if (comm) { - // Look for broadcast value in received remote buffers - // Convention: broadcast operands use input_index = memref count + broadcast index - int bcast_input_idx = n_inputs + i; - auto rb_it = comm->remote_buffers.find(bcast_input_idx); - if (rb_it != comm->remote_buffers.end() && !rb_it->second.empty()) { -#ifdef USE_KOKKOS - // rb.data is a device pointer under USE_KOKKOS - Kokkos::deep_copy( - Kokkos::View>( - &broadcast_values[i]), - Kokkos::View>( - reinterpret_cast(rb_it->second[0].data))); -#else - broadcast_values[i] = *reinterpret_cast(rb_it->second[0].data); -#endif - } - } - broadcast_ptrs[i] = &broadcast_values[i]; - DBG_PRINT("[Chare %d] broadcast[%d] name=%d value=%f\n", - chare_idx, i, bname, (double)broadcast_values[i]); - } - - // For the fragment path, we need the AST regions. extract_regions_nd - // returns array-local coordinates; convert them to global space once we - // know the output decomp so fragment/output mapping stays aligned with the - // communication planner. - ArrayRegion r_out; - std::vector> tmp_inp; - bool has_regions = false; - if (has_comm) { - std::vector tmp_names; - std::array tmp_gs; - has_regions = extract_regions_nd(node, partition->arrays, r_out, tmp_inp, tmp_names, tmp_gs); - } - - if (!has_comm) { - // ---- Fast path: no region operations ---- - DBG_PRINT("[Chare %d] -> fast path (no comm)\n", chare_idx); - int n_memrefs = n_inputs + n_outputs; - std::vector> descs(n_memrefs); - - bool missing_input = false; - for (int i = 0; i < n_inputs; ++i) { - int arr_name = memrefLeaves[i]->result_name; - auto arr_it = partition->arrays.find(arr_name); - if (arr_it == partition->arrays.end()) { - // Input array not on this chare — skip this node - DBG_PRINT("[Chare %d] fast-path input %d: array name=%d not on this chare, skipping (node_id=%d)\n", - chare_idx, i, arr_name, node->id); - missing_input = true; - break; - } - descs[i] = make_input_memref(static_cast*>(arr_it->second)); - } - if (missing_input) - return; - - // Determine output size from first input - std::array out_sizes = {}; - std::array out_gs = {}; - if (n_inputs > 0) { - for (int d = 0; d < N; ++d) - out_sizes[d] = descs[0].sizes[d]; - int first_name = memrefLeaves[0]->result_name; - out_gs = partition->arrays[first_name]->global_shape; - } else if (n_broadcast > 0) { - // Pure scalar-scalar op: output is a single element - for (int d = 0; d < N; ++d) { - out_sizes[d] = 1; - out_gs[d] = 1; - } - } - - DBG_PRINT("[Chare %d] fast-path: n_inputs=%d n_outputs=%d n_broadcast=%d node_id=%d\n", - chare_idx, n_inputs, n_outputs, n_broadcast, node->id); - for (int i = 0; i < n_outputs; ++i) { - int result_name = outputRoots[i]->result_name; - DBG_PRINT("[Chare %d] output[%d] result_name=%d is_temp=%d (node_id=%d)\n", - chare_idx, i, result_name, (int)outputRoots[i]->is_temp, node->id); - std::array out_start = {}; - std::array out_stop, out_step; - for (int d = 0; d < N; ++d) { - out_stop[d] = (int)out_sizes[d]; - out_step[d] = 1; - } - ArrayRegion out_region(out_start, out_stop, out_step); - ArrayDecomp out_decomp_alloc; - { - int lookup_name = result_name; - if (static_cast(outputRoots[i]->opcode) == Opcode::SET_REGION && - !outputRoots[i]->operands.empty()) - lookup_name = outputRoots[i]->operands[0]->result_name; - auto meta_it = array_meta.find(lookup_name); - if (meta_it != array_meta.end()) - out_decomp_alloc = meta_it->second.template decomp(); - } - Array* output = partition->template ensure_array_typed( - out_region, out_gs, result_name, out_decomp_alloc); - descs[n_inputs + i] = make_output_memref(output); - } - - dispatch_kernel(func_ptr, n_memrefs, descs, n_inputs, broadcast_ptrs -#ifdef USE_NVIDIA - , - partition->compute_stream_raw -#endif - ); - return; - } - - // ---- Fragment path: align local input regions + remote buffers ---- - DBG_PRINT("[Chare %d] -> fragment path (with comm)\n", chare_idx); - - // Output decomp: use result array's decomp for output chare region - ArrayDecomp out_decomp; - { - auto root_output_name = [](ASTNode* root) { - if (root == nullptr) - return -1; - if (static_cast(root->opcode) == Opcode::SET_REGION && - !root->operands.empty() && root->operands[0] != nullptr) - return root->operands[0]->result_name; - return root->result_name; - }; - - int result_name = root_output_name(outputRoots[0]); - auto arr_it = partition->arrays.find(result_name); - if (arr_it != partition->arrays.end()) - out_decomp = arr_it->second->decomp; - else { - auto meta_it = array_meta.find(result_name); - if (meta_it != array_meta.end()) - out_decomp = meta_it->second.template decomp(); - else - out_decomp = ArrayDecomp::default_decomp( - partition->arrays.begin()->second->global_shape, - array_tile(array_meta, result_name, N)); - } - } - if (has_regions) - r_out = out_decomp.to_global(r_out); - std::array cs_out; - for (int d = 0; d < N; ++d) - cs_out[d] = out_decomp.chare_start_global(d, nd_idx[d]); - ArrayRegion r_chare_out = out_decomp.chare_region_global(nd_idx); - // Global output sub-region this chare owns (same coordinate system as r_out). - auto [r_myout_global, has_r_myout_global] = intersect(r_out, r_chare_out); - - std::vector>> input_frag_data(n_inputs); - std::vector> parents; - - for (int i = 0; i < n_inputs; ++i) { - int source_name = memrefLeaves[i]->result_name; - auto src_it = partition->arrays.find(source_name); - - auto& my_inp = comm->my_inputs[i]; - - if (src_it != partition->arrays.end()) { - Array* arr = static_cast*>(src_it->second); - - // Use this input array's decomp for its chare region - ArrayRegion r_chare_inp = arr->decomp.chare_region_global(nd_idx); - std::array cs_inp; - for (int d = 0; d < N; ++d) - cs_inp[d] = arr->decomp.chare_start_global(d, nd_idx[d]); - - auto [local_part, has_local] = intersect(my_inp, r_chare_inp); - if (has_local) { - { - std::array arr_strides; - arr_strides[N - 1] = 1; - for (int d = N - 2; d >= 0; --d) - arr_strides[d] = arr_strides[d + 1] * arr->region.size(d + 1); - - int64_t flat_offset = 0; - for (int d = 0; d < N; ++d) - flat_offset += (int64_t)(local_part.start[d] - cs_inp[d]) * arr_strides[d]; - - T* local_ptr = static_cast(arr->device_data_ptr()) + flat_offset; - FragmentData fd; - fd.region = local_part; - fd.data = local_ptr; - for (int d = 0; d < N; ++d) - fd.src_strides[d] = arr_strides[d]; - input_frag_data[i].push_back(fd); - } - } - } - - auto rb_it = comm->remote_buffers.find(i); - if (rb_it != comm->remote_buffers.end()) - for (auto& rb : rb_it->second) - input_frag_data[i].push_back({rb.region, reinterpret_cast(rb.data)}); - - parents.push_back(my_inp); - } - - auto aligned = align_fragments(input_frag_data, parents); - int M = aligned.empty() ? 0 : (int)aligned[0].size(); - - { - std::array first_input_start = {}; - if (comm && !comm->my_inputs.empty()) - first_input_start = comm->my_inputs.front().start; - - for (int i = 0; i < n_outputs; ++i) { - int result_name = outputRoots[i]->result_name; - if (partition->arrays.find(result_name) == partition->arrays.end()) { - std::array out_gs, local_out, zeros = {}, ones; - bool valid = true; - for (int d = 0; d < N; ++d) { - out_gs[d] = r_out.size(d); - // Size this chare's output tile by intersect(r_out, chare tile), not - // min(tile, |r_out| - nd*tile). The latter assumes r_out starts at global 0; - // for slices like [16:48) on chare 2 it gives 0 and skips allocation → null - // out pointer and segfault in the fragment loop. - if (has_r_myout_global) - local_out[d] = r_myout_global.size(d); - else - local_out[d] = 0; - if (local_out[d] <= 0) - valid = false; - ones[d] = 1; - } - if (!valid) - continue; - ArrayRegion out_region(zeros, local_out, ones); - ArrayDecomp result_decomp; - { - int lookup_name = result_name; - if (static_cast(outputRoots[i]->opcode) == Opcode::SET_REGION && - !outputRoots[i]->operands.empty()) - lookup_name = outputRoots[i]->operands[0]->result_name; - auto meta_it = array_meta.find(lookup_name); - if (meta_it != array_meta.end()) - result_decomp = meta_it->second.template decomp(); - else - result_decomp = out_decomp; - } - partition->template ensure_array_typed( - out_region, out_gs, result_name, result_decomp); - } - } - - for (int m = 0; m < M; ++m) { - int n_memrefs = n_inputs + n_outputs; - std::vector> descs(n_memrefs); - - for (int i = 0; i < n_inputs; ++i) - descs[i] = aligned[i][m].memref; - - for (int i = 0; i < n_outputs; ++i) { - auto out_it = partition->arrays.find(outputRoots[i]->result_name); - if (out_it == partition->arrays.end() || out_it->second == nullptr) - CkAbort("[Chare %d] fragment path: missing output array name=%d node=%d", - chare_idx, outputRoots[i]->result_name, node->id); - Array* out = static_cast*>(out_it->second); - int j = n_inputs + i; - - std::array out_strides; - out_strides[N - 1] = 1; - for (int d = N - 2; d >= 0; --d) - out_strides[d] = out_strides[d + 1] * (int64_t)out->region.size(d + 1); - - int64_t flat_offset = 0; - if (n_inputs > 0 && has_r_myout_global) { - ArrayRegion out_corner = - map(aligned[0][m].region, comm->my_inputs[0], r_myout_global); - for (int d = 0; d < N; ++d) { - // Fragment outputs are written into a packed local buffer - // whose origin is the start of this chare's owned output - // slice, not the start of the global chare tile. Using - // raw global coordinates here misplaces strided SET_REGION - // outputs such as prolongation writes. - int64_t out_local = - ((int64_t)out_corner.start[d] - (int64_t)r_myout_global.start[d]) / - (int64_t)r_myout_global.step[d]; - flat_offset += out_local * out_strides[d]; - } - } else { - for (int d = 0; d < N; ++d) - flat_offset += - ((int64_t)(aligned[0][m].region.start[d] - first_input_start[d]) - - (int64_t)cs_out[d]) * - out_strides[d]; - } - - T* ptr = static_cast(out->device_data_ptr()) + flat_offset; - MemRef desc; - desc.allocated = ptr; - desc.aligned = ptr; - desc.offset = 0LL; - for (int d = 0; d < N; ++d) - desc.sizes[d] = aligned[0][m].memref.sizes[d]; - desc.strides[N - 1] = 1LL; - for (int d = N - 2; d >= 0; --d) - desc.strides[d] = desc.strides[d + 1] * (int64_t)out->region.size(d + 1); - descs[j] = desc; - } - - dispatch_kernel(func_ptr, n_memrefs, descs, n_inputs, broadcast_ptrs -#ifdef USE_NVIDIA - , - partition->compute_stream_raw -#endif - ); - - // Debug: verify input and output data for this fragment - DBG_PRINT("[Chare %d] fragment %d/%d: n_inputs=%d n_outputs=%d\n", - chare_idx, m, M, n_inputs, n_outputs); - for (int i = 0; i < n_inputs; ++i) { - T first_val = T(0); - if (descs[i].sizes[0] > 0 && (N < 2 || descs[i].sizes[1] > 0)) - first_val = descs[i].aligned[0]; - DBG_PRINT("[Chare %d] input[%d]: ptr=%p sizes=[%lld", - chare_idx, i, (void*)descs[i].aligned, (long long)descs[i].sizes[0]); - for (int d = 1; d < N; ++d) - DBG_PRINT(",%lld", (long long)descs[i].sizes[d]); - DBG_PRINT("] strides=[%lld", (long long)descs[i].strides[0]); - for (int d = 1; d < N; ++d) - DBG_PRINT(",%lld", (long long)descs[i].strides[d]); - DBG_PRINT("] first_val=%f\n", (double)first_val); - } - for (int i = 0; i < n_outputs; ++i) { - int j = n_inputs + i; - T first_val = T(0); - if (descs[j].sizes[0] > 0 && (N < 2 || descs[j].sizes[1] > 0)) - first_val = descs[j].aligned[0]; - T sum = T(0); - int64_t total = 1; - for (int d = 0; d < N; ++d) - total *= descs[j].sizes[d]; - int check_count = std::min((int64_t)16, total); - for (int c = 0; c < check_count; ++c) { - int64_t flat = 0, rem = c; - for (int d = N - 1; d >= 0; --d) { - flat += (rem % descs[j].sizes[d]) * descs[j].strides[d]; - rem /= descs[j].sizes[d]; - } - sum += descs[j].aligned[flat]; - } - DBG_PRINT("[Chare %d] output[%d]: ptr=%p sizes=[%lld", - chare_idx, i, (void*)descs[j].aligned, (long long)descs[j].sizes[0]); - for (int d = 1; d < N; ++d) - DBG_PRINT(",%lld", (long long)descs[j].sizes[d]); - DBG_PRINT("] strides=[%lld", (long long)descs[j].strides[0]); - for (int d = 1; d < N; ++d) - DBG_PRINT(",%lld", (long long)descs[j].strides[d]); - DBG_PRINT("] first_val=%f sum_first16=%f\n", (double)first_val, (double)sum); - } - } - } - - // Handle SET_REGION side-effects - for (ASTNode* root : node->ast->roots) - if (static_cast(root->opcode) == Opcode::SET_REGION) - partition->executor->ast_visitor(root, dtype_of()); -} - -// Explicit template instantiations for execute_node_nd -template void ArrayDAGGroup::execute_node_nd<1, float>(DAGNode*, PartitionImpl<1>*, PendingComm<1>*); -template void ArrayDAGGroup::execute_node_nd<1, double>(DAGNode*, PartitionImpl<1>*, PendingComm<1>*); -template void ArrayDAGGroup::execute_node_nd<1, int32_t>(DAGNode*, PartitionImpl<1>*, PendingComm<1>*); -template void ArrayDAGGroup::execute_node_nd<1, int64_t>(DAGNode*, PartitionImpl<1>*, PendingComm<1>*); -template void ArrayDAGGroup::execute_node_nd<2, float>(DAGNode*, PartitionImpl<2>*, PendingComm<2>*); -template void ArrayDAGGroup::execute_node_nd<2, double>(DAGNode*, PartitionImpl<2>*, PendingComm<2>*); -template void ArrayDAGGroup::execute_node_nd<2, int32_t>(DAGNode*, PartitionImpl<2>*, PendingComm<2>*); -template void ArrayDAGGroup::execute_node_nd<2, int64_t>(DAGNode*, PartitionImpl<2>*, PendingComm<2>*); -template void ArrayDAGGroup::execute_node_nd<3, float>(DAGNode*, PartitionImpl<3>*, PendingComm<3>*); -template void ArrayDAGGroup::execute_node_nd<3, double>(DAGNode*, PartitionImpl<3>*, PendingComm<3>*); -template void ArrayDAGGroup::execute_node_nd<3, int32_t>(DAGNode*, PartitionImpl<3>*, PendingComm<3>*); -template void ArrayDAGGroup::execute_node_nd<3, int64_t>(DAGNode*, PartitionImpl<3>*, PendingComm<3>*); diff --git a/src/execute_node_regions.cpp b/src/execute_node_regions.cpp deleted file mode 100644 index 6dd1a42..0000000 --- a/src/execute_node_regions.cpp +++ /dev/null @@ -1,142 +0,0 @@ -#include "backend_internal.hpp" - -template -bool extract_regions_nd(DAGNode* node, std::unordered_map*>& arrays, - ArrayRegion& r_out, std::vector>& input_regions, - std::vector& input_source_names, - std::array& global_shape_out) { - bool has_regions = false; - bool has_output = false; - std::array global_shape = {}; - - // Resolve global_shape from the node's operands - for (ASTNode* root : node->ast->roots) { - auto opc = static_cast(root->opcode); - if (opc == Opcode::SET_REGION && root->operands.size() > 0) { - int ref_name = root->operands[0]->result_name; - auto ait = arrays.find(ref_name); - if (ait != arrays.end() && ait->second->global_size > 0) { - global_shape = ait->second->global_shape; - break; - } - } - } - if (global_shape[0] == 0) { - for (ASTNode* root : node->ast->roots) { - for (ASTNode* operand : root->operands) { - if (operand->is_scalar || operand->is_broadcast) - continue; - auto ait = arrays.find(operand->result_name); - if (ait != arrays.end() && ait->second->global_size > 0) { - global_shape = ait->second->global_shape; - break; - } - } - if (global_shape[0] != 0) - break; - } - } - if (global_shape[0] == 0) { - int best_gs = 0; - for (auto& [name, arr] : arrays) { - if (arr->global_size > best_gs) { - best_gs = arr->global_size; - global_shape = arr->global_shape; - } - } - } - global_shape_out = global_shape; - if (global_shape[0] == 0) - return false; - - std::unordered_set internal_names; - for (ASTNode* root : node->ast->roots) - internal_names.insert(root->result_name); - - for (ASTNode* root : node->ast->roots) { - auto opc = static_cast(root->opcode); - - if (opc == Opcode::SET_REGION) { - has_regions = true; - has_output = true; - auto* region = static_cast*>(root->region); - r_out = ArrayRegion(region->start, region->stop, region->step); - - ASTNode* rhs = root->operands[1]; - if (!rhs->is_scalar && !rhs->is_broadcast && - !internal_names.count(rhs->result_name)) { - // Use inline region from the RHS operand if available - Region* rhs_r = root->get_operand_region(1); - if (rhs_r && !rhs_r->is_global) { - auto* inp_r = static_cast*>(rhs_r); - input_regions.emplace_back(inp_r->start, inp_r->stop, inp_r->step); - } else { - std::array inp_start = {}; - std::array inp_stop; - std::array inp_step; - for (int d = 0; d < N; ++d) { - inp_stop[d] = r_out.size(d); - inp_step[d] = 1; - } - input_regions.emplace_back(inp_start, inp_stop, inp_step); - } - input_source_names.push_back(rhs->result_name); - } - } else if (is_elementwise(opc)) { - for (int op_idx = 0; op_idx < (int)root->operands.size(); ++op_idx) { - ASTNode* operand = root->operands[op_idx]; - if (operand->is_scalar || operand->is_broadcast) - continue; - if (internal_names.count(operand->result_name)) - continue; - - // Use inline region from operand_regions if available - Region* op_r = root->get_operand_region(op_idx); - if (op_r && !op_r->is_global) { - has_regions = true; - auto* inp_r = static_cast*>(op_r); - input_regions.emplace_back(inp_r->start, inp_r->stop, inp_r->step); - } else { - std::array inp_start = {}; - std::array inp_step; - for (int d = 0; d < N; ++d) - inp_step[d] = 1; - input_regions.emplace_back(inp_start, global_shape, inp_step); - } - input_source_names.push_back(operand->result_name); - } - - if (!has_output) { - if (has_regions && !input_regions.empty()) { - std::array out_start = {}; - std::array out_stop; - std::array out_step; - for (int d = 0; d < N; ++d) { - out_stop[d] = input_regions.front().size(d); - out_step[d] = 1; - } - r_out = ArrayRegion(out_start, out_stop, out_step); - } else { - std::array out_start = {}; - std::array out_step; - for (int d = 0; d < N; ++d) - out_step[d] = 1; - r_out = ArrayRegion(out_start, global_shape, out_step); - } - has_output = true; - } - } - } - - return has_regions; -} - -template bool extract_regions_nd<1>(DAGNode*, std::unordered_map*>&, - ArrayRegion<1>&, std::vector>&, - std::vector&, std::array&); -template bool extract_regions_nd<2>(DAGNode*, std::unordered_map*>&, - ArrayRegion<2>&, std::vector>&, - std::vector&, std::array&); -template bool extract_regions_nd<3>(DAGNode*, std::unordered_map*>&, - ArrayRegion<3>&, std::vector>&, - std::vector&, std::array&); diff --git a/src/executor_ast_visitor.cpp b/src/executor_ast_visitor.cpp deleted file mode 100644 index 1c5c101..0000000 --- a/src/executor_ast_visitor.cpp +++ /dev/null @@ -1,609 +0,0 @@ -#include "backend_internal.hpp" - -#include -#include - -#ifdef USE_KOKKOS -#include -#include -#endif - -template -int ArrayDAGExecutorND::ast_visitor(ASTNode* node, DType dtype) { - auto nd_idx = partition->nd_index(); - - switch (static_cast(node->opcode)) { - case Opcode::CREATE: { - auto* dag_group = static_cast(group); - auto meta_it = dag_group->array_meta.find(node->result_name); - auto decomp = meta_it->second.template decomp(); - return partition->create(static_cast*>(node->region), node->result_name, - node->dtype, decomp); - } - case Opcode::SET_REGION: { - // Get target array (type-erased) — may not exist on this chare - // if the array is smaller than the chare grid. - auto tgt_it = partition->arrays.find(node->operands[0]->result_name); - if (tgt_it == partition->arrays.end()) - return node->result_name; - CTArrayBase* target = tgt_it->second; - int esz = target->elem_size(); - - auto* region = static_cast*>(node->region); - ASTNode* rhs_node = node->operands.size() > 1 ? node->operands[1] : nullptr; - Region* rhs_region_base = node->get_operand_region(1); - - // Decomp-aware chare region and region-to-global translation - auto tgt_chare = target->decomp.chare_region_global(nd_idx); - std::array g_reg_start, g_reg_stop; - for (int d = 0; d < N; ++d) { - g_reg_start[d] = region->start[d] + target->decomp.offset[d]; - g_reg_stop[d] = region->stop[d] + target->decomp.offset[d]; - } - ArrayRegion dst_region_global(g_reg_start, g_reg_stop, region->step); - auto [overlap, has_overlap] = intersect(dst_region_global, tgt_chare); - - CTArrayBase* source = nullptr; - ArrayRegion source_region_global; - bool has_source_array = false; - bool source_is_result_buffer = false; - auto full_source_region_global = [&](CTArrayBase* candidate) { - std::array start{}, stop{}, step{}; - for (int d = 0; d < N; ++d) { - start[d] = candidate->decomp.offset[d]; - stop[d] = candidate->decomp.offset[d] + candidate->global_shape[d]; - step[d] = 1; - } - return ArrayRegion(start, stop, step); - }; - auto set_source = [&](CTArrayBase* candidate, ArrayRegion const& candidate_region) { - source = candidate; - source_region_global = candidate_region; - has_source_array = true; - source_is_result_buffer = false; - }; - - if (!has_source_array) { - auto temp_it = partition->arrays.find(node->result_name); - if (temp_it != partition->arrays.end()) { - // Prefer the communicated/caller-visible SET_REGION result buffer - // when it exists. For mixed decompositions (for example 65->129 - // prolongation), that buffer is already laid out on the target - // slice's chare grid, while the raw RHS array may only contain - // this chare's local source tile. - source = temp_it->second; - has_source_array = true; - source_is_result_buffer = true; - } - } - - if (!has_source_array && rhs_node) { - auto rhs_it = partition->arrays.find(rhs_node->result_name); - auto rhs_op = static_cast(rhs_node->opcode); - if (rhs_it != partition->arrays.end() && - (rhs_op == Opcode::NOOP || rhs_op == Opcode::CREATE)) { - if (rhs_region_base && !rhs_region_base->is_global) { - auto* rhs_region = static_cast*>(rhs_region_base); - set_source(rhs_it->second, rhs_it->second->decomp.to_global(*rhs_region)); - } else { - set_source(rhs_it->second, full_source_region_global(rhs_it->second)); - } - } - } - - if (!has_source_array && rhs_node) { - auto rhs_it = partition->arrays.find(rhs_node->result_name); - if (rhs_it != partition->arrays.end()) { - if (rhs_region_base && !rhs_region_base->is_global) { - auto* rhs_region = static_cast*>(rhs_region_base); - set_source(rhs_it->second, rhs_it->second->decomp.to_global(*rhs_region)); - } else { - set_source(rhs_it->second, full_source_region_global(rhs_it->second)); - } - } - } - - // Compute row-major strides for target - std::array tgt_strides; - tgt_strides[N - 1] = 1; - for (int d = N - 2; d >= 0; --d) - tgt_strides[d] = tgt_strides[d + 1] * target->region.size(d + 1); - - if (!has_source_array) { - // Scalar fill — convert scalar to target dtype - if (has_overlap && node->operands.size() > 1 && node->operands[1]->is_scalar) { - char scalar_buf[8]; - double s = node->operands[1]->scalar; - switch (dtype) { - case DType::FLOAT32: { - float v = (float)s; - memcpy(scalar_buf, &v, sizeof(v)); - break; - } - case DType::FLOAT64: { - memcpy(scalar_buf, &s, sizeof(s)); - break; - } - case DType::INT32: { - int32_t v = (int32_t)s; - memcpy(scalar_buf, &v, sizeof(v)); - break; - } - case DType::INT64: { - int64_t v = (int64_t)s; - memcpy(scalar_buf, &v, sizeof(v)); - break; - } - } - -#ifdef USE_KOKKOS - // Device-side scalar fill - char* tgt_device = static_cast(target->device_data_ptr()); - // Copy scalar_buf to device (8 bytes as int64_t for capture) - int64_t scalar_bits; - memcpy(&scalar_bits, scalar_buf, sizeof(scalar_bits)); - - int64_t overlap_total = overlap.size(); - int ol_start[N], ol_step[N], tgt_s[N], ol_sizes[N], tgt_cs_arr[N]; - for (int d = 0; d < N; d++) { - ol_start[d] = overlap.start[d]; - ol_step[d] = overlap.step[d]; - tgt_s[d] = tgt_strides[d]; - ol_sizes[d] = overlap.size(d); - tgt_cs_arr[d] = tgt_chare.start[d]; - } - - Kokkos::parallel_for( - CT_COMPUTE_POLICY(partition, overlap_total), KOKKOS_LAMBDA(int flat_idx) { - int remaining = flat_idx; - int flat = 0; - for (int d = N - 1; d >= 0; --d) { - int coord_d = remaining % ol_sizes[d]; - remaining /= ol_sizes[d]; - int global_d = ol_start[d] + coord_d * ol_step[d]; - flat += (global_d - tgt_cs_arr[d]) * tgt_s[d]; - } - char* dst = tgt_device + flat * esz; - const char* src = reinterpret_cast(&scalar_bits); - for (int b = 0; b < esz; b++) - dst[b] = src[b]; - }); -#else - // Host-side odometer over the overlap region - char* tgt_data = static_cast(target->data_ptr()); - std::array idx; - for (int d = 0; d < N; ++d) - idx[d] = overlap.start[d]; - while (true) { - int flat = 0; - for (int d = 0; d < N; ++d) - flat += (idx[d] - tgt_chare.start[d]) * tgt_strides[d]; - memcpy(tgt_data + flat * esz, scalar_buf, esz); - - int d = N - 1; - while (d >= 0) { - idx[d] += overlap.step[d]; - if (idx[d] < overlap.stop[d]) - break; - idx[d] = overlap.start[d]; - --d; - } - if (d < 0) - break; - } -#endif - } - return node->result_name; - } - - if (has_overlap) { - // Compute row-major strides for source - std::array src_strides; - src_strides[N - 1] = 1; - for (int d = N - 2; d >= 0; --d) - src_strides[d] = src_strides[d + 1] * source->region.size(d + 1); - - auto src_chare = source->decomp.chare_region_global(nd_idx); - -#ifdef USE_KOKKOS - // Device-side element copy - char* tgt_device = static_cast(target->device_data_ptr()); - char* src_device = static_cast(source->device_data_ptr()); - - int64_t overlap_total = overlap.size(); - int ol_start[N], ol_step[N], tgt_s[N], src_s[N], ol_sizes[N], tgt_cs_arr2[N]; - int g_reg_s[N], reg_step[N], src_chare_s[N], src_reg_s[N], src_reg_step[N], - src_region_sizes[N]; - for (int d = 0; d < N; d++) { - ol_start[d] = overlap.start[d]; - ol_step[d] = overlap.step[d]; - tgt_s[d] = tgt_strides[d]; - src_s[d] = src_strides[d]; - ol_sizes[d] = overlap.size(d); - tgt_cs_arr2[d] = tgt_chare.start[d]; - g_reg_s[d] = g_reg_start[d]; - reg_step[d] = region->step[d]; - src_chare_s[d] = src_chare.start[d]; - src_reg_s[d] = source_region_global.start[d]; - src_reg_step[d] = source_region_global.step[d]; - src_region_sizes[d] = source->region.size(d); - } - - Kokkos::parallel_for( - CT_COMPUTE_POLICY(partition, overlap_total), KOKKOS_LAMBDA(int flat_idx) { - int remaining = flat_idx; - int dst_flat = 0, src_flat = 0; - bool src_valid = true; - for (int d = N - 1; d >= 0; --d) { - int coord_d = remaining % ol_sizes[d]; - remaining /= ol_sizes[d]; - int global_d = ol_start[d] + coord_d * ol_step[d]; - dst_flat += (global_d - tgt_cs_arr2[d]) * tgt_s[d]; - if (source_is_result_buffer) { - int src_local = (global_d - ol_start[d]) / ol_step[d]; - if (src_local < 0 || src_local >= src_region_sizes[d]) { - src_valid = false; - break; - } - src_flat += src_local * src_s[d]; - } else { - int logical_idx = (global_d - g_reg_s[d]) / reg_step[d]; - int src_global = src_reg_s[d] + logical_idx * src_reg_step[d]; - int src_local = src_global - src_chare_s[d]; - if (src_local < 0 || src_local >= src_region_sizes[d]) { - src_valid = false; - break; - } - src_flat += src_local * src_s[d]; - } - } - if (src_valid) { - char* dst = tgt_device + dst_flat * esz; - char* src = src_device + src_flat * esz; - for (int b = 0; b < esz; b++) - dst[b] = src[b]; - } - }); -#else - // Host-side odometer over the overlap region - char* tgt_data = static_cast(target->data_ptr()); - char* src_data = static_cast(source->data_ptr()); - std::array idx; - for (int d = 0; d < N; ++d) - idx[d] = overlap.start[d]; - while (true) { - int dst_flat = 0, src_flat = 0; - bool src_valid = true; - for (int d = 0; d < N; ++d) { - dst_flat += (idx[d] - tgt_chare.start[d]) * tgt_strides[d]; - if (source_is_result_buffer) { - int src_local = (idx[d] - overlap.start[d]) / overlap.step[d]; - if (src_local < 0 || src_local >= source->region.size(d)) { - src_valid = false; - break; - } - src_flat += src_local * src_strides[d]; - } else { - int logical_idx = (idx[d] - g_reg_start[d]) / region->step[d]; - int src_global = source_region_global.start[d] + - logical_idx * source_region_global.step[d]; - int src_local = src_global - src_chare.start[d]; - if (src_local < 0 || src_local >= source->region.size(d)) { - src_valid = false; - break; - } - src_flat += src_local * src_strides[d]; - } - } - if (src_valid) - memcpy(tgt_data + dst_flat * esz, src_data + src_flat * esz, esz); - - int d = N - 1; - while (d >= 0) { - idx[d] += overlap.step[d]; - if (idx[d] < overlap.stop[d]) - break; - idx[d] = overlap.start[d]; - --d; - } - if (d < 0) - break; - } -#endif - } - return node->result_name; - } - case Opcode::REDUCE: { - // REDUCE is fully handled by execute_reduce_node + Charm++ reduction. - // The ast_visitor should never be reached for REDUCE nodes. - return node->result_name; - } - case Opcode::MATMUL: { - if constexpr (N >= 2) { - // Local matmul in ast_visitor (no-comm path, e.g. single chare) - // This is reached when execute_matmul_node had total_expected == 0 - int mat_name = node->operands[0]->result_name; - int vec_name = node->operands[1]->result_name; - int result_name = node->result_name; - - auto mat_it = partition->arrays.find(mat_name); - auto vec_it = partition->arrays.find(vec_name); - if (mat_it == partition->arrays.end() || vec_it == partition->arrays.end()) { - DBG_PRINT("[Chare %d] MATMUL ast_visitor: missing operand arrays\n", - partition->index[0]); - return result_name; - } - - CTArrayBase* mat_base = mat_it->second; - CTArrayBase* vec_base = vec_it->second; - int local_rows = mat_base->region.size(0); - int local_cols = mat_base->region.size(1); - int vec_len = vec_base->local_size(); - - // Create result array if needed - if (partition->arrays.find(result_name) == partition->arrays.end()) { - std::array out_start = {}, out_stop, out_step, out_gs; - out_stop[0] = local_rows; - out_step[0] = 1; - out_gs[0] = mat_base->global_shape[0]; - if (N > 1) { - out_stop[1] = 1; - out_step[1] = 1; - out_gs[1] = 1; - } - ArrayRegion out_region(out_start, out_stop, out_step); - ArrayDecomp out_decomp = - static_cast(group)->array_meta[result_name].template decomp(); - // Use dtype dispatch to create the correctly typed result array - partition->arrays[result_name] = - partition->allocate_or_reuse(out_region, out_gs, result_name, dtype, out_decomp); - } - - CTArrayBase* result_base = partition->arrays[result_name]; - int actual_cols = std::min(local_cols, vec_len); - -#ifdef USE_KOKKOS - // Perform matmul on device using KokkosBlas::gemv - auto kokkos_gemv = [&](auto dummy) { - using VT = decltype(dummy); - Kokkos::View> - d_mat(static_cast(mat_base->device_data_ptr()), local_rows, local_cols); - auto d_sub = - Kokkos::subview(d_mat, Kokkos::ALL, Kokkos::make_pair(0, actual_cols)); - Kokkos::View> - d_vec(static_cast(vec_base->device_data_ptr()), actual_cols); - Kokkos::View> - d_res(static_cast(result_base->device_data_ptr()), local_rows); - KokkosBlas::gemv("N", VT(1), d_sub, d_vec, VT(0), d_res); - }; - switch (dtype) { - case DType::FLOAT32: - kokkos_gemv(float{}); - break; - case DType::FLOAT64: - kokkos_gemv(double{}); - break; - case DType::INT32: - kokkos_gemv(int32_t{}); - break; - case DType::INT64: - kokkos_gemv(int64_t{}); - break; - } -#else - // Perform matmul on host using Eigen gemv - mat_base->copyToHost(); - vec_base->copyToHost(); - - switch (dtype) { - case DType::FLOAT32: - eigen_gemv(static_cast(mat_base->data_ptr()), local_rows, local_cols, - static_cast(vec_base->data_ptr()), actual_cols, - static_cast(result_base->data_ptr())); - break; - case DType::FLOAT64: - eigen_gemv(static_cast(mat_base->data_ptr()), local_rows, local_cols, - static_cast(vec_base->data_ptr()), actual_cols, - static_cast(result_base->data_ptr())); - break; - case DType::INT32: - eigen_gemv(static_cast(mat_base->data_ptr()), local_rows, local_cols, - static_cast(vec_base->data_ptr()), actual_cols, - static_cast(result_base->data_ptr())); - break; - case DType::INT64: - eigen_gemv(static_cast(mat_base->data_ptr()), local_rows, local_cols, - static_cast(vec_base->data_ptr()), actual_cols, - static_cast(result_base->data_ptr())); - break; - } -#endif - DBG_PRINT("[Chare %d] MATMUL ast_visitor: %dx%d @ %d done\n", partition->index[0], - local_rows, actual_cols, vec_len); - return result_name; - } else { - // N=1: MATMUL on 1D partition is handled via comm path, not ast_visitor - return node->result_name; - } - } - case Opcode::MATMATMUL: { - if constexpr (N == 2) { - // Local matmatmul in ast_visitor (no-comm path, e.g. single chare) - int a_name = node->operands[0]->result_name; - int b_name = node->operands[1]->result_name; - int result_name = node->result_name; - - auto a_it = partition->arrays.find(a_name); - auto b_it = partition->arrays.find(b_name); - if (a_it == partition->arrays.end() || b_it == partition->arrays.end()) { - DBG_PRINT("[Chare %d] MATMATMUL ast_visitor: missing operand arrays\n", - partition->index[0]); - return result_name; - } - - CTArrayBase* a_base = a_it->second; - CTArrayBase* b_base = b_it->second; - - // Determine effective sub-block sizes from regions - int a_local_rows = a_base->region.size(0); - int a_local_cols = a_base->region.size(1); - int b_local_rows = b_base->region.size(0); - int b_local_cols = b_base->region.size(1); - - int sub_rows = a_local_rows; - int k_size = std::min(a_local_cols, b_local_rows); - int sub_cols = b_local_cols; - - // Handle slice regions if present - int a_row_offset = 0, a_col_offset = 0; - int b_row_offset = 0, b_col_offset = 0; - if (node->operand_regions.size() >= 2) { - auto* ar = static_cast*>(node->operand_regions[0]); - auto* br = static_cast*>(node->operand_regions[1]); - // Translate regions to global space using array decomps - auto a_decomp = a_base->decomp; - auto b_decomp = b_base->decomp; - int a_rs = ar->start[0] + a_decomp.offset[0]; - int a_re = ar->stop[0] + a_decomp.offset[0]; - int a_cs = ar->start[1] + a_decomp.offset[1]; - int a_ce = ar->stop[1] + a_decomp.offset[1]; - int b_rs = br->start[0] + b_decomp.offset[0]; - int b_re = br->stop[0] + b_decomp.offset[0]; - int b_cs = br->start[1] + b_decomp.offset[1]; - int b_ce = br->stop[1] + b_decomp.offset[1]; - - auto a_chare = a_decomp.chare_region_global(nd_idx); - int chare_a_row_lo = std::max(a_chare.start[0], a_rs); - int chare_a_row_hi = std::min(a_chare.stop[0], a_re); - int chare_a_col_lo = std::max(a_chare.start[1], a_cs); - int chare_a_col_hi = std::min(a_chare.stop[1], a_ce); - a_row_offset = chare_a_row_lo - a_chare.start[0]; - a_col_offset = chare_a_col_lo - a_chare.start[1]; - sub_rows = chare_a_row_hi - chare_a_row_lo; - - auto b_chare = b_decomp.chare_region_global(nd_idx); - int chare_b_row_lo = std::max(b_chare.start[0], b_rs); - int chare_b_row_hi = std::min(b_chare.stop[0], b_re); - int chare_b_col_lo = std::max(b_chare.start[1], b_cs); - int chare_b_col_hi = std::min(b_chare.stop[1], b_ce); - b_row_offset = chare_b_row_lo - b_chare.start[0]; - b_col_offset = chare_b_col_lo - b_chare.start[1]; - k_size = - std::min(chare_a_col_hi - chare_a_col_lo, chare_b_row_hi - chare_b_row_lo); - sub_cols = chare_b_col_hi - chare_b_col_lo; - } - - if (sub_rows <= 0 || k_size <= 0 || sub_cols <= 0) - return result_name; - - // Create result array - auto* dag_group2 = static_cast(group); - auto c_meta = dag_group2->array_meta.find(result_name); - int c_M = c_meta->second.global_shape[0]; - int c_N = c_meta->second.global_shape[1]; - - if (partition->arrays.find(result_name) == partition->arrays.end()) { - std::array out_start = {0, 0}; - std::array out_stop = {sub_rows, sub_cols}; - std::array out_step = {1, 1}; - std::array out_gs = {c_M, c_N}; - ArrayRegion<2> out_region(out_start, out_stop, out_step); - ArrayDecomp<2> out_decomp = - dag_group2->array_meta[result_name].template decomp<2>(); - partition->arrays[result_name] = - partition->allocate_or_reuse(out_region, out_gs, result_name, dtype, out_decomp); - } - -#ifdef USE_KOKKOS - // Perform matmatmul on device using KokkosBlas::gemm - auto kokkos_gemm = [&](auto dummy) { - using VT = decltype(dummy); - Kokkos::View> - d_a(static_cast(a_base->device_data_ptr()), a_local_rows, a_local_cols); - auto d_a_sub = Kokkos::subview(d_a, - Kokkos::make_pair(a_row_offset, a_row_offset + sub_rows), - Kokkos::make_pair(a_col_offset, a_col_offset + k_size)); - Kokkos::View> - d_b(static_cast(b_base->device_data_ptr()), b_local_rows, b_local_cols); - auto d_b_sub = Kokkos::subview(d_b, - Kokkos::make_pair(b_row_offset, b_row_offset + k_size), - Kokkos::make_pair(b_col_offset, b_col_offset + sub_cols)); - Kokkos::View> - d_c(static_cast(partition->arrays[result_name]->device_data_ptr()), sub_rows, - sub_cols); - KokkosBlas::gemm("N", "N", VT(1), d_a_sub, d_b_sub, VT(1), d_c); - }; - switch (dtype) { - case DType::FLOAT32: - kokkos_gemm(float{}); - break; - case DType::FLOAT64: - kokkos_gemm(double{}); - break; - case DType::INT32: - kokkos_gemm(int32_t{}); - break; - case DType::INT64: - kokkos_gemm(int64_t{}); - break; - } -#else - a_base->copyToHost(); - b_base->copyToHost(); - - switch (dtype) { - case DType::FLOAT32: - eigen_gemm_sub(static_cast(a_base->data_ptr()), a_local_cols, a_row_offset, - a_col_offset, sub_rows, k_size, - static_cast(b_base->data_ptr()), b_local_cols, b_row_offset, - b_col_offset, sub_cols, - static_cast(partition->arrays[result_name]->data_ptr())); - break; - case DType::FLOAT64: - eigen_gemm_sub(static_cast(a_base->data_ptr()), a_local_cols, a_row_offset, - a_col_offset, sub_rows, k_size, - static_cast(b_base->data_ptr()), b_local_cols, b_row_offset, - b_col_offset, sub_cols, - static_cast(partition->arrays[result_name]->data_ptr())); - break; - case DType::INT32: - eigen_gemm_sub(static_cast(a_base->data_ptr()), a_local_cols, a_row_offset, - a_col_offset, sub_rows, k_size, - static_cast(b_base->data_ptr()), b_local_cols, b_row_offset, - b_col_offset, sub_cols, - static_cast(partition->arrays[result_name]->data_ptr())); - break; - case DType::INT64: - eigen_gemm_sub(static_cast(a_base->data_ptr()), a_local_cols, a_row_offset, - a_col_offset, sub_rows, k_size, - static_cast(b_base->data_ptr()), b_local_cols, b_row_offset, - b_col_offset, sub_cols, - static_cast(partition->arrays[result_name]->data_ptr())); - break; - } -#endif - DBG_PRINT("[Chare %d] MATMATMUL ast_visitor: (%dx%d) @ (%dx%d) done\n", - partition->index[0], sub_rows, k_size, k_size, sub_cols); - return result_name; - } else { - return node->result_name; - } - } - case Opcode::DIAG: - // DIAG is always handled via the comm path (execute_diag_node). - return node->result_name; - case Opcode::TILE: - // TILE is always handled via the comm path (execute_tile_node). - return node->result_name; - default: - CkAbort("Unknown opcode in interpreter: %d", node->opcode); - } -} - -template int ArrayDAGExecutorND<1>::ast_visitor(ASTNode*, DType); -template int ArrayDAGExecutorND<2>::ast_visitor(ASTNode*, DType); -template int ArrayDAGExecutorND<3>::ast_visitor(ASTNode*, DType); diff --git a/src/executor_core.cpp b/src/executor_core.cpp deleted file mode 100644 index 70bc8d4..0000000 --- a/src/executor_core.cpp +++ /dev/null @@ -1,787 +0,0 @@ -#include "backend_internal.hpp" -#include "dispatch.hpp" - -#include -#include - -template -static void send_remote_input(PartitionImpl* partition, DAGNode* node, const RemoteSend& send, - int inp_name, const std::array& nd_idx) { - auto arr_it = partition->arrays.find(inp_name); - if (arr_it == partition->arrays.end()) - return; - Array* arr = static_cast*>(arr_it->second); - - // Use the input array's own decomp for chare region and local offset - ArrayRegion r_chare_inp = arr->decomp.chare_region_global(nd_idx); - auto [overlap, has_overlap] = intersect(send.region, r_chare_inp); - if (!has_overlap) - return; - - std::array cs; - for (int d = 0; d < N; ++d) - cs[d] = arr->decomp.chare_start_global(d, nd_idx[d]); - int64_t total_size = overlap.size(); - - { - bool out_of_bounds = false; - for (int d = 0; d < N; ++d) { - int local_start_d = overlap.start[d] - cs[d]; - int phys_extent = local_start_d + (overlap.size(d) - 1) * overlap.step[d] + 1; - if (local_start_d < 0 || phys_extent > arr->region.size(d)) { - out_of_bounds = true; - break; - } - } - if (out_of_bounds) - return; - } - - std::array arr_strides; - arr_strides[N - 1] = 1; - for (int d = N - 2; d >= 0; --d) - arr_strides[d] = arr_strides[d + 1] * arr->region.size(d + 1); - - int region_data[N * 3]; - for (int d = 0; d < N; ++d) { - region_data[d * 3 + 0] = overlap.start[d]; - region_data[d * 3 + 1] = overlap.stop[d]; - region_data[d * 3 + 2] = overlap.step[d]; - } - int64_t byte_size = total_size * sizeof(T); - -#ifndef NDEBUG - partition->comm_bytes_sent += byte_size; -#endif - -#ifdef USE_KOKKOS - // Pack on device and send via direct GPU messaging - T* send_buf = static_cast(Kokkos::kokkos_malloc(total_size * sizeof(T))); - T* src_device = static_cast(arr->device_data_ptr()); - - // Capture strides/offsets in plain arrays for KOKKOS_LAMBDA - int overlap_sizes[N], local_starts[N], src_strides_arr[N], overlap_steps[N]; - for (int d = 0; d < N; d++) { - overlap_sizes[d] = overlap.size(d); - local_starts[d] = overlap.start[d] - cs[d]; - src_strides_arr[d] = arr_strides[d]; - overlap_steps[d] = overlap.step[d]; - } - - Kokkos::parallel_for( - CT_COMM_POLICY(partition, total_size), KOKKOS_LAMBDA(int flat_idx) { - int remaining = flat_idx; - int src_flat = 0; - for (int d = N - 1; d >= 0; --d) { - int coord_d = remaining % overlap_sizes[d]; - remaining /= overlap_sizes[d]; - src_flat += (local_starts[d] + coord_d * overlap_steps[d]) * src_strides_arr[d]; - } - send_buf[flat_idx] = src_device[src_flat]; - }); - - device_pack_send(partition, partition->thisProxy, send.target, - node->id, send.input_index, inp_name, region_data, - byte_size, send_buf); -#else - // Host path: pack on CPU and send via regular Charm++ messaging - arr->copyToHost(); - T* send_buf = new T[total_size]; - T* host_data = static_cast(arr->data_ptr()); - - if (overlap.step[N - 1] == 1) { - // Fast path: innermost dimension is contiguous, use memcpy - int64_t inner_size = overlap.size(N - 1); - int64_t buf_offset = 0; - std::array idx = {}; - while (true) { - int64_t src_flat = 0; - for (int d = 0; d < N; ++d) - src_flat += (int64_t)(overlap.start[d] + idx[d] * overlap.step[d] - cs[d]) * - arr_strides[d]; - memcpy(send_buf + buf_offset, host_data + src_flat, inner_size * sizeof(T)); - buf_offset += inner_size; - - int d = N - 2; - while (d >= 0) { - if (++idx[d] < overlap.size(d)) - break; - idx[d] = 0; - --d; - } - if (d < 0) - break; - } - } else { - // Element-wise packing for non-unit innermost step - int64_t buf_offset = 0; - std::array idx = {}; - while (true) { - int64_t src_flat = 0; - for (int d = 0; d < N; ++d) - src_flat += (int64_t)(overlap.start[d] + idx[d] * overlap.step[d] - cs[d]) * - arr_strides[d]; - send_buf[buf_offset++] = host_data[src_flat]; - - int d = N - 1; - while (d >= 0) { - if (++idx[d] < overlap.size(d)) - break; - idx[d] = 0; - --d; - } - if (d < 0) - break; - } - } - - proxy_at(partition->thisProxy, send.target) - .receive_data(node->id, send.input_index, inp_name, N, region_data, byte_size, - reinterpret_cast(send_buf)); - delete[] send_buf; -#endif -} - -template -static void dispatch_send(DType dt, PartitionImpl* partition, DAGNode* node, - const RemoteSend& send, int inp_name, - const std::array& nd_idx) { - switch (dt) { - case DType::FLOAT32: - send_remote_input(partition, node, send, inp_name, nd_idx); - break; - case DType::FLOAT64: - send_remote_input(partition, node, send, inp_name, nd_idx); - break; - case DType::INT32: - send_remote_input(partition, node, send, inp_name, nd_idx); - break; - case DType::INT64: - send_remote_input(partition, node, send, inp_name, nd_idx); - break; - } -} - -template -static void dispatch_execute(DType dt, ArrayDAGGroup* group, DAGNode* node, - PartitionImpl* partition, PendingComm* comm) { - switch (dt) { - case DType::FLOAT32: - group->execute_node_nd(node, partition, comm); - break; - case DType::FLOAT64: - group->execute_node_nd(node, partition, comm); - break; - case DType::INT32: - group->execute_node_nd(node, partition, comm); - break; - case DType::INT64: - group->execute_node_nd(node, partition, comm); - break; - } -} - -template -void ArrayDAGExecutorND::delete_array(int name) { - auto* dag_group = static_cast(group); - const int had_live_meta = dag_group->live_array_meta.count(name) ? 1 : 0; - DBG_PRINT("[PE %d] Partition<%d> delete_array epoch=%d name=%d live_meta=%d\n", - CkMyPe(), N, epoch, name, had_live_meta); - partition->retire_array(name); - dag_group->live_array_meta.erase(name); -} - -template -void ArrayDAGExecutorND::execute_dag_node(DAGNode* node) { - // Check if this node is relevant to our partition dimensionality - bool relevant = false; - bool is_matmul_node = false; - bool is_matmatmul_node = false; - bool is_reduce_node = false; - bool is_cross_set_region = false; - bool is_diag_node = false; - bool is_tile_node = false; - // Broadcast operands: size-1 arrays from another partition used as scalars - struct BroadcastInfo { - int source_name; - int source_ndims; - int target_ndims; - }; - std::vector broadcast_ops; - auto* dag_group_meta = static_cast(group); - auto nd_idx = partition->nd_index(); - - auto root_output_name = [](ASTNode* root) { - if (root == nullptr) - return -1; - if (static_cast(root->opcode) == Opcode::SET_REGION && - !root->operands.empty() && root->operands[0] != nullptr) - return root->operands[0]->result_name; - return root->result_name; - }; - - auto chare_owns_result = [&](ASTNode* root) -> bool { - if (root == nullptr || root->ndims != N) - return false; - - int output_name = root_output_name(root); - auto meta_it = dag_group_meta->array_meta.find(output_name); - if (meta_it == dag_group_meta->array_meta.end()) - return false; - - auto decomp = meta_it->second.template decomp(); - auto chare_region = decomp.chare_region_global(nd_idx); - if (chare_region.size() <= 0) - return false; - - if (static_cast(root->opcode) == Opcode::SET_REGION && root->region != nullptr && - !root->region->is_global) { - auto* out_region = static_cast*>(root->region); - auto out_region_global = decomp.to_global(*out_region); - auto [overlap, has_overlap] = intersect(out_region_global, chare_region); - return has_overlap && overlap.size() > 0; - } - - return true; - }; - - for (ASTNode* root : node->ast->roots) { - auto opc = static_cast(root->opcode); - if (opc == Opcode::CREATE) { - if (root->ndims == N) - relevant = true; - continue; - } - if (opc == Opcode::REDUCE) { - is_reduce_node = true; - if (N == 1) - relevant = true; - continue; - } - if (opc == Opcode::MATMUL) { - is_matmul_node = true; - // MATMUL is relevant for 1D (vector send/result receive), - // 2D (matrix compute), and 3D (dimension-dropped matvec) - if (N == 1 || N == 2 || N == 3) - relevant = true; - continue; - } - if (opc == Opcode::MATMATMUL) { - is_matmatmul_node = true; - // MATMATMUL is relevant for 2D chares that hold A, B, or will hold C, - // and 3D chares with dim-dropped operands - if constexpr (N == 2) { - // Check if this chare holds A or B - for (ASTNode* operand : root->operands) - if (!operand->is_scalar && !operand->is_broadcast && - partition->arrays.count(operand->result_name)) - relevant = true; - // Check if this chare will hold part of C - auto* dag_group_tmp = static_cast(group); - auto c_meta = dag_group_tmp->array_meta.find(root->result_name); - if (c_meta != dag_group_tmp->array_meta.end()) { - auto nd = partition->nd_index(); - auto c_decomp = c_meta->second.template decomp<2>(); - auto c_chare = c_decomp.chare_region_global(nd); - if (c_chare.stop[0] > c_chare.start[0] && - c_chare.stop[1] > c_chare.start[1]) - relevant = true; - } - } else if constexpr (N == 3) { - relevant = true; - } - continue; - } - if (opc == Opcode::DIAG) { - is_diag_node = true; - // DIAG is relevant for 1D and 2D partitions - if (N == 1 || N == 2) - relevant = true; - continue; - } - if (opc == Opcode::TILE) { - is_tile_node = true; - int input_ndims = root->operands[0]->ndims; - int result_ndims = root->ndims; - if (N == input_ndims || N == result_ndims) - relevant = true; - continue; - } - if (opc == Opcode::SET_REGION && root->operands.size() >= 2 && - !root->operands[1]->is_scalar && !root->operands[1]->is_broadcast) { - int source_ndims = root->operands[1]->ndims; - int target_ndims = root->ndims; - if (source_ndims != target_ndims) { - is_cross_set_region = true; - // Relevant if we are the source or target partition - if (N == source_ndims || N == target_ndims) - relevant = true; - continue; - } - } - // Detect broadcast operands (cross-partition or same-partition) - for (ASTNode* operand : root->operands) { - if (operand->is_broadcast) { - auto* dag_group = static_cast(group); - auto meta_it = dag_group->array_meta.find(operand->result_name); - if (meta_it != dag_group->array_meta.end()) { - int src_nd = meta_it->second.ndims; - int tgt_nd = root->ndims; - broadcast_ops.push_back({operand->result_name, src_nd, tgt_nd}); - // Relevant for source and target partitions - if (N == src_nd || N == tgt_nd) - relevant = true; - } - } - } - // For other ops, check if any operand arrays exist on this partition - for (ASTNode* operand : root->operands) - if (!operand->is_scalar && !operand->is_broadcast && - partition->arrays.count(operand->result_name)) - relevant = true; - if (partition->arrays.count(root->result_name)) - relevant = true; - if (chare_owns_result(root)) - relevant = true; - } - if (!relevant) { - DBG_PRINT("[PE %d] Partition<%d> chare %d: IRRELEVANT node %d\n", - CkMyPe(), N, partition->index[0], node->id); - node_finished(node->id); - return; - } - - // Check for REDUCE nodes — 1D dot product with cross-chare reduction - if (is_reduce_node) { - execute_reduce_node(node); - return; - } - - // Check for MATMUL nodes — these use a custom communication pattern - if (is_matmul_node) { - execute_matmul_node(node); - return; - } - - // Check for MATMATMUL nodes — SUMMA matrix-matrix multiply - if (is_matmatmul_node) { - execute_matmatmul_node(node); - return; - } - - // Check for cross-partition SET_REGION - if (is_cross_set_region) { - execute_cross_set_region_node(node); - return; - } - - // Check for DIAG nodes — cross-partition diagonal construction/extraction - if (is_diag_node) { - execute_diag_node(node); - return; - } - - // Check for TILE nodes — numpy.tile repetition with possible ndims change - if (is_tile_node) { - execute_tile_node(node); - return; - } - - // Handle cross-partition broadcast: send scalar values from source to target - if (!broadcast_ops.empty()) { - auto* dag_group = static_cast(group); - bool is_source_only = true; // true if this partition only sends, doesn't compute - bool has_local_broadcast = false; // true if broadcast source is on this partition+chare - - // Compute broadcast input index for each broadcast op (shared logic). - // Must match the receiver convention: n_memref_leaves + bcast_index. - // First count all memref leaves, then find the broadcast index. - auto compute_bcast_input_idx = [&](int source_name) -> int { - int n_memrefs = 0; - for (ASTNode* root : node->ast->roots) { - auto root_opc = static_cast(root->opcode); - for (ASTNode* operand : root->operands) { - if (operand->is_scalar || operand->is_broadcast) - continue; - if (root_opc == Opcode::SET_REGION && operand == root->operands[0]) - continue; - n_memrefs++; - } - } - - // Second pass: find the broadcast index for source_name - std::unordered_set seen; - int bcast_idx = 0; - for (ASTNode* root : node->ast->roots) { - for (ASTNode* operand : root->operands) { - if (operand->is_scalar || seen.count(operand->result_name)) - continue; - seen.insert(operand->result_name); - if (operand->is_broadcast) { - if (operand->result_name == source_name) - return n_memrefs + bcast_idx; - bcast_idx++; - } - } - } - return n_memrefs; - }; - - // Helper: send a broadcast scalar to a target chare identified by ND index - auto send_broadcast = [&](int tgt_nd, const std::array& tgt_idx, - int bcast_input_idx, int source_name, char* val_ptr, - int elem_sz) { - std::vector region_data(tgt_nd * 3, 0); - for (int d = 0; d < tgt_nd; ++d) { - region_data[d * 3 + 1] = 1; - region_data[d * 3 + 2] = 1; - } - char* send_buf = new char[elem_sz]; - memcpy(send_buf, val_ptr, elem_sz); -#ifndef NDEBUG - partition->comm_bytes_sent += elem_sz; -#endif - switch (tgt_nd) { - case 1: { - ChareIndex<1> ci; - ci.idx[0] = tgt_idx[0]; - proxy_at<1>(dag_group->partition_proxy_1, ci) - .receive_data(node->id, bcast_input_idx, source_name, tgt_nd, - region_data.data(), elem_sz, send_buf); - break; - } - case 2: { - ChareIndex<2> ci; - ci.idx[0] = tgt_idx[0]; - ci.idx[1] = tgt_idx[1]; - proxy_at<2>(dag_group->partition_proxy_2, ci) - .receive_data(node->id, bcast_input_idx, source_name, tgt_nd, - region_data.data(), elem_sz, send_buf); - break; - } - case 3: { - ChareIndex<3> ci; - ci.idx[0] = tgt_idx[0]; - ci.idx[1] = tgt_idx[1]; - ci.idx[2] = tgt_idx[2]; - proxy_at<3>(dag_group->partition_proxy_3, ci) - .receive_data(node->id, bcast_input_idx, source_name, tgt_nd, - region_data.data(), elem_sz, send_buf); - break; - } - } - delete[] send_buf; - }; - - for (auto& bcast : broadcast_ops) { - int bcast_input_idx = compute_bcast_input_idx(bcast.source_name); - - if (N == bcast.source_ndims) { - // We are on the source partition — send if we have the array - auto arr_it = partition->arrays.find(bcast.source_name); - DBG_PRINT("[Chare %d] bcast send check: name=%d found=%d local_size=%d\n", - partition->index[0], bcast.source_name, - (int)(arr_it != partition->arrays.end()), - (arr_it != partition->arrays.end()) ? arr_it->second->local_size() : -1); - if (arr_it != partition->arrays.end() && arr_it->second->local_size() > 0) { - int elem_sz = arr_it->second->elem_size(); -#ifdef USE_KOKKOS - char host_scalar[8]; - Kokkos::deep_copy( - Kokkos::View>( - host_scalar, elem_sz), - Kokkos::View>( - static_cast(arr_it->second->device_data_ptr()), elem_sz)); - char* val_ptr = host_scalar; -#else - arr_it->second->copyToHost(); - char* val_ptr = static_cast(arr_it->second->data_ptr()); -#endif - - if (bcast.source_ndims == bcast.target_ndims) { - // Same-partition broadcast: send to all OTHER chares - has_local_broadcast = true; - auto pg_it = dag_group->partition_grid.find(N); - if (pg_it != dag_group->partition_grid.end()) { - int total_chares = 1; - for (int d = 0; d < N; ++d) - total_chares *= pg_it->second.grid[d]; - for (int t = 0; t < total_chares; ++t) { - std::array tgt_idx = {}; - int rem = t; - for (int d = N - 1; d >= 0; --d) { - tgt_idx[d] = rem % pg_it->second.grid[d]; - rem /= pg_it->second.grid[d]; - } - // Skip self - bool is_self = true; - for (int d = 0; d < N; ++d) - if (tgt_idx[d] != nd_idx[d]) - is_self = false; - if (is_self) - continue; - send_broadcast(N, tgt_idx, bcast_input_idx, bcast.source_name, - val_ptr, elem_sz); - } - } - } else { - // Cross-partition broadcast: send to all chares on target partition - int tgt_nd = bcast.target_ndims; - auto pg_it = dag_group->partition_grid.find(tgt_nd); - if (pg_it != dag_group->partition_grid.end()) { - int total_chares = 1; - for (int d = 0; d < tgt_nd; ++d) - total_chares *= pg_it->second.grid[d]; - for (int t = 0; t < total_chares; ++t) { - std::array tgt_idx = {}; - int rem = t; - for (int d = tgt_nd - 1; d >= 0; --d) { - tgt_idx[d] = rem % pg_it->second.grid[d]; - rem /= pg_it->second.grid[d]; - } - send_broadcast(tgt_nd, tgt_idx, bcast_input_idx, - bcast.source_name, val_ptr, elem_sz); - } - } - } - } - } - - if (N == bcast.target_ndims) - is_source_only = false; - } - - // If this partition is only the broadcast source (cross-partition), we're done - if (is_source_only) { - node_finished(node->id); - return; - } - - // Count expected broadcast messages for this chare - // For same-partition: chares that DON'T have the source array expect 1 message per broadcast - // For cross-partition: all target chares expect 1 message per broadcast - int n_expected_broadcasts = 0; - for (auto& bcast : broadcast_ops) { - if (N != bcast.target_ndims) - continue; - if (bcast.source_ndims == bcast.target_ndims) { - // Same-partition: only expect a message if we DON'T have the array locally - bool have_it = partition->arrays.count(bcast.source_name) > 0; - DBG_PRINT("[Chare %d] bcast count: name=%d src_nd=%d tgt_nd=%d has_local_broadcast=%d have_it=%d\n", - partition->index[0], bcast.source_name, bcast.source_ndims, bcast.target_ndims, - (int)has_local_broadcast, (int)have_it); - if (!has_local_broadcast) - n_expected_broadcasts++; - } else { - // Cross-partition: we always expect a message - n_expected_broadcasts++; - } - } - DBG_PRINT("[Chare %d] n_expected_broadcasts=%d is_source_only=%d\n", - partition->index[0], n_expected_broadcasts, (int)is_source_only); - - if (n_expected_broadcasts > 0) { - auto pre_it = pending.find(node->id); - int pre_arrived = 0; - std::unordered_map>> pre_buffers; - if (pre_it != pending.end()) { - pre_arrived = -(pre_it->second.expected_msgs); - pre_buffers = std::move(pre_it->second.remote_buffers); - } - - int remaining = n_expected_broadcasts - pre_arrived; - pending[node->id] = {node, remaining, {}, std::move(pre_buffers)}; - - if (remaining <= 0) { - on_comm_done(node->id); - } else { - DBG_PRINT("[Chare %d] -> waiting for %d broadcast values\n", - partition->index[0], remaining); - } - return; - } - // If no messages expected (source chare in same-partition broadcast), - // fall through to normal execution — broadcast value is already local. - } - - ArrayRegion r_out; - std::vector> input_regions; - std::vector input_source_names; - std::array global_shape; - - DType dt = determine_dtype(node, partition->arrays); - - if (!extract_regions_nd(node, partition->arrays, r_out, input_regions, input_source_names, - global_shape) || - input_regions.empty()) { - dispatch_execute(dt, static_cast(group), node, partition, nullptr); -#ifdef USE_NVIDIA - { - auto* p = new ComputeDoneParam{partition, node->id, false}; - CkCallback hcb(compute_done_cb, p); - hapiAddCallback(partition->compute_stream_raw, &hcb); - } -#else - CT_KOKKOS_FENCE(); - node_finished(node->id); -#endif - return; - } - - // Build per-input decomps and output decomp from array metadata - auto* dag_group = static_cast(group); - std::vector> input_decomps; - for (int i = 0; i < (int)input_source_names.size(); ++i) { - auto arr_it = partition->arrays.find(input_source_names[i]); - if (arr_it != partition->arrays.end()) { - input_decomps.push_back(arr_it->second->decomp); - } else { - auto meta_it = dag_group->array_meta.find(input_source_names[i]); - if (meta_it != dag_group->array_meta.end()) - input_decomps.push_back(meta_it->second.template decomp()); - else - input_decomps.push_back(ArrayDecomp::default_decomp(global_shape, - array_tile(dag_group->array_meta, input_source_names[i], N))); - - } - } - - // Output decomp: use the first non-temp root output, not the first AST root. - // Fused ASTs often keep temporary internal roots ahead of the real result, - // and those temps can legitimately have a different decomposition offset. - ArrayDecomp output_decomp; - { - int result_name = root_output_name(node->ast->roots[0]); - for (ASTNode* root : node->ast->roots) { - auto opc = static_cast(root->opcode); - if (opc != Opcode::NOOP && opc != Opcode::CREATE && !root->is_temp) { - result_name = root_output_name(root); - break; - } - } - auto arr_it = partition->arrays.find(result_name); - if (arr_it != partition->arrays.end()) { - output_decomp = arr_it->second->decomp; - } else { - auto meta_it = dag_group->array_meta.find(result_name); - if (meta_it != dag_group->array_meta.end()) - output_decomp = meta_it->second.template decomp(); - else - output_decomp = ArrayDecomp::default_decomp(global_shape, - array_tile(dag_group->array_meta, result_name, N)); - - } - } - - // Translate regions from local space to global space - ArrayRegion r_out_global = output_decomp.to_global(r_out); - std::vector> input_regions_global; - for (int i = 0; i < (int)input_regions.size(); ++i) - input_regions_global.push_back(input_decomps[i].to_global(input_regions[i])); - - ArrayRegion r_chare_out = output_decomp.chare_region_global(nd_idx); - - // Send remote inputs (type-dispatched) - auto sends = send_remote_inputs(r_out_global, input_regions_global, nd_idx, output_decomp, - input_decomps); - for (auto& send : sends) { - int inp_name = input_source_names[send.input_index]; - dispatch_send(dt, partition, node, send, inp_name, nd_idx); - } - - // Determine expected messages - auto li = local_inputs(r_out_global, input_regions_global, nd_idx, output_decomp, - input_decomps); - - if (li.expected_msgs == 0) { - if (li.my_inputs.empty()) { - auto [r_myout, has_out] = intersect(r_out_global, r_chare_out); - if (has_out && r_myout.size() > 0) { - PendingComm comm_local = {node, 0, {}, {}}; - dispatch_execute(dt, static_cast(group), node, partition, - &comm_local); - } -#ifdef USE_NVIDIA - { - auto* p = new ComputeDoneParam{partition, node->id, false}; - CkCallback hcb(compute_done_cb, p); - hapiAddCallback(partition->compute_stream_raw, &hcb); - } -#else - CT_KOKKOS_FENCE(); - node_finished(node->id); -#endif - } else { - PendingComm comm_local = {node, 0, std::move(li.my_inputs), {}}; - dispatch_execute(dt, static_cast(group), node, partition, - &comm_local); -#ifdef USE_NVIDIA - { - auto* p = new ComputeDoneParam{partition, node->id, false}; - CkCallback hcb(compute_done_cb, p); - hapiAddCallback(partition->compute_stream_raw, &hcb); - } -#else - CT_KOKKOS_FENCE(); - node_finished(node->id); -#endif - } - } else { - auto pre_it = pending.find(node->id); - int pre_arrived = 0; - std::unordered_map>> pre_buffers; - if (pre_it != pending.end()) { - pre_arrived = -(pre_it->second.expected_msgs); - pre_buffers = std::move(pre_it->second.remote_buffers); - } - - int remaining = li.expected_msgs - pre_arrived; - pending[node->id] = {node, remaining, std::move(li.my_inputs), std::move(pre_buffers)}; - - if (remaining <= 0) { - on_comm_done(node->id); - } else { - DBG_PRINT("[Chare %d] -> waiting for %d remote messages\n", partition->index[0], - remaining); - } - } -} - -template -void ArrayDAGExecutorND::on_comm_done(int node_id) { - auto it = pending.find(node_id); - if (it == pending.end()) - return; - - PendingComm& comm = it->second; - DAGNode* node = comm.node; - DType dt = determine_dtype(node, partition->arrays); - dispatch_execute(dt, static_cast(group), node, partition, &comm); - -#ifdef USE_NVIDIA - // Defer cleanup + node_finished until GPU work completes on compute_stream - { - auto* p = new ComputeDoneParam{partition, node_id, true}; - CkCallback hcb(compute_done_cb, p); - hapiAddCallback(partition->compute_stream_raw, &hcb); - } -#else - CT_KOKKOS_FENCE(); - comm.clear_remote_buffers(); - pending.erase(it); - // on_comm_done is high-volume; node_finished already prints - node_finished(node_id); -#endif -} - -template void ArrayDAGExecutorND<1>::delete_array(int); -template void ArrayDAGExecutorND<2>::delete_array(int); -template void ArrayDAGExecutorND<3>::delete_array(int); - -template void ArrayDAGExecutorND<1>::execute_dag_node(DAGNode*); -template void ArrayDAGExecutorND<2>::execute_dag_node(DAGNode*); -template void ArrayDAGExecutorND<3>::execute_dag_node(DAGNode*); - -template void ArrayDAGExecutorND<1>::on_comm_done(int); -template void ArrayDAGExecutorND<2>::on_comm_done(int); -template void ArrayDAGExecutorND<3>::on_comm_done(int); diff --git a/src/executor_incremental.cpp b/src/executor_incremental.cpp deleted file mode 100644 index 2baf3fc..0000000 --- a/src/executor_incremental.cpp +++ /dev/null @@ -1,181 +0,0 @@ -#include "backend_internal.hpp" - -#include -#include - -#ifdef USE_KOKKOS -#include -#endif - -template -void ArrayDAGExecutorND::on_matmatmul_receive(int node_id, int input_index) { - if constexpr (N != 2) - return; - - auto it = pending.find(node_id); - if (it == pending.end()) - return; - PendingComm& comm = it->second; - - // The just-arrived panel is the last element in remote_buffers[input_index] - auto& arrived_bufs = comm.remote_buffers[input_index]; - if (arrived_bufs.empty()) - return; - auto& new_panel = arrived_bufs.back(); - - // Scan all panels on the opposite side for k-overlap - int other_side = 1 - input_index; - auto other_it = comm.remote_buffers.find(other_side); - if (other_it == comm.remote_buffers.end() || other_it->second.empty()) - return; - - DAGNode* node = comm.node; - ASTNode* mm_root = nullptr; - for (ASTNode* root : node->ast->roots) { - if (static_cast(root->opcode) == Opcode::MATMATMUL) { - mm_root = root; - break; - } - } - if (!mm_root) - return; - - int result_name = mm_root->result_name; - auto arr_it = partition->arrays.find(result_name); - if (arr_it == partition->arrays.end()) - return; - - DType dt = arr_it->second->dtype; - auto nd_idx = partition->nd_index(); - int c_row_lo = arr_it->second->decomp.chare_start_global(0, nd_idx[0]); - int c_col_lo = arr_it->second->decomp.chare_start_global(1, nd_idx[1]); - int sub_rows = arr_it->second->region.size(0); - int sub_cols = arr_it->second->region.size(1); - - // Dispatch by dtype for typed computation - auto compute_pair = [&](auto* dummy) { - using T = std::remove_pointer_t; - - for (auto& other_panel : other_it->second) { - // Determine which is A (input_index=0) and which is B (input_index=1) - auto& a_buf = (input_index == 0) ? new_panel : other_panel; - auto& b_buf = (input_index == 0) ? other_panel : new_panel; - - int a_k_start = a_buf.region.start[0]; - int a_k_end = a_buf.region.stop[0]; - int a_c_row_start = a_buf.region.start[1]; - int a_c_row_end = a_buf.region.stop[1]; - int a_rows = a_c_row_end - a_c_row_start; - int a_k_size = a_k_end - a_k_start; - - int b_k_start = b_buf.region.start[0]; - int b_k_end = b_buf.region.stop[0]; - int b_c_col_start = b_buf.region.start[1]; - int b_c_col_end = b_buf.region.stop[1]; - int b_cols = b_c_col_end - b_c_col_start; - - // Check k-range overlap - int k_lo = std::max(a_k_start, b_k_start); - int k_hi = std::min(a_k_end, b_k_end); - if (k_lo >= k_hi) - continue; - - int c_local_row = a_c_row_start - c_row_lo; - int c_local_col = b_c_col_start - c_col_lo; - - T* a_data = reinterpret_cast(a_buf.data); - T* b_data = reinterpret_cast(b_buf.data); - - int k_size = k_hi - k_lo; - int a_col_off = k_lo - a_k_start; - int b_row_off = k_lo - b_k_start; - -#ifdef USE_KOKKOS - // Device-side GEMM: rb.data is a device pointer under USE_KOKKOS - T* c_device = static_cast(arr_it->second->device_data_ptr()); - - Kokkos::View> - d_A_full(a_data, a_rows, a_k_size); - auto d_A_sub = Kokkos::subview(d_A_full, Kokkos::ALL, - Kokkos::make_pair(a_col_off, a_col_off + k_size)); - - int b_k_size = b_k_end - b_k_start; - Kokkos::View> - d_B_full(b_data, b_k_size, b_cols); - auto d_B_sub = Kokkos::subview(d_B_full, - Kokkos::make_pair(b_row_off, b_row_off + k_size), - Kokkos::ALL); - - Kokkos::View> - d_C_full(c_device, sub_rows, sub_cols); - auto d_C_sub = Kokkos::subview(d_C_full, - Kokkos::make_pair(c_local_row, c_local_row + a_rows), - Kokkos::make_pair(c_local_col, c_local_col + b_cols)); - - KokkosBlas::gemm("N", "N", T(1), d_A_sub, d_B_sub, T(1), d_C_sub); -#else - T* c_data = static_cast(arr_it->second->data_ptr()); - T* c_sub = c_data + c_local_row * sub_cols + c_local_col; - - using RMat = Eigen::Matrix; - using CStride = Eigen::Stride; - - Eigen::Map - A_map(a_data + a_col_off, a_rows, k_size, CStride(a_k_size, 1)); - Eigen::Map - B_map(b_data + b_row_off * b_cols, k_size, b_cols, CStride(b_cols, 1)); - Eigen::Map - C_map(c_sub, a_rows, b_cols, CStride(sub_cols, 1)); - C_map.noalias() += A_map * B_map; -#endif - } - }; - - switch (dt) { - case DType::FLOAT32: { - float* d = nullptr; - compute_pair(d); - break; - } - case DType::FLOAT64: { - double* d = nullptr; - compute_pair(d); - break; - } - case DType::INT32: { - int32_t* d = nullptr; - compute_pair(d); - break; - } - case DType::INT64: { - int64_t* d = nullptr; - compute_pair(d); - break; - } - } -} - -template -void ArrayDAGExecutorND::on_matmul_partial(int node_id, RemoteBuffer& partial_buf) { - // No longer used — partials are accumulated on 1D partition directly -} - -template -void ArrayDAGExecutorND::matmul_check_finalize(int node_id) { - // No longer used — partials are accumulated on 1D partition directly -} - -template void ArrayDAGExecutorND<1>::on_matmatmul_receive(int, int); -template void ArrayDAGExecutorND<2>::on_matmatmul_receive(int, int); -template void ArrayDAGExecutorND<3>::on_matmatmul_receive(int, int); - -template void ArrayDAGExecutorND<1>::on_matmul_partial(int, RemoteBuffer<1>&); -template void ArrayDAGExecutorND<2>::on_matmul_partial(int, RemoteBuffer<2>&); -template void ArrayDAGExecutorND<3>::on_matmul_partial(int, RemoteBuffer<3>&); - -template void ArrayDAGExecutorND<1>::matmul_check_finalize(int); -template void ArrayDAGExecutorND<2>::matmul_check_finalize(int); -template void ArrayDAGExecutorND<3>::matmul_check_finalize(int); diff --git a/src/executor_matmatmul.cpp b/src/executor_matmatmul.cpp deleted file mode 100644 index 74a1ca0..0000000 --- a/src/executor_matmatmul.cpp +++ /dev/null @@ -1,822 +0,0 @@ -#include "backend_internal.hpp" -#include "dispatch.hpp" - -#include - -template -static void cross_matmatmul_send_panel_2d( - PartitionImpl<2>* partition, DAGNode* node, int arr_name, int input_index, - const ChareIndex<2>& target_ci, T* local_data, int local_cols, - int row_offset, int col_offset, int sub_rows, int sub_cols, - int k_start, int k_end, int c_pos, - CProxy_Partition2D& proxy_2d) { - int64_t send_count = (int64_t)sub_rows * sub_cols; - int64_t byte_size = send_count * sizeof(T); - -#ifndef NDEBUG - partition->comm_bytes_sent += byte_size; -#endif - - // For A panels (input_index=0): dim 1 = [c_pos, c_pos + sub_rows) = C-space row range - // For B panels (input_index=1): dim 1 = [c_pos, c_pos + sub_cols) = C-space col range - int region_data[2 * 3]; - region_data[0] = k_start; - region_data[1] = k_end; - region_data[2] = 1; - region_data[3] = c_pos; - region_data[4] = c_pos + ((input_index == 0) ? sub_rows : sub_cols); - region_data[5] = 1; - -#ifndef USE_KOKKOS - // Pack sub-block into contiguous send buffer - T* send_buf = new T[send_count]; - for (int r = 0; r < sub_rows; r++) - memcpy(send_buf + r * sub_cols, - local_data + (row_offset + r) * local_cols + col_offset, - sub_cols * sizeof(T)); - proxy_at<2>(proxy_2d, target_ci) - .receive_data(node->id, input_index, arr_name, 2, region_data, byte_size, - reinterpret_cast(send_buf)); - delete[] send_buf; -#else - T* send_buf = static_cast(Kokkos::kokkos_malloc(byte_size)); - T* src = local_data + row_offset * local_cols + col_offset; - int lc = local_cols; - int sc = sub_cols; - Kokkos::parallel_for( - CT_COMM_POLICY(partition, sub_rows), - KOKKOS_LAMBDA(int r) { - for (int c = 0; c < sc; c++) - send_buf[r * sc + c] = src[r * lc + c]; - }); - device_pack_send<2>(partition, proxy_2d, target_ci, node->id, input_index, arr_name, - region_data, byte_size, send_buf); -#endif -} - -static void dispatch_cross_matmatmul_send_panel_2d( - DType dt, PartitionImpl<2>* partition, DAGNode* node, int arr_name, int input_index, - const ChareIndex<2>& target_ci, int local_cols, - int row_offset, int col_offset, int sub_rows, int sub_cols, - int k_start, int k_end, int c_pos, - CProxy_Partition2D& proxy_2d) { - // Under USE_KOKKOS the packing lambda operates on device, so pass device_data_ptr. - // Under CPU-only, pass host data_ptr. - auto arr_ptr = [&](auto* dummy) -> void* { - (void)dummy; -#ifdef USE_KOKKOS - return partition->arrays[arr_name]->device_data_ptr(); -#else - return partition->arrays[arr_name]->data_ptr(); -#endif - }; - switch (dt) { - case DType::FLOAT32: - cross_matmatmul_send_panel_2d( - partition, node, arr_name, input_index, target_ci, - static_cast(arr_ptr((float*)nullptr)), - local_cols, row_offset, col_offset, sub_rows, sub_cols, - k_start, k_end, c_pos, proxy_2d); - break; - case DType::FLOAT64: - cross_matmatmul_send_panel_2d( - partition, node, arr_name, input_index, target_ci, - static_cast(arr_ptr((double*)nullptr)), - local_cols, row_offset, col_offset, sub_rows, sub_cols, - k_start, k_end, c_pos, proxy_2d); - break; - case DType::INT32: - cross_matmatmul_send_panel_2d( - partition, node, arr_name, input_index, target_ci, - static_cast(arr_ptr((int32_t*)nullptr)), - local_cols, row_offset, col_offset, sub_rows, sub_cols, - k_start, k_end, c_pos, proxy_2d); - break; - case DType::INT64: - cross_matmatmul_send_panel_2d( - partition, node, arr_name, input_index, target_ci, - static_cast(arr_ptr((int64_t*)nullptr)), - local_cols, row_offset, col_offset, sub_rows, sub_cols, - k_start, k_end, c_pos, proxy_2d); - break; - } -} - -template -void ArrayDAGExecutorND::execute_matmatmul_node(DAGNode* node) { - ASTNode* mm_root = nullptr; - for (ASTNode* root : node->ast->roots) { - if (static_cast(root->opcode) == Opcode::MATMATMUL) { - mm_root = root; - break; - } - } - int a_name = mm_root->operands[0]->result_name; - int b_name = mm_root->operands[1]->result_name; - int result_name = mm_root->result_name; - - auto* dag_group = static_cast(group); - - if constexpr (N == 2) { - int tile_2d = array_tile(dag_group->array_meta, a_name, 2); - - auto nd_idx = partition->nd_index(); - - DType dt = determine_dtype(node, partition->arrays); - - // Look up operand metadata - auto a_meta_it = dag_group->array_meta.find(a_name); - auto b_meta_it = dag_group->array_meta.find(b_name); - auto c_meta_it = dag_group->array_meta.find(result_name); - - // Extract A and B slice regions (or full range) - int a_rs = 0, a_re, a_cs = 0, a_ce; - int b_rs = 0, b_re, b_cs = 0, b_ce; - if (a_meta_it != dag_group->array_meta.end()) { - a_re = a_meta_it->second.global_shape[0]; - a_ce = a_meta_it->second.global_shape[1]; - } else { - CkAbort("MATMATMUL: could not find A metadata for name=%d", a_name); - return; - } - if (b_meta_it != dag_group->array_meta.end()) { - b_re = b_meta_it->second.global_shape[0]; - b_ce = b_meta_it->second.global_shape[1]; - } else { - CkAbort("MATMATMUL: could not find B metadata for name=%d", b_name); - return; - } - - if (mm_root->operand_regions.size() >= 2) { - // Extract effective 2D ranges from operand regions. - // Operands may be 3D (dimension-dropped), so check ndims. - int a_ndims = a_meta_it->second.ndims; - int b_ndims = b_meta_it->second.ndims; - - if (a_ndims == 3) { - auto* ar = static_cast*>(mm_root->operand_regions[0]); - int dd = -1; - for (int d = 0; d < 3; d++) - if (ar->stop[d] - ar->start[d] == 1) { dd = d; break; } - int rd = (dd == 0) ? 1 : 0; - int cd = (dd <= 1) ? 2 : 1; - a_rs = ar->start[rd]; a_re = ar->stop[rd]; - a_cs = ar->start[cd]; a_ce = ar->stop[cd]; - } else { - auto* ar = static_cast*>(mm_root->operand_regions[0]); - a_rs = ar->start[0]; a_re = ar->stop[0]; - a_cs = ar->start[1]; a_ce = ar->stop[1]; - } - - if (b_ndims == 3) { - auto* br = static_cast*>(mm_root->operand_regions[1]); - int dd = -1; - for (int d = 0; d < 3; d++) - if (br->stop[d] - br->start[d] == 1) { dd = d; break; } - int rd = (dd == 0) ? 1 : 0; - int cd = (dd <= 1) ? 2 : 1; - b_rs = br->start[rd]; b_re = br->stop[rd]; - b_cs = br->start[cd]; b_ce = br->stop[cd]; - } else { - auto* br = static_cast*>(mm_root->operand_regions[1]); - b_rs = br->start[0]; b_re = br->stop[0]; - b_cs = br->start[1]; b_ce = br->stop[1]; - } - } - - // Translate A/B regions to global space - if (mm_root->operand_regions.size() >= 2) { - int a_ndims = a_meta_it->second.ndims; - if (a_ndims == 3) { - auto ad = a_meta_it->second.template decomp<3>(); - auto* ar3 = static_cast*>(mm_root->operand_regions[0]); - int dd = -1; - for (int d = 0; d < 3; d++) - if (ar3->stop[d] - ar3->start[d] == 1) { dd = d; break; } - int rd = (dd == 0) ? 1 : 0, cd = (dd <= 1) ? 2 : 1; - a_rs += ad.offset[rd]; a_re += ad.offset[rd]; - a_cs += ad.offset[cd]; a_ce += ad.offset[cd]; - } else { - auto ad = a_meta_it->second.template decomp<2>(); - a_rs += ad.offset[0]; a_re += ad.offset[0]; - a_cs += ad.offset[1]; a_ce += ad.offset[1]; - } - int b_ndims = b_meta_it->second.ndims; - if (b_ndims == 3) { - auto bd = b_meta_it->second.template decomp<3>(); - auto* br3 = static_cast*>(mm_root->operand_regions[1]); - int dd = -1; - for (int d = 0; d < 3; d++) - if (br3->stop[d] - br3->start[d] == 1) { dd = d; break; } - int rd = (dd == 0) ? 1 : 0, cd = (dd <= 1) ? 2 : 1; - b_rs += bd.offset[rd]; b_re += bd.offset[rd]; - b_cs += bd.offset[cd]; b_ce += bd.offset[cd]; - } else { - auto bd = b_meta_it->second.template decomp<2>(); - b_rs += bd.offset[0]; b_re += bd.offset[0]; - b_cs += bd.offset[1]; b_ce += bd.offset[1]; - } - } else { - // No operand regions — both A and B are 2D - auto ad = a_meta_it->second.template decomp<2>(); - a_rs += ad.offset[0]; a_re += ad.offset[0]; - a_cs += ad.offset[1]; a_ce += ad.offset[1]; - auto bd = b_meta_it->second.template decomp<2>(); - b_rs += bd.offset[0]; b_re += bd.offset[0]; - b_cs += bd.offset[1]; b_ce += bd.offset[1]; - } - - int M = a_re - a_rs; // result rows - int K = a_ce - a_cs; // shared dimension - int N_cols = b_ce - b_cs; // result cols - - // C result range in global coords - ArrayDecomp<2> c_decomp; - if (c_meta_it != dag_group->array_meta.end()) - c_decomp = c_meta_it->second.template decomp<2>(); - int c_rs = c_decomp.offset[0], c_re = c_decomp.offset[0] + M; - int c_cs = c_decomp.offset[1], c_ce = c_decomp.offset[1] + N_cols; - - // Chare ranges for A, B, and C - int first_a_row_chare = a_rs / tile_2d; - int last_a_row_chare = (a_re > 0) ? (a_re - 1) / tile_2d : 0; - int first_a_col_chare = a_cs / tile_2d; - int last_a_col_chare = (a_ce > 0) ? (a_ce - 1) / tile_2d : 0; - int first_b_row_chare = b_rs / tile_2d; - int last_b_row_chare = (b_re > 0) ? (b_re - 1) / tile_2d : 0; - int first_b_col_chare = b_cs / tile_2d; - int last_b_col_chare = (b_ce > 0) ? (b_ce - 1) / tile_2d : 0; - int first_c_row_chare = c_rs / tile_2d; - int last_c_row_chare = (c_re > 0) ? (c_re - 1) / tile_2d : 0; - int first_c_col_chare = c_cs / tile_2d; - int last_c_col_chare = (c_ce > 0) ? (c_ce - 1) / tile_2d : 0; - - // --- Phase 1: Send A panels --- - auto a_it = partition->arrays.find(a_name); - if (a_it != partition->arrays.end() && a_it->second->local_size() > 0) { - int local_cols = a_it->second->region.size(1); - - // Overlap of this chare's tile with A's slice region - auto a_decomp = a_it->second->decomp; - auto a_chare = a_decomp.chare_region_global(nd_idx); - int a_row_lo = std::max(a_chare.start[0], a_rs); - int a_row_hi = std::min(a_chare.stop[0], a_re); - int a_col_lo = std::max(a_chare.start[1], a_cs); - int a_col_hi = std::min(a_chare.stop[1], a_ce); - - if (a_row_lo < a_row_hi && a_col_lo < a_col_hi) { - int sub_rows = a_row_hi - a_row_lo; - int sub_cols = a_col_hi - a_col_lo; - int row_offset = a_row_lo - a_chare.start[0]; - int col_offset = a_col_lo - a_chare.start[1]; - - // k-range in K-space: [a_col_lo - a_cs, a_col_hi - a_cs) - int k_start = a_col_lo - a_cs; - int k_end = a_col_hi - a_cs; - - // Which result row chare does this A block map to? - int c_row_start = c_rs + (a_row_lo - a_rs); - int c_row_end = c_rs + (a_row_hi - a_rs); - int first_dest_row_chare = c_row_start / tile_2d; - int last_dest_row_chare = (c_row_end - 1) / tile_2d; - -#ifndef USE_KOKKOS - a_it->second->copyToHost(); -#endif - for (int tj = first_c_col_chare; tj <= last_c_col_chare; tj++) { - for (int ti = first_dest_row_chare; ti <= last_dest_row_chare; ti++) { - ChareIndex<2> target_ci; - target_ci.idx[0] = ti; - target_ci.idx[1] = tj; - - int dest_row_start = std::max(c_row_start, ti * tile_2d); - int dest_row_end = std::min(c_row_end, (ti + 1) * tile_2d); - int send_rows = dest_row_end - dest_row_start; - int send_row_offset = row_offset + (dest_row_start - c_row_start); - - dispatch_cross_matmatmul_send_panel_2d( - dt, partition, node, a_name, /*input_index=*/0, - target_ci, local_cols, - send_row_offset, col_offset, send_rows, sub_cols, - k_start, k_end, /*c_pos=*/dest_row_start, - dag_group->partition_proxy_2); - } - } - } - } - - // --- Phase 2: Send B panels --- - auto b_it = partition->arrays.find(b_name); - if (b_it != partition->arrays.end() && b_it->second->local_size() > 0) { - int local_cols = b_it->second->region.size(1); - - auto b_decomp = b_it->second->decomp; - auto b_chare = b_decomp.chare_region_global(nd_idx); - int b_row_lo = std::max(b_chare.start[0], b_rs); - int b_row_hi = std::min(b_chare.stop[0], b_re); - int b_col_lo = std::max(b_chare.start[1], b_cs); - int b_col_hi = std::min(b_chare.stop[1], b_ce); - - if (b_row_lo < b_row_hi && b_col_lo < b_col_hi) { - int sub_rows = b_row_hi - b_row_lo; - int sub_cols = b_col_hi - b_col_lo; - int row_offset = b_row_lo - b_chare.start[0]; - int col_offset = b_col_lo - b_chare.start[1]; - - // k-range in K-space: [b_row_lo - b_rs, b_row_hi - b_rs) - int k_start = b_row_lo - b_rs; - int k_end = b_row_hi - b_rs; - - // B cols [b_col_lo, b_col_hi) map to C cols in global C space - int c_col_start = c_cs + (b_col_lo - b_cs); - int c_col_end = c_cs + (b_col_hi - b_cs); - int first_dest_col_chare = c_col_start / tile_2d; - int last_dest_col_chare = (c_col_end - 1) / tile_2d; - -#ifndef USE_KOKKOS - b_it->second->copyToHost(); -#endif - for (int ti = first_c_row_chare; ti <= last_c_row_chare; ti++) { - for (int tj = first_dest_col_chare; tj <= last_dest_col_chare; tj++) { - ChareIndex<2> target_ci; - target_ci.idx[0] = ti; - target_ci.idx[1] = tj; - - int dest_col_start = std::max(c_col_start, tj * tile_2d); - int dest_col_end = std::min(c_col_end, (tj + 1) * tile_2d); - int send_cols = dest_col_end - dest_col_start; - int send_col_offset = col_offset + (dest_col_start - c_col_start); - - dispatch_cross_matmatmul_send_panel_2d( - dt, partition, node, b_name, /*input_index=*/1, - target_ci, local_cols, - row_offset, send_col_offset, sub_rows, send_cols, - k_start, k_end, /*c_pos=*/dest_col_start, - dag_group->partition_proxy_2); - } - } - } - } - - // --- Phase 3: Allocate C, count expected messages, enable incremental --- - auto c_chare = c_decomp.chare_region_global(nd_idx); - int my_c_row_lo = std::max(c_chare.start[0], c_rs); - int my_c_row_hi = std::min(c_chare.stop[0], c_re); - int my_c_col_lo = std::max(c_chare.start[1], c_cs); - int my_c_col_hi = std::min(c_chare.stop[1], c_ce); - - int total_expected = 0; - if (my_c_row_lo < my_c_row_hi && my_c_col_lo < my_c_col_hi) { - int sub_rows = my_c_row_hi - my_c_row_lo; - int sub_cols = my_c_col_hi - my_c_col_lo; - - // Allocate C upfront so incremental receives can accumulate into it - if (partition->arrays.find(result_name) == partition->arrays.end()) { - std::array out_start = {0, 0}; - std::array out_stop = {sub_rows, sub_cols}; - std::array out_step = {1, 1}; - std::array out_gs = {M, N_cols}; - ArrayRegion<2> out_region(out_start, out_stop, out_step); - ArrayDecomp<2> out_decomp = dag_group->array_meta[result_name].template decomp<2>(); - partition->arrays[result_name] = - partition->allocate_or_reuse(out_region, out_gs, result_name, dt, out_decomp); - } - - int a_k_chares = 0; - for (int tk = first_a_col_chare; tk <= last_a_col_chare; tk++) { - int a_row_for_c_lo = a_rs + my_c_row_lo; - int a_row_for_c_hi = a_rs + my_c_row_hi; - int fa = a_row_for_c_lo / tile_2d; - int la = (a_row_for_c_hi - 1) / tile_2d; - a_k_chares += (la - fa + 1); - } - - int b_k_chares = 0; - for (int tk = first_b_row_chare; tk <= last_b_row_chare; tk++) { - int b_col_for_c_lo = b_cs + my_c_col_lo; - int b_col_for_c_hi = b_cs + my_c_col_hi; - int fb = b_col_for_c_lo / tile_2d; - int lb = (b_col_for_c_hi - 1) / tile_2d; - b_k_chares += (lb - fb + 1); - } - - total_expected = a_k_chares + b_k_chares; - } - - auto pre_it = pending.find(node->id); - int pre_arrived = 0; - std::unordered_map>> pre_buffers; - if (pre_it != pending.end()) { - pre_arrived = -(pre_it->second.expected_msgs); - pre_buffers = std::move(pre_it->second.remote_buffers); - } - int remaining = total_expected - pre_arrived; - pending[node->id] = {node, remaining, {}, std::move(pre_buffers), /*incremental=*/true}; - - DBG_PRINT("[2D Chare %d] matmatmul node=%d: A=[%d:%d,%d:%d] B=[%d:%d,%d:%d] " - "C_tile=[%d:%d,%d:%d] total_expected=%d remaining=%d\n", - partition->index[0], node->id, - a_rs, a_re, a_cs, a_ce, b_rs, b_re, b_cs, b_ce, - my_c_row_lo, my_c_row_hi, my_c_col_lo, my_c_col_hi, - total_expected, remaining); - - // Process any pre-arrived panels that came before incremental mode was set - if (pre_arrived > 0) { - PendingComm<2>& comm = pending[node->id]; - auto a_it2 = comm.remote_buffers.find(0); - auto b_it2 = comm.remote_buffers.find(1); - if (a_it2 != comm.remote_buffers.end() && b_it2 != comm.remote_buffers.end()) { - auto& a_bufs = a_it2->second; - auto& b_bufs = b_it2->second; - auto arr_it = partition->arrays.find(result_name); - if (arr_it != partition->arrays.end()) { - auto nd_idx2 = partition->nd_index(); - int c_row_lo2 = c_decomp.chare_start_global(0, nd_idx2[0]); - int c_col_lo2 = c_decomp.chare_start_global(1, nd_idx2[1]); - int sub_cols2 = arr_it->second->region.size(1); - - auto compute_pre = [&](auto* dummy) { - using T = std::remove_pointer_t; - T* c_data = static_cast(arr_it->second->data_ptr()); - using RMat = Eigen::Matrix; - using CStride = Eigen::Stride; - for (auto& a_buf : a_bufs) { - for (auto& b_buf : b_bufs) { - int a_k_s = a_buf.region.start[0], a_k_e = a_buf.region.stop[0]; - int b_k_s = b_buf.region.start[0], b_k_e = b_buf.region.stop[0]; - int k_lo = std::max(a_k_s, b_k_s); - int k_hi = std::min(a_k_e, b_k_e); - if (k_lo >= k_hi) - continue; - int a_rows = a_buf.region.stop[1] - a_buf.region.start[1]; - int a_k_size = a_k_e - a_k_s; - int b_cols = b_buf.region.stop[1] - b_buf.region.start[1]; - int c_lr = a_buf.region.start[1] - c_row_lo2; - int c_lc = b_buf.region.start[1] - c_col_lo2; - T* a_data = reinterpret_cast(a_buf.data); - T* b_data = reinterpret_cast(b_buf.data); - T* c_sub = c_data + c_lr * sub_cols2 + c_lc; - int k_size = k_hi - k_lo; - int a_co = k_lo - a_k_s; - int b_ro = k_lo - b_k_s; - Eigen::Map - A_map(a_data + a_co, a_rows, k_size, CStride(a_k_size, 1)); - Eigen::Map - B_map(b_data + b_ro * b_cols, k_size, b_cols, CStride(b_cols, 1)); - Eigen::Map - C_map(c_sub, a_rows, b_cols, CStride(sub_cols2, 1)); - C_map.noalias() += A_map * B_map; - } - } - }; - switch (dt) { - case DType::FLOAT32: { float* d = nullptr; compute_pre(d); break; } - case DType::FLOAT64: { double* d = nullptr; compute_pre(d); break; } - case DType::INT32: { int32_t* d = nullptr; compute_pre(d); break; } - case DType::INT64: { int64_t* d = nullptr; compute_pre(d); break; } - } - } - } - } - - if (remaining <= 0) - on_comm_done(node->id); - - } else if constexpr (N == 3) { - // 3D partition: one operand is a dimension-dropped 3D array. - int tile_2d = array_tile(dag_group->array_meta, result_name, 2); - - auto nd_idx = partition->nd_index(); - DType dt = determine_dtype(node, partition->arrays); - - // Determine which operand (0=A, 1=B) is the 3D one - for (int op_idx = 0; op_idx < 2; op_idx++) { - int op_name = mm_root->operands[op_idx]->result_name; - auto op_it = partition->arrays.find(op_name); - if (op_it == partition->arrays.end() || op_it->second->local_size() <= 0) - continue; - - if (mm_root->operand_regions.size() < 2) - continue; - auto* r3d = static_cast*>(mm_root->operand_regions[op_idx]); - - int dd = -1; - for (int d = 0; d < 3; d++) { - if (r3d->stop[d] - r3d->start[d] == 1) { dd = d; break; } - } - if (dd < 0) - continue; - int rd = (dd == 0) ? 1 : 0; - int cd = (dd <= 1) ? 2 : 1; - - int dd_s = r3d->start[dd], dd_e = r3d->stop[dd]; - int rd_s = r3d->start[rd], rd_e = r3d->stop[rd]; - int cd_s = r3d->start[cd], cd_e = r3d->stop[cd]; - - // Translate AST regions to global space - auto op_decomp = op_it->second->decomp; - dd_s += op_decomp.offset[dd]; - dd_e += op_decomp.offset[dd]; - rd_s += op_decomp.offset[rd]; - rd_e += op_decomp.offset[rd]; - cd_s += op_decomp.offset[cd]; - cd_e += op_decomp.offset[cd]; - - int local_dims[3]; - for (int d = 0; d < 3; d++) - local_dims[d] = op_it->second->region.size(d); - - auto op_chare = op_decomp.chare_region_global(nd_idx); - int dd_lo = std::max(op_chare.start[dd], dd_s); - int dd_hi = std::min(op_chare.stop[dd], dd_e); - int row_lo = std::max(op_chare.start[rd], rd_s); - int row_hi = std::min(op_chare.stop[rd], rd_e); - int col_lo = std::max(op_chare.start[cd], cd_s); - int col_hi = std::min(op_chare.stop[cd], cd_e); - - if (dd_lo >= dd_hi || row_lo >= row_hi || col_lo >= col_hi) - continue; - - int sub_rows = row_hi - row_lo; - int sub_cols = col_hi - col_lo; - int dd_local_offset = dd_lo - op_chare.start[dd]; - int row_offset = row_lo - op_chare.start[rd]; - int col_offset = col_lo - op_chare.start[cd]; - -#ifndef USE_KOKKOS - op_it->second->copyToHost(); -#endif - - auto c3_meta_it = dag_group->array_meta.find(result_name); - ArrayDecomp<2> c3_decomp; - if (c3_meta_it != dag_group->array_meta.end()) - c3_decomp = c3_meta_it->second.template decomp<2>(); - int c_M = c3_decomp.offset[0] + c3_decomp.global_shape[0]; - int c_N_cols_g = c3_decomp.offset[1] + c3_decomp.global_shape[1]; - - int first_c_row_chare = c3_decomp.offset[0] / tile_2d; - int last_c_row_chare = (c_M > 0) ? (c_M - 1) / tile_2d : 0; - int first_c_col_chare = c3_decomp.offset[1] / tile_2d; - int last_c_col_chare = (c_N_cols_g > 0) ? (c_N_cols_g - 1) / tile_2d : 0; - - if (op_idx == 0) { - // This 3D chare holds A. Send A panels to C chares. - int k_start = col_lo - cd_s; - int k_end = col_hi - cd_s; - int c_row_start = c3_decomp.offset[0] + (row_lo - rd_s); - int c_row_end = c3_decomp.offset[0] + (row_hi - rd_s); - int first_dest_row_chare = c_row_start / tile_2d; - int last_dest_row_chare = (c_row_end - 1) / tile_2d; - - auto send_3d_panel = [&](auto* dummy) { - using T = std::remove_pointer_t; - int s0 = local_dims[1] * local_dims[2]; - int s1 = local_dims[2]; - int s2 = 1; - int strides[3] = {s0, s1, s2}; - -#ifdef USE_KOKKOS - T* d_data = static_cast(op_it->second->device_data_ptr()); - int cap_dd = dd_local_offset, cap_rd = row_offset, cap_cd = col_offset; - int cap_dd_dim = dd, cap_rd_dim = rd, cap_cd_dim = cd; - int cap_sub_cols = sub_cols; - - for (int tj = first_c_col_chare; tj <= last_c_col_chare; tj++) { - for (int ti = first_dest_row_chare; ti <= last_dest_row_chare; ti++) { - int dest_row_start_c = std::max(c_row_start, ti * tile_2d); - int dest_row_end_c = std::min(c_row_end, (ti + 1) * tile_2d); - int send_rows = dest_row_end_c - dest_row_start_c; - int send_row_off = dest_row_start_c - c_row_start; - - ChareIndex<2> target_ci; - target_ci.idx[0] = ti; - target_ci.idx[1] = tj; - - int64_t panel_count = (int64_t)send_rows * sub_cols; - int64_t panel_bytes = panel_count * sizeof(T); - T* panel_buf = static_cast( - Kokkos::kokkos_malloc(panel_bytes)); - Kokkos::parallel_for( - CT_COMM_POLICY(partition, panel_count), - KOKKOS_LAMBDA(int i) { - int r = i / cap_sub_cols; - int c = i % cap_sub_cols; - int idx3d[3]; - idx3d[cap_dd_dim] = cap_dd; - idx3d[cap_rd_dim] = cap_rd + send_row_off + r; - idx3d[cap_cd_dim] = cap_cd + c; - panel_buf[i] = d_data[idx3d[0] * strides[0] + - idx3d[1] * strides[1] + - idx3d[2] * strides[2]]; - }); - - int region_data[2 * 3] = {k_start, k_end, 1, - dest_row_start_c, dest_row_end_c, 1}; - -#ifndef NDEBUG - partition->comm_bytes_sent += panel_bytes; -#endif - device_pack_send<3, 2>(partition, dag_group->partition_proxy_2, - target_ci, node->id, 0, op_name, - region_data, panel_bytes, panel_buf); - } - } -#else - T* data = static_cast(op_it->second->data_ptr()); - int64_t send_count = (int64_t)sub_rows * sub_cols; - T* send_buf = new T[send_count]; - - for (int r = 0; r < sub_rows; r++) { - for (int c = 0; c < sub_cols; c++) { - int idx3d[3]; - idx3d[dd] = dd_local_offset; - idx3d[rd] = row_offset + r; - idx3d[cd] = col_offset + c; - send_buf[r * sub_cols + c] = - data[idx3d[0] * strides[0] + idx3d[1] * strides[1] + idx3d[2] * strides[2]]; - } - } - - for (int tj = first_c_col_chare; tj <= last_c_col_chare; tj++) { - for (int ti = first_dest_row_chare; ti <= last_dest_row_chare; ti++) { - int dest_row_start_c = std::max(c_row_start, ti * tile_2d); - int dest_row_end_c = std::min(c_row_end, (ti + 1) * tile_2d); - int send_rows = dest_row_end_c - dest_row_start_c; - int send_row_off = dest_row_start_c - c_row_start; - - ChareIndex<2> target_ci; - target_ci.idx[0] = ti; - target_ci.idx[1] = tj; - - int region_data[2 * 3]; - region_data[0] = k_start; - region_data[1] = k_end; - region_data[2] = 1; - region_data[3] = dest_row_start_c; - region_data[4] = dest_row_end_c; - region_data[5] = 1; - - int64_t panel_count = (int64_t)send_rows * sub_cols; - int64_t panel_bytes = panel_count * sizeof(T); - T* panel_buf = new T[panel_count]; - for (int r = 0; r < send_rows; r++) - memcpy(panel_buf + r * sub_cols, - send_buf + (send_row_off + r) * sub_cols, - sub_cols * sizeof(T)); - -#ifndef NDEBUG - partition->comm_bytes_sent += panel_bytes; -#endif - proxy_at<2>(dag_group->partition_proxy_2, target_ci) - .receive_data(node->id, /*input_index=*/0, op_name, 2, - region_data, panel_bytes, - reinterpret_cast(panel_buf)); - delete[] panel_buf; - } - } - delete[] send_buf; -#endif - }; - switch (dt) { - case DType::FLOAT32: { float* d = nullptr; send_3d_panel(d); break; } - case DType::FLOAT64: { double* d = nullptr; send_3d_panel(d); break; } - case DType::INT32: { int32_t* d = nullptr; send_3d_panel(d); break; } - case DType::INT64: { int64_t* d = nullptr; send_3d_panel(d); break; } - } - } else { - // This 3D chare holds B. Send B panels to C chares. - int k_start = row_lo - rd_s; - int k_end = row_hi - rd_s; - int c_col_start = c3_decomp.offset[1] + (col_lo - cd_s); - int c_col_end = c3_decomp.offset[1] + (col_hi - cd_s); - int first_dest_col_chare = c_col_start / tile_2d; - int last_dest_col_chare = (c_col_end - 1) / tile_2d; - - auto send_3d_panel = [&](auto* dummy) { - using T = std::remove_pointer_t; - int s0 = local_dims[1] * local_dims[2]; - int s1 = local_dims[2]; - int s2 = 1; - int strides[3] = {s0, s1, s2}; - -#ifdef USE_KOKKOS - T* d_data = static_cast(op_it->second->device_data_ptr()); - int cap_dd = dd_local_offset, cap_rd = row_offset, cap_cd = col_offset; - int cap_dd_dim = dd, cap_rd_dim = rd, cap_cd_dim = cd; - - for (int ti = first_c_row_chare; ti <= last_c_row_chare; ti++) { - for (int tj = first_dest_col_chare; tj <= last_dest_col_chare; tj++) { - int dest_col_start_c = std::max(c_col_start, tj * tile_2d); - int dest_col_end_c = std::min(c_col_end, (tj + 1) * tile_2d); - int send_cols = dest_col_end_c - dest_col_start_c; - int send_col_off = dest_col_start_c - c_col_start; - - ChareIndex<2> target_ci; - target_ci.idx[0] = ti; - target_ci.idx[1] = tj; - - int64_t panel_count = (int64_t)sub_rows * send_cols; - int64_t panel_bytes = panel_count * sizeof(T); - T* panel_buf = static_cast( - Kokkos::kokkos_malloc(panel_bytes)); - int cap_send_cols = send_cols; - Kokkos::parallel_for( - CT_COMM_POLICY(partition, panel_count), - KOKKOS_LAMBDA(int i) { - int r = i / cap_send_cols; - int c = i % cap_send_cols; - int idx3d[3]; - idx3d[cap_dd_dim] = cap_dd; - idx3d[cap_rd_dim] = cap_rd + r; - idx3d[cap_cd_dim] = cap_cd + send_col_off + c; - panel_buf[i] = d_data[idx3d[0] * strides[0] + - idx3d[1] * strides[1] + - idx3d[2] * strides[2]]; - }); - - int region_data[2 * 3] = {k_start, k_end, 1, - dest_col_start_c, dest_col_end_c, 1}; -#ifndef NDEBUG - partition->comm_bytes_sent += panel_bytes; -#endif - device_pack_send<3, 2>(partition, dag_group->partition_proxy_2, - target_ci, node->id, 1, op_name, - region_data, panel_bytes, panel_buf); - } - } -#else - T* data = static_cast(op_it->second->data_ptr()); - int64_t send_count = (int64_t)sub_rows * sub_cols; - T* send_buf = new T[send_count]; - - for (int r = 0; r < sub_rows; r++) { - for (int c = 0; c < sub_cols; c++) { - int idx3d[3]; - idx3d[dd] = dd_local_offset; - idx3d[rd] = row_offset + r; - idx3d[cd] = col_offset + c; - send_buf[r * sub_cols + c] = - data[idx3d[0] * strides[0] + idx3d[1] * strides[1] + idx3d[2] * strides[2]]; - } - } - - for (int ti = first_c_row_chare; ti <= last_c_row_chare; ti++) { - for (int tj = first_dest_col_chare; tj <= last_dest_col_chare; tj++) { - int dest_col_start_c = std::max(c_col_start, tj * tile_2d); - int dest_col_end_c = std::min(c_col_end, (tj + 1) * tile_2d); - int send_cols = dest_col_end_c - dest_col_start_c; - int send_col_off = dest_col_start_c - c_col_start; - - ChareIndex<2> target_ci; - target_ci.idx[0] = ti; - target_ci.idx[1] = tj; - - int region_data[2 * 3]; - region_data[0] = k_start; - region_data[1] = k_end; - region_data[2] = 1; - region_data[3] = dest_col_start_c; - region_data[4] = dest_col_end_c; - region_data[5] = 1; - - int64_t panel_count = (int64_t)sub_rows * send_cols; - int64_t panel_bytes = panel_count * sizeof(T); - T* panel_buf = new T[panel_count]; - for (int r = 0; r < sub_rows; r++) - for (int c = 0; c < send_cols; c++) - panel_buf[r * send_cols + c] = - send_buf[r * sub_cols + (send_col_off + c)]; - -#ifndef NDEBUG - partition->comm_bytes_sent += panel_bytes; -#endif - proxy_at<2>(dag_group->partition_proxy_2, target_ci) - .receive_data(node->id, /*input_index=*/1, op_name, 2, - region_data, panel_bytes, - reinterpret_cast(panel_buf)); - delete[] panel_buf; - } - } - delete[] send_buf; -#endif - }; - switch (dt) { - case DType::FLOAT32: { float* d = nullptr; send_3d_panel(d); break; } - case DType::FLOAT64: { double* d = nullptr; send_3d_panel(d); break; } - case DType::INT32: { int32_t* d = nullptr; send_3d_panel(d); break; } - case DType::INT64: { int64_t* d = nullptr; send_3d_panel(d); break; } - } - } - } - - // 3D chares only send — they don't hold C, so just finish - node_finished(node->id); - - } else { - CkAbort("MATMATMUL only supported for 2D and 3D partitions (N=%d)", N); - } -} - -template void ArrayDAGExecutorND<1>::execute_matmatmul_node(DAGNode*); -template void ArrayDAGExecutorND<2>::execute_matmatmul_node(DAGNode*); -template void ArrayDAGExecutorND<3>::execute_matmatmul_node(DAGNode*); diff --git a/src/executor_matmul.cpp b/src/executor_matmul.cpp deleted file mode 100644 index 63e73b4..0000000 --- a/src/executor_matmul.cpp +++ /dev/null @@ -1,593 +0,0 @@ -#include "backend_internal.hpp" -#include "dispatch.hpp" - -#include - -template -static void cross_matmul_send_vector_slice_1d_to_2d( - PartitionImpl<1>* partition, DAGNode* node, int vec_name, - const ChareIndex<2>& target_ci, int local_offset, int send_len, int col_start, int col_end, - CProxy_Partition2D& proxy_2d) { - auto vec_it = partition->arrays.find(vec_name); - if (vec_it == partition->arrays.end()) - return; - Array<1, T>* vec = static_cast*>(vec_it->second); - int64_t byte_size = send_len * sizeof(T); - - int region_data[2 * 3]; - region_data[0] = col_start; - region_data[1] = col_end; - region_data[2] = 1; - region_data[3] = 0; - region_data[4] = 1; - region_data[5] = 1; - -#ifndef NDEBUG - partition->comm_bytes_sent += byte_size; -#endif - -#ifdef USE_KOKKOS - T* send_buf = static_cast(Kokkos::kokkos_malloc(byte_size)); - T* src_device = static_cast(vec->device_data_ptr()) + local_offset; - Kokkos::parallel_for( - CT_COMM_POLICY(partition, send_len), - KOKKOS_LAMBDA(int i) { send_buf[i] = src_device[i]; }); - device_pack_send<1, 2>(partition, proxy_2d, target_ci, node->id, 0, vec_name, region_data, - byte_size, send_buf); -#else - vec->copyToHost(); - T* send_buf = new T[send_len]; - memcpy(send_buf, static_cast(vec->data_ptr()) + local_offset, byte_size); - proxy_at<2>(proxy_2d, target_ci) - .receive_data(node->id, /*input_index=*/0, vec_name, 2, region_data, byte_size, - reinterpret_cast(send_buf)); - delete[] send_buf; -#endif -} - -template -static void cross_matmul_send_vector_slice_1d_to_3d( - PartitionImpl<1>* partition, DAGNode* node, int vec_name, - const ChareIndex<3>& target_ci, int local_offset, int send_len, int col_start, int col_end, - CProxy_Partition3D& proxy_3d) { - auto vec_it = partition->arrays.find(vec_name); - if (vec_it == partition->arrays.end()) - return; - Array<1, T>* vec = static_cast*>(vec_it->second); - int64_t byte_size = send_len * sizeof(T); - - int region_data[3 * 3]; - region_data[0] = col_start; - region_data[1] = col_end; - region_data[2] = 1; - region_data[3] = 0; - region_data[4] = 1; - region_data[5] = 1; - region_data[6] = 0; - region_data[7] = 1; - region_data[8] = 1; - -#ifndef NDEBUG - partition->comm_bytes_sent += byte_size; -#endif - -#ifdef USE_KOKKOS - T* send_buf = static_cast(Kokkos::kokkos_malloc(byte_size)); - T* src_device = static_cast(vec->device_data_ptr()) + local_offset; - Kokkos::parallel_for( - CT_COMM_POLICY(partition, send_len), - KOKKOS_LAMBDA(int i) { send_buf[i] = src_device[i]; }); - device_pack_send<1, 3>(partition, proxy_3d, target_ci, node->id, 0, vec_name, region_data, - byte_size, send_buf); -#else - vec->copyToHost(); - T* send_buf = new T[send_len]; - memcpy(send_buf, static_cast(vec->data_ptr()) + local_offset, byte_size); - proxy_at<3>(proxy_3d, target_ci) - .receive_data(node->id, /*input_index=*/0, vec_name, 3, region_data, byte_size, - reinterpret_cast(send_buf)); - delete[] send_buf; -#endif -} - -static void dispatch_cross_matmul_send_vector_slice_3d( - DType dt, PartitionImpl<1>* partition, DAGNode* node, int vec_name, - const ChareIndex<3>& target_ci, int local_offset, int send_len, int col_start, int col_end, - CProxy_Partition3D& proxy_3d) { - switch (dt) { - case DType::FLOAT32: - cross_matmul_send_vector_slice_1d_to_3d(partition, node, vec_name, target_ci, - local_offset, send_len, col_start, col_end, - proxy_3d); - break; - case DType::FLOAT64: - cross_matmul_send_vector_slice_1d_to_3d(partition, node, vec_name, target_ci, - local_offset, send_len, col_start, - col_end, proxy_3d); - break; - case DType::INT32: - cross_matmul_send_vector_slice_1d_to_3d(partition, node, vec_name, target_ci, - local_offset, send_len, col_start, - col_end, proxy_3d); - break; - case DType::INT64: - cross_matmul_send_vector_slice_1d_to_3d(partition, node, vec_name, target_ci, - local_offset, send_len, col_start, - col_end, proxy_3d); - break; - } -} - -static void dispatch_cross_matmul_send_vector_slice( - DType dt, PartitionImpl<1>* partition, DAGNode* node, int vec_name, - const ChareIndex<2>& target_ci, int local_offset, int send_len, int col_start, int col_end, - CProxy_Partition2D& proxy_2d) { - switch (dt) { - case DType::FLOAT32: - cross_matmul_send_vector_slice_1d_to_2d(partition, node, vec_name, target_ci, - local_offset, send_len, col_start, col_end, - proxy_2d); - break; - case DType::FLOAT64: - cross_matmul_send_vector_slice_1d_to_2d(partition, node, vec_name, target_ci, - local_offset, send_len, col_start, - col_end, proxy_2d); - break; - case DType::INT32: - cross_matmul_send_vector_slice_1d_to_2d(partition, node, vec_name, target_ci, - local_offset, send_len, col_start, - col_end, proxy_2d); - break; - case DType::INT64: - cross_matmul_send_vector_slice_1d_to_2d(partition, node, vec_name, target_ci, - local_offset, send_len, col_start, - col_end, proxy_2d); - break; - } -} - -template -void ArrayDAGExecutorND::execute_matmul_node(DAGNode* node) { - // Parse AST to find operand names and detect cross-partition case - ASTNode* matmul_root = nullptr; - for (ASTNode* root : node->ast->roots) { - if (static_cast(root->opcode) == Opcode::MATMUL) { - matmul_root = root; - break; - } - } - int mat_name = matmul_root->operands[0]->result_name; - int vec_name = matmul_root->operands[1]->result_name; - int result_name = matmul_root->result_name; - - auto* dag_group = static_cast(group); - - if constexpr (N == 1) { - // --- 1D partition role in cross-partition matmul --- - int tile_1d = array_tile(dag_group->array_meta, vec_name, 1); - - auto nd_idx = partition->nd_index(); - int k = nd_idx[0]; - - DType dt = determine_dtype<1>(node, partition->arrays); - - // Look up the matrix global shape from array_meta - auto meta_it = dag_group->array_meta.find(mat_name); - if (meta_it == dag_group->array_meta.end()) { - CkAbort("Cross-partition MATMUL: could not find matrix metadata for name=%d", mat_name); - return; - } - int mat_ndims = meta_it->second.ndims; - - auto vec_it = partition->arrays.find(vec_name); - bool has_vec = (vec_it != partition->arrays.end() && vec_it->second->local_size() > 0); - int local_size = has_vec ? vec_it->second->local_size() : 0; - - // Compute vec global size from array_meta - int vec_s = 0, vec_e = 0; - auto vec_meta_it = dag_group->array_meta.find(vec_name); - if (vec_meta_it != dag_group->array_meta.end()) - vec_e = vec_meta_it->second.global_shape[0]; - - // Get decomps for global-space translation - ArrayDecomp<1> vec_decomp; - if (vec_meta_it != dag_group->array_meta.end()) - vec_decomp = vec_meta_it->second.template decomp<1>(); - auto result_meta_1d = dag_group->array_meta.find(result_name); - ArrayDecomp<1> result_decomp; - if (result_meta_1d != dag_group->array_meta.end()) - result_decomp = result_meta_1d->second.template decomp<1>(); - - if (mat_ndims == 3) { - // --- 3D matrix path (dimension-dropped matvec) --- - int tile_3d = array_tile(dag_group->array_meta, mat_name, 3); - - // operand_regions[0] is a 3D region, operand_regions[1] is 1D vec region - auto* mat_region_3d = static_cast*>(matmul_root->operand_regions[0]); - auto* vec_region = static_cast*>(matmul_root->operand_regions[1]); - vec_s = vec_region->start[0]; - vec_e = vec_region->stop[0]; - - // Identify dropped dim (size-1 range), row dim, col dim - int dd = -1; - for (int d = 0; d < 3; d++) { - if (mat_region_3d->stop[d] - mat_region_3d->start[d] == 1) { - dd = d; - break; - } - } - int rd = (dd == 0) ? 1 : 0; - int cd = (dd <= 1) ? 2 : 1; - - int dd_s = mat_region_3d->start[dd], dd_e = mat_region_3d->stop[dd]; - int rd_s = mat_region_3d->start[rd], rd_e = mat_region_3d->stop[rd]; - int cd_s = mat_region_3d->start[cd], cd_e = mat_region_3d->stop[cd]; - - // Translate AST regions to global space - auto mat_decomp_3d = meta_it->second.template decomp<3>(); - dd_s += mat_decomp_3d.offset[dd]; - dd_e += mat_decomp_3d.offset[dd]; - rd_s += mat_decomp_3d.offset[rd]; - rd_e += mat_decomp_3d.offset[rd]; - cd_s += mat_decomp_3d.offset[cd]; - cd_e += mat_decomp_3d.offset[cd]; - vec_s += vec_decomp.offset[0]; - vec_e += vec_decomp.offset[0]; - - int slice_rows = rd_e - rd_s; - - // Determine which 3D chares overlap the region - int first_dd_chare = dd_s / tile_3d; - int last_dd_chare = (dd_e > 0) ? (dd_e - 1) / tile_3d : 0; - int first_row_chare = rd_s / tile_3d; - int last_row_chare = (rd_e > 0) ? (rd_e - 1) / tile_3d : 0; - int first_col_chare = cd_s / tile_3d; - int last_col_chare = (cd_e > 0) ? (cd_e - 1) / tile_3d : 0; - - DBG_PRINT("[1D Chare %d] cross_matmul_3d: node=%d dd=%d rd=%d cd=%d " - "dd=[%d:%d] rd=[%d:%d] cd=[%d:%d] vec=[%d:%d]\n", - partition->index[0], node->id, dd, rd, cd, - dd_s, dd_e, rd_s, rd_e, cd_s, cd_e, vec_s, vec_e); - - // Phase 1: Send vector data to relevant 3D chares - if (has_vec) { - int my_vec_start = vec_decomp.chare_start_global(0, k); - int my_vec_end = my_vec_start + local_size; - int rel_start = std::max(my_vec_start, vec_s); - int rel_end = std::min(my_vec_end, vec_e); - if (rel_start < rel_end) { - int mat_col_lo = cd_s + (rel_start - vec_s); - int mat_col_hi = cd_s + (rel_end - vec_s); - - int fc = mat_col_lo / tile_3d; - int lc = (mat_col_hi - 1) / tile_3d; - - for (int tc = fc; tc <= lc; tc++) { - int chare_col_start = std::max(mat_col_lo, tc * tile_3d); - int chare_col_end = std::min(mat_col_hi, (tc + 1) * tile_3d); - int send_len = chare_col_end - chare_col_start; - int local_offset = (chare_col_start - cd_s + vec_s) - my_vec_start; - - for (int tr = first_row_chare; tr <= last_row_chare; tr++) { - for (int tb = first_dd_chare; tb <= last_dd_chare; tb++) { - ChareIndex<3> target_ci; - target_ci.idx[dd] = tb; - target_ci.idx[rd] = tr; - target_ci.idx[cd] = tc; - dispatch_cross_matmul_send_vector_slice_3d( - dt, partition, node, vec_name, target_ci, - local_offset, send_len, chare_col_start, chare_col_end, - dag_group->partition_proxy_3); - } - } - } - } - } - - // Phase 2: Wait for partial results from 3D chares - auto r_local_3d = result_decomp.chare_region_local(nd_idx); - int my_result_start = r_local_3d.start[0]; - int my_result_end = std::min(r_local_3d.stop[0], slice_rows); - int total_expected = 0; - if (my_result_start < slice_rows) { - for (int tb = first_dd_chare; tb <= last_dd_chare; tb++) { - for (int tc = first_col_chare; tc <= last_col_chare; tc++) { - for (int tr = first_row_chare; tr <= last_row_chare; tr++) { - int global_row_lo = std::max(tr * tile_3d, rd_s); - int global_row_hi = std::min((tr + 1) * tile_3d, rd_e); - int result_lo = global_row_lo - rd_s; - int result_hi = global_row_hi - rd_s; - if (result_lo < my_result_end && result_hi > my_result_start) - total_expected++; - } - } - } - } - - auto pre_it = pending.find(node->id); - int pre_arrived = 0; - std::unordered_map>> pre_buffers; - if (pre_it != pending.end()) { - pre_arrived = -(pre_it->second.expected_msgs); - pre_buffers = std::move(pre_it->second.remote_buffers); - } - int remaining = total_expected - pre_arrived; - pending[node->id] = {node, remaining, {}, std::move(pre_buffers)}; - DBG_PRINT("[1D Chare %d] cross_matmul_3d: total_expected=%d remaining=%d\n", - partition->index[0], total_expected, remaining); - if (remaining <= 0) - on_comm_done(node->id); - - } else { - // --- 2D matrix path (existing) --- - int tile_2d = array_tile(dag_group->array_meta, mat_name, 2); - - int mat_global_rows = meta_it->second.global_shape[0]; - int mat_global_cols = meta_it->second.global_shape[1]; - - int mat_rs = 0, mat_re = mat_global_rows; - int mat_cs = 0, mat_ce = mat_global_cols; - if (matmul_root->operand_regions.size() >= 2) { - auto* mat_region = static_cast*>(matmul_root->operand_regions[0]); - auto* vec_region = static_cast*>(matmul_root->operand_regions[1]); - mat_rs = mat_region->start[0]; - mat_re = mat_region->stop[0]; - mat_cs = mat_region->start[1]; - mat_ce = mat_region->stop[1]; - vec_s = vec_region->start[0]; - vec_e = vec_region->stop[0]; - } - - // Translate AST regions to global space - auto mat_decomp_2d = meta_it->second.template decomp<2>(); - mat_rs += mat_decomp_2d.offset[0]; - mat_re += mat_decomp_2d.offset[0]; - mat_cs += mat_decomp_2d.offset[1]; - mat_ce += mat_decomp_2d.offset[1]; - vec_s += vec_decomp.offset[0]; - vec_e += vec_decomp.offset[0]; - - int slice_rows = mat_re - mat_rs; - - int first_row_chare = mat_rs / tile_2d; - int last_row_chare = (mat_re > 0) ? (mat_re - 1) / tile_2d : 0; - int first_col_chare = mat_cs / tile_2d; - int last_col_chare = (mat_ce > 0) ? (mat_ce - 1) / tile_2d : 0; - - DBG_PRINT("[1D Chare %d] cross_matmul: node=%d vec=%d mat=%d result=%d " - "slice=[%d:%d,%d:%d] vec=[%d:%d] has_vec=%d\n", - partition->index[0], node->id, vec_name, mat_name, result_name, - mat_rs, mat_re, mat_cs, mat_ce, vec_s, vec_e, (int)has_vec); - - if (has_vec) { - int my_vec_start = vec_decomp.chare_start_global(0, k); - int my_vec_end = my_vec_start + local_size; - int rel_start = std::max(my_vec_start, vec_s); - int rel_end = std::min(my_vec_end, vec_e); - if (rel_start < rel_end) { - int mat_col_lo = mat_cs + (rel_start - vec_s); - int mat_col_hi = mat_cs + (rel_end - vec_s); - - int fc = mat_col_lo / tile_2d; - int lc = (mat_col_hi - 1) / tile_2d; - for (int tc = fc; tc <= lc; tc++) { - int chare_col_start = std::max(mat_col_lo, tc * tile_2d); - int chare_col_end = std::min(mat_col_hi, (tc + 1) * tile_2d); - int send_len = chare_col_end - chare_col_start; - int local_offset = (chare_col_start - mat_cs + vec_s) - my_vec_start; - - for (int tr = first_row_chare; tr <= last_row_chare; tr++) { - ChareIndex<2> target_ci; - target_ci.idx[0] = tr; - target_ci.idx[1] = tc; - dispatch_cross_matmul_send_vector_slice( - dt, partition, node, vec_name, target_ci, - local_offset, send_len, chare_col_start, chare_col_end, - dag_group->partition_proxy_2); - } - } - } - } - - auto r_local_2d = result_decomp.chare_region_local(nd_idx); - int my_result_start = r_local_2d.start[0]; - int my_result_end = std::min(r_local_2d.stop[0], slice_rows); - int total_expected = 0; - if (my_result_start < slice_rows) { - for (int tc = first_col_chare; tc <= last_col_chare; tc++) { - for (int tr = first_row_chare; tr <= last_row_chare; tr++) { - int global_row_lo = std::max(tr * tile_2d, mat_rs); - int global_row_hi = std::min((tr + 1) * tile_2d, mat_re); - int result_lo = global_row_lo - mat_rs; - int result_hi = global_row_hi - mat_rs; - if (result_lo < my_result_end && result_hi > my_result_start) - total_expected++; - } - } - } - - auto pre_it = pending.find(node->id); - int pre_arrived = 0; - std::unordered_map>> pre_buffers; - if (pre_it != pending.end()) { - pre_arrived = -(pre_it->second.expected_msgs); - pre_buffers = std::move(pre_it->second.remote_buffers); - } - int remaining = total_expected - pre_arrived; - pending[node->id] = {node, remaining, {}, std::move(pre_buffers)}; - DBG_PRINT("[1D Chare %d] cross_matmul: node=%d total_expected=%d pre_arrived=%d remaining=%d\n", - partition->index[0], node->id, total_expected, pre_arrived, remaining); - if (remaining <= 0) - on_comm_done(node->id); - } - - } else if constexpr (N == 2) { - int tile_1d = array_tile(dag_group->array_meta, vec_name, 1); - - auto nd_idx = partition->nd_index(); - - DType dt = determine_dtype(node, partition->arrays); - - auto mat_it = partition->arrays.find(mat_name); - if (mat_it == partition->arrays.end()) { - node_finished(node->id); - return; - } - int tile_2d = mat_it->second->decomp.tile; - int mat_global_rows = mat_it->second->global_shape[0]; - int mat_global_cols = mat_it->second->global_shape[1]; - - // Extract slice regions if present - int mat_rs = 0, mat_re = mat_global_rows; - int mat_cs = 0, mat_ce = mat_global_cols; - int vec_s = 0, vec_e = mat_ce - mat_cs; // default: full range - if (matmul_root->operand_regions.size() >= 2) { - auto* mat_region = static_cast*>(matmul_root->operand_regions[0]); - auto* vec_region = static_cast*>(matmul_root->operand_regions[1]); - mat_rs = mat_region->start[0]; - mat_re = mat_region->stop[0]; - mat_cs = mat_region->start[1]; - mat_ce = mat_region->stop[1]; - vec_s = vec_region->start[0]; - vec_e = vec_region->stop[0]; - } - - // Translate AST regions to global space - auto mat_decomp = mat_it->second->decomp; - mat_rs += mat_decomp.offset[0]; - mat_re += mat_decomp.offset[0]; - mat_cs += mat_decomp.offset[1]; - mat_ce += mat_decomp.offset[1]; - auto vec_decomp_1d = dag_group->array_meta[vec_name].template decomp<1>(); - vec_s += vec_decomp_1d.offset[0]; - vec_e += vec_decomp_1d.offset[0]; - - // Check if this 2D chare's tile overlaps the matrix slice - auto mat_chare = mat_decomp.chare_region_global(nd_idx); - int row_lo = std::max(mat_chare.start[0], mat_rs); - int row_hi = std::min(mat_chare.stop[0], mat_re); - int col_lo = std::max(mat_chare.start[1], mat_cs); - int col_hi = std::min(mat_chare.stop[1], mat_ce); - if (row_lo >= row_hi || col_lo >= col_hi) { - node_finished(node->id); - return; - } - - // Count how many vector messages to expect. - int vec_idx_start = vec_s + (col_lo - mat_cs); - int vec_idx_end = vec_s + (col_hi - mat_cs); - int first_vec_chare = vec_idx_start / tile_1d; - int last_vec_chare = (vec_idx_end - 1) / tile_1d; - int vec_msgs = last_vec_chare - first_vec_chare + 1; - - auto pre_it = pending.find(node->id); - int pre_arrived = 0; - std::unordered_map>> pre_buffers; - if (pre_it != pending.end()) { - pre_arrived = -(pre_it->second.expected_msgs); - pre_buffers = std::move(pre_it->second.remote_buffers); - } - int remaining = vec_msgs - pre_arrived; - pending[node->id] = {node, remaining, {}, std::move(pre_buffers)}; - if (pre_arrived > 0 || remaining <= 0) { - DBG_PRINT("[2D Chare %d] matmul node=%d: pre_arrived=%d remaining=%d vec_msgs=%d\n", - partition->index[0], node->id, pre_arrived, remaining, vec_msgs); - } - - if (remaining <= 0) - on_comm_done(node->id); - - } else if constexpr (N == 3) { - // --- 3D partition role in cross-partition matmul (dimension-dropped) --- - int tile_1d = array_tile(dag_group->array_meta, vec_name, 1); - - auto nd_idx = partition->nd_index(); - - DType dt = determine_dtype(node, partition->arrays); - - auto mat_it = partition->arrays.find(mat_name); - if (mat_it == partition->arrays.end()) { - node_finished(node->id); - return; - } - int tile_3d = mat_it->second->decomp.tile; - - // operand_regions[0] is a 3D region - auto* mat_region_3d = static_cast*>(matmul_root->operand_regions[0]); - auto* vec_region = static_cast*>(matmul_root->operand_regions[1]); - - // Identify dropped dim (size-1 range), row dim, col dim - int dd = -1; - for (int d = 0; d < 3; d++) { - if (mat_region_3d->stop[d] - mat_region_3d->start[d] == 1) { - dd = d; - break; - } - } - int rd = (dd == 0) ? 1 : 0; - int cd = (dd <= 1) ? 2 : 1; - - int dd_s = mat_region_3d->start[dd], dd_e = mat_region_3d->stop[dd]; - int rd_s = mat_region_3d->start[rd], rd_e = mat_region_3d->stop[rd]; - int cd_s = mat_region_3d->start[cd], cd_e = mat_region_3d->stop[cd]; - int vec_s = vec_region->start[0], vec_e = vec_region->stop[0]; - - // Translate AST regions to global space - auto mat_decomp = mat_it->second->decomp; - dd_s += mat_decomp.offset[dd]; - dd_e += mat_decomp.offset[dd]; - rd_s += mat_decomp.offset[rd]; - rd_e += mat_decomp.offset[rd]; - cd_s += mat_decomp.offset[cd]; - cd_e += mat_decomp.offset[cd]; - auto vec_decomp_1d = dag_group->array_meta[vec_name].template decomp<1>(); - vec_s += vec_decomp_1d.offset[0]; - vec_e += vec_decomp_1d.offset[0]; - - // Check if this 3D chare overlaps the region in ALL dims - auto mat_chare = mat_decomp.chare_region_global(nd_idx); - int dd_lo = std::max(mat_chare.start[dd], dd_s); - int dd_hi = std::min(mat_chare.stop[dd], dd_e); - int row_lo = std::max(mat_chare.start[rd], rd_s); - int row_hi = std::min(mat_chare.stop[rd], rd_e); - int col_lo = std::max(mat_chare.start[cd], cd_s); - int col_hi = std::min(mat_chare.stop[cd], cd_e); - if (dd_lo >= dd_hi || row_lo >= row_hi || col_lo >= col_hi) { - node_finished(node->id); - return; - } - - // Count how many vector messages to expect - int vec_idx_start = vec_s + (col_lo - cd_s); - int vec_idx_end = vec_s + (col_hi - cd_s); - int first_vec_chare = vec_idx_start / tile_1d; - int last_vec_chare = (vec_idx_end - 1) / tile_1d; - int vec_msgs = last_vec_chare - first_vec_chare + 1; - - auto pre_it = pending.find(node->id); - int pre_arrived = 0; - std::unordered_map>> pre_buffers; - if (pre_it != pending.end()) { - pre_arrived = -(pre_it->second.expected_msgs); - pre_buffers = std::move(pre_it->second.remote_buffers); - } - int remaining = vec_msgs - pre_arrived; - pending[node->id] = {node, remaining, {}, std::move(pre_buffers)}; - DBG_PRINT("[3D Chare %d] matmul node=%d: dd=%d rd=%d cd=%d " - "dd_lo=%d dd_hi=%d row_lo=%d row_hi=%d col_lo=%d col_hi=%d " - "vec_msgs=%d remaining=%d\n", - partition->index[0], node->id, dd, rd, cd, - dd_lo, dd_hi, row_lo, row_hi, col_lo, col_hi, - vec_msgs, remaining); - - if (remaining <= 0) - on_comm_done(node->id); - - } else { - CkAbort("MATMUL only supported for 1D, 2D, and 3D partitions (N=%d)", N); - } -} - -template void ArrayDAGExecutorND<1>::execute_matmul_node(DAGNode*); -template void ArrayDAGExecutorND<2>::execute_matmul_node(DAGNode*); -template void ArrayDAGExecutorND<3>::execute_matmul_node(DAGNode*); diff --git a/src/executor_reduce.cpp b/src/executor_reduce.cpp deleted file mode 100644 index c68ca0a..0000000 --- a/src/executor_reduce.cpp +++ /dev/null @@ -1,123 +0,0 @@ -#include "backend_internal.hpp" - -#ifdef USE_KOKKOS -#include -#endif - -template -void ArrayDAGExecutorND::execute_reduce_node(DAGNode* node) { - if constexpr (N != 1) { - CkAbort("REDUCE only supported for 1D partitions (N=%d)", N); - return; - } - - if constexpr (N == 1) { - ASTNode* reduce_root = nullptr; - for (ASTNode* root : node->ast->roots) { - if (static_cast(root->opcode) == Opcode::REDUCE) { - reduce_root = root; - break; - } - } - int lhs_name = reduce_root->operands[0]->result_name; - int rhs_name = reduce_root->operands[1]->result_name; - int result_name = reduce_root->result_name; - - DType dt = determine_dtype<1>(node, partition->arrays); - - DBG_PRINT("[1D Chare %d] execute_reduce_node: lhs=%d rhs=%d result=%d\n", - partition->index[0], lhs_name, rhs_name, result_name); - - // Compute local dot product (0 if this chare has no data for either operand) - auto lhs_it = partition->arrays.find(lhs_name); - auto rhs_it = partition->arrays.find(rhs_name); - - // Build ReduceContrib with metadata + local partial value - ReduceContrib contrib = {}; - contrib.node_id = node->id; - contrib.result_name = result_name; - contrib.dtype_int = static_cast(dt); - - if (lhs_it != partition->arrays.end() && rhs_it != partition->arrays.end()) { - int local_size = - std::min(lhs_it->second->local_size(), rhs_it->second->local_size()); - -#ifdef USE_KOKKOS - auto kokkos_dot = [&](auto dummy) { - using VT = decltype(dummy); - Kokkos::View> - d_lhs(static_cast(lhs_it->second->device_data_ptr()), local_size); - Kokkos::View> - d_rhs(static_cast(rhs_it->second->device_data_ptr()), local_size); - VT val = KokkosBlas::dot(d_lhs, d_rhs); - memcpy(contrib.value, &val, sizeof(VT)); - }; - switch (dt) { - case DType::FLOAT32: kokkos_dot(float{}); break; - case DType::FLOAT64: kokkos_dot(double{}); break; - case DType::INT32: kokkos_dot(int32_t{}); break; - case DType::INT64: kokkos_dot(int64_t{}); break; - } -#else - lhs_it->second->copyToHost(); - rhs_it->second->copyToHost(); - - switch (dt) { - case DType::FLOAT32: { - float val = eigen_dot(static_cast(lhs_it->second->data_ptr()), - static_cast(rhs_it->second->data_ptr()), - local_size); - memcpy(contrib.value, &val, sizeof(float)); - break; - } - case DType::FLOAT64: { - double val = eigen_dot(static_cast(lhs_it->second->data_ptr()), - static_cast(rhs_it->second->data_ptr()), - local_size); - memcpy(contrib.value, &val, sizeof(double)); - break; - } - case DType::INT32: { - int32_t val = eigen_dot(static_cast(lhs_it->second->data_ptr()), - static_cast(rhs_it->second->data_ptr()), - local_size); - memcpy(contrib.value, &val, sizeof(int32_t)); - break; - } - case DType::INT64: { - int64_t val = eigen_dot(static_cast(lhs_it->second->data_ptr()), - static_cast(rhs_it->second->data_ptr()), - local_size); - memcpy(contrib.value, &val, sizeof(int64_t)); - break; - } - } -#endif - - DBG_PRINT("[1D Chare %d] local dot computed (local_size=%d)\n", - partition->index[0], local_size); - } - // else: contrib.value is all zeros — contributes 0 to the reduction - - // All chares contribute to the Charm++ reduction using the custom - // reducer that carries metadata (node_id, result_name, dtype). - ChareIndex<1> target_ci; - target_ci.idx[0] = 0; - - auto* dag_group = static_cast(group); - CkCallback cb(CkIndex_Partition1D::reduce_result(nullptr), - proxy_at<1>(dag_group->partition_proxy_1, target_ci)); - partition->contribute(sizeof(ReduceContrib), &contrib, reduce_dot_sum_type, cb); - - // Non-zero chares are done — their partial is in flight via the reduction. - auto nd_idx = partition->nd_index(); - if (nd_idx[0] != 0) { - node_finished(node->id); - } - // Chare 0 waits for reduce_result callback before calling node_finished. - } -} - -template void ArrayDAGExecutorND<1>::execute_reduce_node(DAGNode*); -template void ArrayDAGExecutorND<2>::execute_reduce_node(DAGNode*); -template void ArrayDAGExecutorND<3>::execute_reduce_node(DAGNode*); diff --git a/src/executor_reducer.cpp b/src/executor_reducer.cpp deleted file mode 100644 index 0f599b5..0000000 --- a/src/executor_reducer.cpp +++ /dev/null @@ -1,82 +0,0 @@ -#include "backend_internal.hpp" - -#include - -#ifdef USE_KOKKOS -#include -#endif - -CkReduction::reducerType reduce_dot_sum_type; - -CkReductionMsg* reduce_dot_sum(int nMsg, CkReductionMsg** msgs) { - // All messages have the same layout: ReduceContrib (24 bytes). - // Sum the value field based on dtype, preserve metadata from first msg. - ReduceContrib result; - memcpy(&result, msgs[0]->getData(), sizeof(ReduceContrib)); - - DType dt = static_cast(result.dtype_int); - for (int i = 1; i < nMsg; ++i) { - ReduceContrib other; - memcpy(&other, msgs[i]->getData(), sizeof(ReduceContrib)); - switch (dt) { - case DType::FLOAT32: { - float a, b; - memcpy(&a, result.value, sizeof(float)); - memcpy(&b, other.value, sizeof(float)); - a += b; - memcpy(result.value, &a, sizeof(float)); - break; - } - case DType::FLOAT64: { - double a, b; - memcpy(&a, result.value, sizeof(double)); - memcpy(&b, other.value, sizeof(double)); - a += b; - memcpy(result.value, &a, sizeof(double)); - break; - } - case DType::INT32: { - int32_t a, b; - memcpy(&a, result.value, sizeof(int32_t)); - memcpy(&b, other.value, sizeof(int32_t)); - a += b; - memcpy(result.value, &a, sizeof(int32_t)); - break; - } - case DType::INT64: { - int64_t a, b; - memcpy(&a, result.value, sizeof(int64_t)); - memcpy(&b, other.value, sizeof(int64_t)); - a += b; - memcpy(result.value, &a, sizeof(int64_t)); - break; - } - } - } - return CkReductionMsg::buildNew(sizeof(ReduceContrib), &result); -} - -void register_reduce_dot_sum() { - reduce_dot_sum_type = CkReduction::addReducer(reduce_dot_sum); -} - -template -DType determine_dtype(DAGNode* node, std::unordered_map*>& arrays) { - // Check AST roots for dtype info - for (ASTNode* root : node->ast->roots) - if (root->dtype != DType::FLOAT32) - return root->dtype; - // Fall back to first referenced array's dtype - for (ASTNode* root : node->ast->roots) - for (ASTNode* operand : root->operands) - if (!operand->is_scalar && !operand->is_broadcast) { - auto it = arrays.find(operand->result_name); - if (it != arrays.end()) - return it->second->dtype; - } - return DType::FLOAT32; -} - -template DType determine_dtype<1>(DAGNode*, std::unordered_map*>&); -template DType determine_dtype<2>(DAGNode*, std::unordered_map*>&); -template DType determine_dtype<3>(DAGNode*, std::unordered_map*>&); diff --git a/src/executor_transfer.cpp b/src/executor_transfer.cpp deleted file mode 100644 index 7171586..0000000 --- a/src/executor_transfer.cpp +++ /dev/null @@ -1,876 +0,0 @@ -#include "backend_internal.hpp" -#include "dispatch.hpp" - -#include -#include -#include -#include - -template -static typename PartitionTraits::ProxyType& get_partition_proxy(ArrayDAGGroup* dag_group); - -template <> -CProxy_Partition1D& get_partition_proxy<1>(ArrayDAGGroup* dag_group) { - return dag_group->partition_proxy_1; -} -template <> -CProxy_Partition2D& get_partition_proxy<2>(ArrayDAGGroup* dag_group) { - return dag_group->partition_proxy_2; -} -template <> -CProxy_Partition3D& get_partition_proxy<3>(ArrayDAGGroup* dag_group) { - return dag_group->partition_proxy_3; -} - -template -static void dispatch_cross_send(DType dt, PartitionImpl* partition, DAGNode* node, int source_name, - int target_name, ASTNode* sr_root, ArrayDAGGroup* dag_group) { - int target_ndims = sr_root->ndims; - Region* tgt_region = sr_root->region; - - // Compute dimension mapping from source metadata - auto src_it = partition->arrays.find(source_name); - if (src_it == partition->arrays.end() || src_it->second->local_size() == 0) - return; - - int src_global_shape[3] = {}; - for (int d = 0; d < N; ++d) - src_global_shape[d] = src_it->second->global_shape[d]; - - // Look up target array decomposition from array_meta - auto tgt_meta_it = dag_group->array_meta.find(target_name); - - // Dispatch to the correct N_tgt template - if (target_ndims == 1) { - int dim_map[1]; - compute_dim_map(src_global_shape, tgt_region, dim_map); - auto tgt_decomp = tgt_meta_it->second.template decomp<1>(); - dispatch_cross_set_region_send(dt, partition, node, source_name, dim_map, tgt_region, - get_partition_proxy<1>(dag_group), tgt_decomp); - } else if (target_ndims == 2) { - int dim_map[2]; - compute_dim_map(src_global_shape, tgt_region, dim_map); - auto tgt_decomp = tgt_meta_it->second.template decomp<2>(); - dispatch_cross_set_region_send(dt, partition, node, source_name, dim_map, tgt_region, - get_partition_proxy<2>(dag_group), tgt_decomp); - } else if (target_ndims == 3) { - int dim_map[3]; - compute_dim_map(src_global_shape, tgt_region, dim_map); - auto tgt_decomp = tgt_meta_it->second.template decomp<3>(); - dispatch_cross_set_region_send(dt, partition, node, source_name, dim_map, tgt_region, - get_partition_proxy<3>(dag_group), tgt_decomp); - } -} - -template -void ArrayDAGExecutorND::execute_cross_set_region_node(DAGNode* node) { - ASTNode* sr_root = nullptr; - for (ASTNode* root : node->ast->roots) { - if (static_cast(root->opcode) == Opcode::SET_REGION) { - sr_root = root; - break; - } - } - if (!sr_root) { - node_finished(node->id); - return; - } - - int target_name = sr_root->operands[0]->result_name; - int source_name = sr_root->operands[1]->result_name; - int target_ndims = sr_root->ndims; - int source_ndims = sr_root->operands[1]->ndims; - - auto* dag_group = static_cast(group); - DType dt = determine_dtype(node, partition->arrays); - - DBG_PRINT("[Chare %d] execute_cross_set_region_node<%d>: node_id=%d target=%d(nd=%d) " - "source=%d(nd=%d)\n", - partition->index[0], N, node->id, target_name, target_ndims, source_name, - source_ndims); - - if (N == source_ndims) { - // === SOURCE SIDE: send data to target partition === - dispatch_cross_send(dt, partition, node, source_name, target_name, sr_root, dag_group); - node_finished(node->id); - - } else if (N == target_ndims) { - // === TARGET SIDE: set up to receive data from source partition === - auto tgt_it = partition->arrays.find(target_name); - if (tgt_it == partition->arrays.end()) { - node_finished(node->id); - return; - } - - auto nd_idx = partition->nd_index(); - auto tgt_decomp = tgt_it->second->decomp; - auto r_chare_global = tgt_decomp.chare_region_global(nd_idx); - auto* tgt_region = static_cast*>(sr_root->region); - - // Preserve the target lattice for stepped target slices. - auto [my_input, has_overlap] = intersect(*tgt_region, r_chare_global); - if (!has_overlap) { - node_finished(node->id); - return; - } - - // Compute expected messages from source partition - auto meta_it = dag_group->array_meta.find(source_name); - if (meta_it == dag_group->array_meta.end()) { - CkAbort("Cross-partition SET_REGION: could not find source metadata for name=%d", - source_name); - return; - } - int src_ndims = meta_it->second.ndims; - int src_tile = array_tile(dag_group->array_meta, source_name, src_ndims); - - // Build dimension mapping (source dim → target dim) - int dim_map[3] = {}; // dim_map[td] = source dim - if (src_ndims > N) { - // Higher→lower: non-singleton source dims map to target dims - int td = 0; - for (int sd = 0; sd < src_ndims && td < N; ++sd) { - if (meta_it->second.global_shape[sd] > 1) - dim_map[td++] = sd; - } - } else { - // Lower→higher: non-singleton target region dims receive source dims - int sd = 0; - for (int td = 0; td < N && sd < src_ndims; ++td) { - if (tgt_region->size(td) > 1) - dim_map[td] = sd++; - else - dim_map[td] = -1; - } - } - - // Count expected messages: for each target dim, how many source chares overlap - int expected_msgs = 1; - for (int td = 0; td < N; ++td) { - if (dim_map[td] < 0) - continue; - int range_start = my_input.start[td]; - int range_stop = my_input.stop[td]; - int src_chare_start = range_start / src_tile; - int src_chare_stop = (range_stop + src_tile - 1) / src_tile; - expected_msgs *= (src_chare_stop - src_chare_start); - } - - DBG_PRINT("[Chare %d] target side: expected_msgs=%d\n", partition->index[0], - expected_msgs); - - auto pre_it = pending.find(node->id); - int pre_arrived = 0; - std::unordered_map>> pre_buffers; - if (pre_it != pending.end()) { - pre_arrived = -(pre_it->second.expected_msgs); - pre_buffers = std::move(pre_it->second.remote_buffers); - } - int remaining = expected_msgs - pre_arrived; - pending[node->id] = {node, remaining, {my_input}, std::move(pre_buffers)}; - if (remaining <= 0) - on_comm_done(node->id); - - } else { - node_finished(node->id); - } -} - -template -static void diag_send(PartitionImpl* partition, DAGNode* node, int source_name, - int result_name, int k_offset, ArrayDAGGroup* dag_group) { - auto src_it = partition->arrays.find(source_name); - if (src_it == partition->arrays.end() || src_it->second->local_size() == 0) - return; - - auto* src = static_cast*>(src_it->second); - - if constexpr (N_src == 1) { - int tile_2d = array_tile(dag_group->array_meta, result_name, 2); - - auto nd_idx = partition->nd_index(); - int my_start = src_it->second->decomp.chare_start_global(0, nd_idx[0]); - int local_size = src->region.size(0); - - struct ChareRange { - ChareIndex<2> target_ci; - int first_vec_idx, last_vec_idx; - int first_local, count; - }; - std::vector ranges; - for (int i = 0; i < local_size;) { - int vec_idx = my_start + i; - int row = (k_offset >= 0) ? vec_idx : vec_idx - k_offset; - int col = (k_offset >= 0) ? vec_idx + k_offset : vec_idx; - int target_row_chare = row / tile_2d; - int target_col_chare = col / tile_2d; - ChareIndex<2> target_ci; - target_ci.idx[0] = target_row_chare; - target_ci.idx[1] = target_col_chare; - - int j = i + 1; - while (j < local_size) { - int vj = my_start + j; - int rj = (k_offset >= 0) ? vj : vj - k_offset; - int cj = (k_offset >= 0) ? vj + k_offset : vj; - if (rj / tile_2d != target_row_chare || cj / tile_2d != target_col_chare) - break; - ++j; - } - ranges.push_back({target_ci, vec_idx, my_start + j - 1, i, j - i}); - i = j; - } - - for (auto& r : ranges) { - int first_row = (k_offset >= 0) ? r.first_vec_idx : r.first_vec_idx - k_offset; - int last_row = (k_offset >= 0) ? r.last_vec_idx : r.last_vec_idx - k_offset; - int first_col = (k_offset >= 0) ? r.first_vec_idx + k_offset : r.first_vec_idx; - int last_col = (k_offset >= 0) ? r.last_vec_idx + k_offset : r.last_vec_idx; - - int region_data[2 * 3]; - region_data[0] = first_row; - region_data[1] = last_row + 1; - region_data[2] = 1; - region_data[3] = first_col; - region_data[4] = last_col + 1; - region_data[5] = 1; - - int64_t byte_size = r.count * sizeof(T); -#ifndef NDEBUG - partition->comm_bytes_sent += byte_size; -#endif -#ifdef USE_KOKKOS - T* send_buf = static_cast(Kokkos::kokkos_malloc(byte_size)); - T* src_device = static_cast(src->device_data_ptr()) + r.first_local; - Kokkos::parallel_for( - CT_COMM_POLICY(partition, r.count), - KOKKOS_LAMBDA(int j) { send_buf[j] = src_device[j]; }); - device_pack_send<1, 2>(partition, dag_group->partition_proxy_2, r.target_ci, - node->id, 0, source_name, region_data, byte_size, send_buf); -#else - src->copyToHost(); - T* src_data = static_cast(src->data_ptr()); - T* send_buf = new T[r.count]; - for (int j = 0; j < r.count; ++j) - send_buf[j] = src_data[r.first_local + j]; - proxy_at<2>(dag_group->partition_proxy_2, r.target_ci) - .receive_data(node->id, /*input_index=*/0, source_name, 2, region_data, - byte_size, reinterpret_cast(send_buf)); - delete[] send_buf; -#endif - } - } else if constexpr (N_src == 2) { - int tile_1d = array_tile(dag_group->array_meta, result_name, 1); - - auto nd_idx = partition->nd_index(); - auto src_decomp = src_it->second->decomp; - int row_start = src_decomp.chare_start_global(0, nd_idx[0]); - int col_start = src_decomp.chare_start_global(1, nd_idx[1]); - int local_rows = src->region.size(0); - int local_cols = src->region.size(1); - - int row_lo = row_start; - int row_hi = row_start + local_rows; - int col_lo = col_start; - int col_hi = col_start + local_cols; - - int diag_lo, diag_hi; - if (k_offset >= 0) { - diag_lo = std::max(row_lo, col_lo - k_offset); - diag_hi = std::min(row_hi, col_hi - k_offset); - } else { - diag_lo = std::max(row_lo + k_offset, col_lo); - diag_hi = std::min(row_hi + k_offset, col_hi); - } - - if (diag_lo >= diag_hi) - return; - - struct ChareRange { - ChareIndex<1> target_ci; - int first_vec_idx; - int count; - }; - std::vector ranges; - for (int vec_idx = diag_lo; vec_idx < diag_hi;) { - int target_chare = vec_idx / tile_1d; - ChareIndex<1> target_ci; - target_ci.idx[0] = target_chare; - int end = std::min(diag_hi, (target_chare + 1) * tile_1d); - ranges.push_back({target_ci, vec_idx, end - vec_idx}); - vec_idx = end; - } - - for (auto& r : ranges) { - int region_data[1 * 3]; - region_data[0] = r.first_vec_idx; - region_data[1] = r.first_vec_idx + r.count; - region_data[2] = 1; - - int64_t byte_size = r.count * sizeof(T); -#ifndef NDEBUG - partition->comm_bytes_sent += byte_size; -#endif -#ifdef USE_KOKKOS - T* send_buf = static_cast(Kokkos::kokkos_malloc(byte_size)); - T* src_device = static_cast(src->device_data_ptr()); - int rs_cap = row_start, cs_cap = col_start, lc_cap = local_cols; - int ko_cap = k_offset; - int fvi_cap = r.first_vec_idx; - Kokkos::parallel_for( - CT_COMM_POLICY(partition, r.count), - KOKKOS_LAMBDA(int j) { - int vec_idx = fvi_cap + j; - int row = (ko_cap >= 0) ? vec_idx : vec_idx - ko_cap; - int col = (ko_cap >= 0) ? vec_idx + ko_cap : vec_idx; - int local_row = row - rs_cap; - int local_col = col - cs_cap; - send_buf[j] = src_device[local_row * lc_cap + local_col]; - }); - device_pack_send<2, 1>(partition, dag_group->partition_proxy_1, r.target_ci, - node->id, 0, source_name, region_data, byte_size, send_buf); -#else - src->copyToHost(); - T* src_data = static_cast(src->data_ptr()); - T* send_buf = new T[r.count]; - for (int j = 0; j < r.count; ++j) { - int vec_idx = r.first_vec_idx + j; - int row, col; - if (k_offset >= 0) { - row = vec_idx; - col = vec_idx + k_offset; - } else { - row = vec_idx - k_offset; - col = vec_idx; - } - int local_row = row - row_start; - int local_col = col - col_start; - send_buf[j] = src_data[local_row * local_cols + local_col]; - } - proxy_at<1>(dag_group->partition_proxy_1, r.target_ci) - .receive_data(node->id, /*input_index=*/0, source_name, 1, region_data, - byte_size, reinterpret_cast(send_buf)); - delete[] send_buf; -#endif - } - } -} - -template -static void dispatch_diag_send(DType dt, PartitionImpl* partition, DAGNode* node, - int source_name, int result_name, int k_offset, - ArrayDAGGroup* dag_group) { - switch (dt) { - case DType::FLOAT32: - diag_send(partition, node, source_name, result_name, k_offset, dag_group); - break; - case DType::FLOAT64: - diag_send(partition, node, source_name, result_name, k_offset, dag_group); - break; - case DType::INT32: - diag_send(partition, node, source_name, result_name, k_offset, dag_group); - break; - case DType::INT64: - diag_send(partition, node, source_name, result_name, k_offset, dag_group); - break; - } -} - -template -void ArrayDAGExecutorND::execute_diag_node(DAGNode* node) { - ASTNode* diag_root = nullptr; - for (ASTNode* root : node->ast->roots) { - if (static_cast(root->opcode) == Opcode::DIAG) { - diag_root = root; - break; - } - } - if (!diag_root) { - node_finished(node->id); - return; - } - - int input_name = diag_root->operands[0]->result_name; - int result_name = diag_root->result_name; - int k_offset = 0; - if (diag_root->operands.size() >= 2 && diag_root->operands[1]->is_scalar) - k_offset = (int)diag_root->operands[1]->scalar; - - auto* dag_group = static_cast(group); - DType dt = determine_dtype(node, partition->arrays); - - auto meta_it = dag_group->array_meta.find(input_name); - if (meta_it == dag_group->array_meta.end()) { - node_finished(node->id); - return; - } - int input_ndims = meta_it->second.ndims; - - DBG_PRINT("[Chare %d] execute_diag_node<%d>: node_id=%d input=%d(nd=%d) k=%d\n", - partition->index[0], N, node->id, input_name, input_ndims, k_offset); - - if (input_ndims == 1) { - if constexpr (N == 1) { - dispatch_diag_send<1>(dt, partition, node, input_name, result_name, k_offset, dag_group); - node_finished(node->id); - } else if constexpr (N == 2) { - int tile_1d = array_tile(dag_group->array_meta, input_name, 1); - - auto nd_idx = partition->nd_index(); - auto result_meta_it = dag_group->array_meta.find(result_name); - if (result_meta_it == dag_group->array_meta.end()) { - node_finished(node->id); - return; - } - int out_n = result_meta_it->second.global_shape[0]; - auto result_decomp = result_meta_it->second.template decomp<2>(); - auto result_chare = result_decomp.chare_region_global(nd_idx); - int row_start = result_chare.start[0]; - int col_start = result_chare.start[1]; - int local_rows = result_chare.stop[0] - row_start; - int local_cols = result_chare.stop[1] - col_start; - if (local_rows <= 0 || local_cols <= 0) { - node_finished(node->id); - return; - } - - int vec_len = meta_it->second.global_shape[0]; - int diag_lo, diag_hi; - if (k_offset >= 0) { - diag_lo = std::max(row_start, col_start - k_offset); - diag_hi = std::min(row_start + local_rows, col_start + local_cols - k_offset); - } else { - diag_lo = std::max(row_start + k_offset, col_start); - diag_hi = std::min(row_start + local_rows + k_offset, col_start + local_cols); - } - diag_lo = std::max(diag_lo, 0); - diag_hi = std::min(diag_hi, vec_len); - - if (diag_lo >= diag_hi) { - std::array out_start = {0, 0}; - std::array out_stop = {local_rows, local_cols}; - std::array out_step = {1, 1}; - std::array out_gs = {out_n, out_n}; - ArrayRegion<2> out_region(out_start, out_stop, out_step); - ArrayDecomp<2> out_decomp = dag_group->array_meta[result_name].template decomp<2>(); - partition->arrays[result_name] = - partition->allocate_or_reuse(out_region, out_gs, result_name, dt, out_decomp); -#ifndef USE_KOKKOS - std::memset(partition->arrays[result_name]->data_ptr(), 0, - local_rows * local_cols * dtype_size(dt)); -#endif - node_finished(node->id); - return; - } - - int first_1d_chare = diag_lo / tile_1d; - int last_1d_chare = (diag_hi - 1) / tile_1d; - int expected_msgs = last_1d_chare - first_1d_chare + 1; - - std::array my_start = {row_start, col_start}; - std::array my_stop = {row_start + local_rows, col_start + local_cols}; - std::array my_step = {1, 1}; - ArrayRegion<2> my_input(my_start, my_stop, my_step); - - auto pre_it = pending.find(node->id); - int pre_arrived = 0; - std::unordered_map>> pre_buffers; - if (pre_it != pending.end()) { - pre_arrived = -(pre_it->second.expected_msgs); - pre_buffers = std::move(pre_it->second.remote_buffers); - } - int remaining = expected_msgs - pre_arrived; - pending[node->id] = {node, remaining, {my_input}, std::move(pre_buffers)}; - if (remaining <= 0) - on_comm_done(node->id); - } - } else if (input_ndims == 2) { - if constexpr (N == 2) { - dispatch_diag_send<2>(dt, partition, node, input_name, result_name, k_offset, dag_group); - node_finished(node->id); - } else if constexpr (N == 1) { - int tile_2d = array_tile(dag_group->array_meta, input_name, 2); - - auto nd_idx = partition->nd_index(); - auto result_meta_it = dag_group->array_meta.find(result_name); - if (result_meta_it == dag_group->array_meta.end()) { - node_finished(node->id); - return; - } - int diag_len = result_meta_it->second.global_shape[0]; - auto result_decomp = result_meta_it->second.template decomp<1>(); - auto result_chare = result_decomp.chare_region_global(nd_idx); - int my_start_idx = result_chare.start[0]; - int local_size = result_chare.stop[0] - my_start_idx; - if (local_size <= 0) { - node_finished(node->id); - return; - } - - int my_end_idx = my_start_idx + local_size; - - std::unordered_set, ChareIndexHash<2>> source_chares; - int M = meta_it->second.global_shape[0]; - int N_cols = meta_it->second.global_shape[1]; - for (int vec_idx = my_start_idx; vec_idx < my_end_idx; ++vec_idx) { - int row, col; - if (k_offset >= 0) { - row = vec_idx; - col = vec_idx + k_offset; - } else { - row = vec_idx - k_offset; - col = vec_idx; - } - if (row >= 0 && row < M && col >= 0 && col < N_cols) { - int src_row_chare = row / tile_2d; - int src_col_chare = col / tile_2d; - ChareIndex<2> src_ci; - src_ci.idx[0] = src_row_chare; - src_ci.idx[1] = src_col_chare; - source_chares.insert(src_ci); - } - } - int expected_msgs = (int)source_chares.size(); - - std::array my_start_arr = {my_start_idx}; - std::array my_stop_arr = {my_end_idx}; - std::array my_step_arr = {1}; - ArrayRegion<1> my_input(my_start_arr, my_stop_arr, my_step_arr); - - auto pre_it = pending.find(node->id); - int pre_arrived = 0; - std::unordered_map>> pre_buffers; - if (pre_it != pending.end()) { - pre_arrived = -(pre_it->second.expected_msgs); - pre_buffers = std::move(pre_it->second.remote_buffers); - } - int remaining = expected_msgs - pre_arrived; - pending[node->id] = {node, remaining, {my_input}, std::move(pre_buffers)}; - if (remaining <= 0) - on_comm_done(node->id); - } - } else { - node_finished(node->id); - } -} - -template -struct PartitionProxyHelper; - -template <> -struct PartitionProxyHelper<1> { - static auto& get(ArrayDAGGroup* g) { return g->partition_proxy_1; } -}; -template <> -struct PartitionProxyHelper<2> { - static auto& get(ArrayDAGGroup* g) { return g->partition_proxy_2; } -}; -template <> -struct PartitionProxyHelper<3> { - static auto& get(ArrayDAGGroup* g) { return g->partition_proxy_3; } -}; - -template -static void tile_send(PartitionImpl* partition, DAGNode* node, int source_name, - int result_name, std::array const& reps, - std::array const& input_shape, int out_ndims, - ArrayDAGGroup* dag_group) { - auto src_it = partition->arrays.find(source_name); - if (src_it == partition->arrays.end() || src_it->second->local_size() == 0) - return; - - auto* src = static_cast*>(src_it->second); - auto nd_idx = partition->nd_index(); - auto src_decomp = src_it->second->decomp; - - std::array src_start, src_stop; - for (int d = 0; d < N_src; ++d) { - src_start[d] = src_decomp.chare_start_global(d, nd_idx[d]); - src_stop[d] = std::min(src_start[d] + src->region.size(d), - src_decomp.offset[d] + src_decomp.global_shape[d]); - } - - auto result_meta_it = dag_group->array_meta.find(result_name); - if (result_meta_it == dag_group->array_meta.end()) - return; - auto result_decomp = result_meta_it->second.template decomp(); - - int delta = out_ndims - N_src; - std::array out_shape; - for (int d = 0; d < N_tgt; ++d) - out_shape[d] = result_meta_it->second.global_shape[d]; - - int result_tile = result_meta_it->second.tile; - - std::array, N_tgt> chare_indices_per_dim; - for (int d = 0; d < N_tgt; ++d) { - int num_chares_d = result_decomp.num_chares(d); - if (d < delta) { - for (int ci = 0; ci < num_chares_d; ++ci) - chare_indices_per_dim[d].push_back(ci); - } else { - int src_d = d - delta; - int in_size = input_shape[src_d]; - for (int ci = 0; ci < num_chares_d; ++ci) { - int chare_lo = std::max(ci * result_tile, result_decomp.offset[d]); - int chare_hi = std::min((ci + 1) * result_tile, - result_decomp.offset[d] + out_shape[d]); - if (chare_lo >= chare_hi) - continue; - bool overlaps = false; - if (chare_hi - chare_lo >= in_size) { - overlaps = true; - } else { - int mod_lo = chare_lo % in_size; - int mod_hi = (chare_hi - 1) % in_size; - if (mod_lo <= mod_hi) { - overlaps = !(mod_hi < src_start[src_d] || mod_lo >= src_stop[src_d]); - } else { - overlaps = !(mod_hi < src_start[src_d] && mod_lo >= src_stop[src_d]); - } - } - if (overlaps) - chare_indices_per_dim[d].push_back(ci); - } - } - } - - int local_size = src->local_size(); - int64_t byte_size = local_size * sizeof(T); - - std::function&)> send_to_chares; - send_to_chares = [&](int dim, ChareIndex& ci) { - if (dim == N_tgt) { - int region_data[N_tgt * 3]; - for (int d = 0; d < N_tgt; ++d) { - if (d < delta) { - region_data[d * 3 + 0] = 0; - region_data[d * 3 + 1] = 1; - region_data[d * 3 + 2] = 1; - } else { - int src_d = d - delta; - region_data[d * 3 + 0] = src_start[src_d]; - region_data[d * 3 + 1] = src_stop[src_d]; - region_data[d * 3 + 2] = 1; - } - } - -#ifndef USE_KOKKOS - src->copyToHost(); - T* src_data = static_cast(src->data_ptr()); - T* send_buf = new T[local_size]; - std::memcpy(send_buf, src_data, byte_size); -#ifndef NDEBUG - partition->comm_bytes_sent += byte_size; -#endif - proxy_at(PartitionProxyHelper::get(dag_group), ci) - .receive_data(node->id, /*input_index=*/0, source_name, N_tgt, region_data, - byte_size, reinterpret_cast(send_buf)); - delete[] send_buf; -#else - T* send_buf = static_cast(Kokkos::kokkos_malloc(byte_size)); - T* src_device = static_cast(src->device_data_ptr()); - Kokkos::parallel_for( - CT_COMM_POLICY(partition, local_size), - KOKKOS_LAMBDA(int j) { send_buf[j] = src_device[j]; }); - device_pack_send(partition, - PartitionProxyHelper::get(dag_group), ci, - node->id, 0, source_name, region_data, byte_size, - send_buf); -#endif - return; - } - for (int idx : chare_indices_per_dim[dim]) { - ci.idx[dim] = idx; - send_to_chares(dim + 1, ci); - } - }; - - ChareIndex ci; - send_to_chares(0, ci); -} - -template -static void dispatch_tile_send(DType dt, PartitionImpl* partition, DAGNode* node, - int source_name, int result_name, - std::array const& reps, - std::array const& input_shape, int out_ndims, - ArrayDAGGroup* dag_group) { - switch (dt) { - case DType::FLOAT32: - tile_send(partition, node, source_name, result_name, reps, - input_shape, out_ndims, dag_group); - break; - case DType::FLOAT64: - tile_send(partition, node, source_name, result_name, reps, - input_shape, out_ndims, dag_group); - break; - case DType::INT32: - tile_send(partition, node, source_name, result_name, reps, - input_shape, out_ndims, dag_group); - break; - case DType::INT64: - tile_send(partition, node, source_name, result_name, reps, - input_shape, out_ndims, dag_group); - break; - } -} - -template -void ArrayDAGExecutorND::execute_tile_node(DAGNode* node) { - ASTNode* tile_root = nullptr; - for (ASTNode* root : node->ast->roots) { - if (static_cast(root->opcode) == Opcode::TILE) { - tile_root = root; - break; - } - } - if (!tile_root) { - node_finished(node->id); - return; - } - - int input_name = tile_root->operands[0]->result_name; - int result_name = tile_root->result_name; - int out_ndims = tile_root->ndims; - - std::array reps = {1, 1, 1}; - int num_reps = (int)tile_root->operands.size() - 1; - for (int d = 0; d < num_reps && d < 3; d++) { - if (tile_root->operands[d + 1]->is_scalar) - reps[d] = (int)tile_root->operands[d + 1]->scalar; - } - - auto* dag_group = static_cast(group); - DType dt = determine_dtype(node, partition->arrays); - - auto meta_it = dag_group->array_meta.find(input_name); - if (meta_it == dag_group->array_meta.end()) { - node_finished(node->id); - return; - } - int input_ndims = meta_it->second.ndims; - std::array input_shape = meta_it->second.global_shape; - - DBG_PRINT("[Chare %d] execute_tile_node<%d>: node_id=%d input=%d(nd=%d) result_nd=%d " - "reps=(%d,%d,%d)\n", - partition->index[0], N, node->id, input_name, input_ndims, out_ndims, - reps[0], reps[1], reps[2]); - - if (N == input_ndims) { - if (input_ndims == out_ndims) { - switch (N) { - case 1: - dispatch_tile_send<1, 1>(dt, reinterpret_cast*>(partition), - node, input_name, result_name, reps, input_shape, - out_ndims, dag_group); - break; - case 2: - dispatch_tile_send<2, 2>(dt, reinterpret_cast*>(partition), - node, input_name, result_name, reps, input_shape, - out_ndims, dag_group); - break; - case 3: - dispatch_tile_send<3, 3>(dt, reinterpret_cast*>(partition), - node, input_name, result_name, reps, input_shape, - out_ndims, dag_group); - break; - } - } else { - if (input_ndims == 1 && out_ndims == 2) { - dispatch_tile_send<1, 2>(dt, reinterpret_cast*>(partition), - node, input_name, result_name, reps, input_shape, - out_ndims, dag_group); - } else if (input_ndims == 1 && out_ndims == 3) { - dispatch_tile_send<1, 3>(dt, reinterpret_cast*>(partition), - node, input_name, result_name, reps, input_shape, - out_ndims, dag_group); - } else if (input_ndims == 2 && out_ndims == 3) { - dispatch_tile_send<2, 3>(dt, reinterpret_cast*>(partition), - node, input_name, result_name, reps, input_shape, - out_ndims, dag_group); - } - } - if (N != out_ndims) { - node_finished(node->id); - } - } - - if (N == out_ndims) { - auto nd_idx = partition->nd_index(); - - auto result_meta_it = dag_group->array_meta.find(result_name); - if (result_meta_it == dag_group->array_meta.end()) { - node_finished(node->id); - return; - } - auto result_decomp = result_meta_it->second.template decomp(); - auto result_chare = result_decomp.chare_region_global(nd_idx); - - bool has_extent = true; - for (int d = 0; d < N; ++d) { - if (result_chare.stop[d] <= result_chare.start[d]) { - has_extent = false; - break; - } - } - if (!has_extent) { - node_finished(node->id); - return; - } - - int delta = out_ndims - input_ndims; - auto input_decomp_meta = dag_group->array_meta.find(input_name); - int input_tile = input_decomp_meta->second.tile; - - int expected_msgs = 1; - for (int src_d = 0; src_d < input_ndims; ++src_d) { - int out_d = src_d + delta; - int in_size = input_shape[src_d]; - int chare_lo = result_chare.start[out_d]; - int chare_hi = result_chare.stop[out_d]; - - std::set needed_src_chares; - for (int p = chare_lo; p < chare_hi;) { - int inp_pos = p % in_size; - int src_chare = inp_pos / input_tile; - needed_src_chares.insert(src_chare); - int next_boundary = (src_chare + 1) * input_tile - inp_pos + p; - int next_wrap = p + (in_size - inp_pos); - p = std::min({next_boundary, next_wrap, chare_hi}); - } - expected_msgs *= (int)needed_src_chares.size(); - } - - ArrayRegion my_input; - for (int d = 0; d < N; ++d) { - my_input.start[d] = result_chare.start[d]; - my_input.stop[d] = result_chare.stop[d]; - my_input.step[d] = 1; - } - - auto pre_it = pending.find(node->id); - int pre_arrived = 0; - std::unordered_map>> pre_buffers; - if (pre_it != pending.end()) { - pre_arrived = -(pre_it->second.expected_msgs); - pre_buffers = std::move(pre_it->second.remote_buffers); - } - int remaining = expected_msgs - pre_arrived; - pending[node->id] = {node, remaining, {my_input}, std::move(pre_buffers)}; - if (remaining <= 0) - on_comm_done(node->id); - } -} - -template void ArrayDAGExecutorND<1>::execute_cross_set_region_node(DAGNode*); -template void ArrayDAGExecutorND<2>::execute_cross_set_region_node(DAGNode*); -template void ArrayDAGExecutorND<3>::execute_cross_set_region_node(DAGNode*); - -template void ArrayDAGExecutorND<1>::execute_diag_node(DAGNode*); -template void ArrayDAGExecutorND<2>::execute_diag_node(DAGNode*); -template void ArrayDAGExecutorND<3>::execute_diag_node(DAGNode*); - -template void ArrayDAGExecutorND<1>::execute_tile_node(DAGNode*); -template void ArrayDAGExecutorND<2>::execute_tile_node(DAGNode*); -template void ArrayDAGExecutorND<3>::execute_tile_node(DAGNode*); diff --git a/src/jit.cpp b/src/jit.cpp deleted file mode 100644 index b1b1e04..0000000 --- a/src/jit.cpp +++ /dev/null @@ -1,716 +0,0 @@ -#include "jit.hpp" - -MLIRJitCompiler::MLIRJitCompiler() { - // 1. Load required dialects - context.getOrLoadDialect(); - context.getOrLoadDialect(); - context.getOrLoadDialect(); - context.getOrLoadDialect(); - context.getOrLoadDialect(); - context.getOrLoadDialect(); - context.getOrLoadDialect(); - context.getOrLoadDialect(); - context.getOrLoadDialect(); - context.getOrLoadDialect(); - context.getOrLoadDialect(); - context.getOrLoadDialect(); - - // 2. Register LLVM IR translation interfaces (needed by ExecutionEngine) - mlir::DialectRegistry registry; - mlir::registerAllToLLVMIRTranslations(registry); - context.appendDialectRegistry(registry); - - // 3. Initialize Module - module = mlir::ModuleOp::create(mlir::UnknownLoc::get(&context)); - builder = std::make_unique(&context); -} - -void MLIRJitCompiler::buildFromAST(void* astPtr) { - using namespace mlir; - AST* ast = static_cast(astPtr); - auto loc = builder->getUnknownLoc(); - - // --- 1. Collect leaf nodes --- - std::vector memrefLeaves; // non-scalar leaves → memref func args - std::vector scalarLeaves; // scalar leaves → constants - std::vector broadcastLeaves; // size-1 arrays → scalar func args - collectLeaves(ast, memrefLeaves, scalarLeaves, broadcastLeaves); - - // Collect non-temp root operations as output slots. - // Each gets a caller-provided output memref appended after the input args. - std::vector outputRoots; - for (ASTNode* op : ast->roots) { - auto opc = static_cast(op->opcode); - if (opc != Opcode::NOOP && opc != Opcode::CREATE && !op->is_temp) - outputRoots.push_back(op); - } - - // Determine element type from the AST's dtype - mlir::Type elemType; - DType astDtype = DType::FLOAT32; - for (ASTNode* op : ast->roots) - if (op->dtype != DType::FLOAT32) { - astDtype = op->dtype; - break; - } - switch (astDtype) { - case DType::FLOAT64: - elemType = Float64Type::get(&context); - break; - case DType::INT32: - elemType = IntegerType::get(&context, 32); - break; - case DType::INT64: - elemType = IntegerType::get(&context, 64); - break; - default: - elemType = Float32Type::get(&context); - break; - } - - auto dynamicStridedMemRefType = [&](int64_t rank) { - SmallVector dynShape(rank, ShapedType::kDynamic); - SmallVector dynStrides(rank, ShapedType::kDynamic); - auto layout = StridedLayoutAttr::get(&context, ShapedType::kDynamic, dynStrides); - return MemRefType::get(dynShape, elemType, layout); - }; - - // --- 2. Create a func.func: memref inputs, broadcast scalars, then output memrefs --- - SmallVector argTypes; - for (const LeafUse& leaf : memrefLeaves) - argTypes.push_back(dynamicStridedMemRefType(leaf.operand->ndims)); - // Broadcast leaves are scalar function arguments (runtime values, not constants) - for (size_t i = 0; i < broadcastLeaves.size(); ++i) - argTypes.push_back(elemType); - for (ASTNode* root : outputRoots) - argTypes.push_back(dynamicStridedMemRefType(root->ndims)); - auto funcType = FunctionType::get(&context, argTypes, /*results=*/{}); - auto funcOp = func::FuncOp::create(loc, "fused_kernel", funcType); - module->push_back(funcOp); - - Block* entry = funcOp.addEntryBlock(); - builder->setInsertionPointToStart(entry); - - // Seed memref leaves with block arguments - nodeValues.clear(); - leafValues.clear(); - tempValues.clear(); - for (size_t i = 0; i < memrefLeaves.size(); ++i) - leafValues[leafUseKey(memrefLeaves[i].operand, memrefLeaves[i].region)] = - entry->getArgument(i); - - // Seed broadcast leaves with their scalar block arguments - for (size_t i = 0; i < broadcastLeaves.size(); ++i) - nodeValues[broadcastLeaves[i]->result_name] = - entry->getArgument(memrefLeaves.size() + i); - - // Seed scalar leaves as arith.constant values - for (ASTNode* sn : scalarLeaves) { - mlir::TypedAttr attr; - if (mlir::isa(elemType)) - attr = builder->getIntegerAttr(elemType, static_cast(sn->scalar)); - else - attr = builder->getFloatAttr(elemType, static_cast(sn->scalar)); - auto val = arith::ConstantOp::create(*builder, loc, attr); - nodeValues[sn->result_name] = val.getResult(); - } - - // --- 3. Process each operation sequentially (flat list) --- - for (ASTNode* op : ast->roots) - emitOp(op, loc); - - // --- 3b. Copy each non-temp root result into its caller-provided output arg --- - // - // Use linalg.generic instead of memref.copy: the default-layout memref - // type (memref) causes memref.copy to lower to a flat memcpy, - // which is incorrect when source and destination have different strides - // (e.g., a 14×14 contiguous temp copied into a 14×14 view of a 16×16 - // array with stride [16,1]). A linalg.generic copy lowers to affine - // loops that respect each operand's stride descriptor and participates - // in affine loop fusion, eliminating the intermediate allocation. - size_t outArgIdx = memrefLeaves.size() + broadcastLeaves.size(); - for (ASTNode* root : outputRoots) { - Value result = nodeValues[root->result_name]; - Value outArg = entry->getArgument(outArgIdx++); - if (mlir::isa(result.getType())) { - auto outType = mlir::cast(outArg.getType()); - int64_t rank = outType.getRank(); - auto identityMap = builder->getMultiDimIdentityMap(rank); - SmallVector copyMaps = {identityMap, identityMap}; - SmallVector copyIters(rank, utils::IteratorType::parallel); - linalg::GenericOp::create(*builder, loc, TypeRange{}, ValueRange{result}, - ValueRange{outArg}, copyMaps, copyIters, - [](OpBuilder& nb, Location nl, ValueRange args) { - linalg::YieldOp::create(nb, nl, args[0]); - }); - } else { - // Scalar result (e.g., SET_REGION with scalar RHS) — fill output - linalg::FillOp::create(*builder, loc, result, outArg); - } - } - - // --- 4. Deallocate temporaries (after all consumers have read) --- - for (mlir::Value tmp : tempValues) - memref::DeallocOp::create(*builder, loc, tmp); - - // --- 5. Terminate the function body --- - func::ReturnOp::create(*builder, loc); -} - -bool MLIRJitCompiler::optimizeAndFuse() { - mlir::PassManager pm(&context); - auto& funcPM = pm.nest(); - -#if defined(USE_NVIDIA) || defined(USE_AMD) || defined(USE_INTEL) - // GPU path: fuse at Linalg level, keeping linalg.generic ops intact - // for later conversion to scf.parallel → gpu.launch. - funcPM.addPass(mlir::createLinalgElementwiseOpFusionPass()); -#else - // CPU path: convert to Affine and use the powerful Affine Fusion pass. - // A. Convert Linalg to Affine loops - funcPM.addPass(mlir::createConvertLinalgToAffineLoopsPass()); - - // B. Maximal Loop Fusion: removes intermediate array results - // by merging producer loops into consumer loops. - funcPM.addPass(mlir::affine::createLoopFusionPass(0, /* fastMemorySpace */ - UINT64_MAX, /* localBufSizeThreshold */ - true, /* maximalFusion */ - mlir::affine::FusionMode::ProducerConsumer)); - - // C. Scalar Replacement: Promotes intermediate MemRefs to registers/SSA - funcPM.addPass(mlir::affine::createAffineScalarReplacementPass()); -#endif - - return mlir::succeeded(pm.run(*module)); -} - -std::unique_ptr MLIRJitCompiler::generateCPU() { - if (!lowerToCPU()) - return nullptr; - - // Initialize native target (required for JIT on the host CPU) - llvm::InitializeNativeTarget(); - llvm::InitializeNativeTargetAsmPrinter(); - - mlir::ExecutionEngineOptions engineOptions; - auto maybeEngine = mlir::ExecutionEngine::create(*module, engineOptions); - if (!maybeEngine) { - llvm::errs() << "Failed to create ExecutionEngine\n"; - return nullptr; - } - - return std::move(*maybeEngine); -} - -bool MLIRJitCompiler::loadCPU(std::unique_ptr& engine, - const char* kernelName, void** moduleOut, void** functionOut) { - if (!engine) - return false; - - auto funcPtrOrErr = engine->lookupPacked(kernelName); - if (!funcPtrOrErr) - return false; - - *moduleOut = (void*)engine.release(); - *functionOut = (void*)*funcPtrOrErr; - return true; -} - -void MLIRJitCompiler::dumpMLIR() { module->dump(); } - -// ---- Private helpers ---- - -void MLIRJitCompiler::collectLeaves(AST* ast, std::vector& memrefLeaves, - std::vector& scalarLeaves, - std::vector& broadcastLeaves) { - // Build set of result_names produced internally by roots in this fused kernel. - // Operands referencing these names are internal intermediates, not external inputs. - std::unordered_set internalResults; - for (ASTNode* root : ast->roots) - internalResults.insert(root->result_name); - - std::unordered_set seenBroadcasts; - - for (ASTNode* root : ast->roots) { - auto root_opc = static_cast(root->opcode); - - for (int op_idx = 0; op_idx < (int)root->operands.size(); ++op_idx) { - ASTNode* operand = root->operands[op_idx]; - if (!operand) - continue; - - if (internalResults.count(operand->result_name)) - continue; - - // Skip SET_REGION's target operand (operands[0]) — it's a write - // destination handled by the copy-back, not a JIT input. - if (root_opc == Opcode::SET_REGION && op_idx == 0) - continue; - - auto op = static_cast(operand->opcode); - if (op == Opcode::NOOP || op == Opcode::CREATE) { - if (operand->is_scalar) { - scalarLeaves.push_back(operand); - } else if (operand->is_broadcast) { - if (seenBroadcasts.insert(operand->result_name).second) - broadcastLeaves.push_back(operand); - } else { - memrefLeaves.push_back({operand, root->get_operand_region(op_idx)}); - } - } - } - } -} - -std::string MLIRJitCompiler::leafUseKey(ASTNode* operand, Region* region) const { - std::string key = std::to_string(operand->result_name) + "|" + std::to_string(operand->ndims); - if (region == nullptr) - return key + "|null"; - if (region->is_global) - return key + "|global"; - - switch (operand->ndims) { - case 1: - return key + "|" + fmt_region(*static_cast*>(region)); - case 2: - return key + "|" + fmt_region(*static_cast*>(region)); - case 3: - return key + "|" + fmt_region(*static_cast*>(region)); - default: - return key + "|nd=" + std::to_string(operand->ndims); - } -} - -mlir::Value MLIRJitCompiler::lookupOperand(ASTNode* operand, Region* region) { - if (region != nullptr) { - auto leafIt = leafValues.find(leafUseKey(operand, region)); - if (leafIt != leafValues.end()) - return leafIt->second; - } - auto it = nodeValues.find(operand->result_name); - if (it != nodeValues.end()) - return it->second; - return {}; // should not happen if AST is well-formed -} - -void MLIRJitCompiler::emitOp(ASTNode* node, mlir::Location loc) { - using namespace mlir; - - auto op = static_cast(node->opcode); - Value result; - - switch (op) { - case Opcode::NOOP: - case Opcode::CREATE: - return; - - case Opcode::SET_REGION: { - result = lookupOperand(node->operands[1], node->get_operand_region(1)); - break; - } - - default: { - // Binary elementwise ops - auto& binReg = binaryOpRegistry(); - auto binIt = binReg.find(node->opcode); - if (binIt != binReg.end()) { - Value lhs = lookupOperand(node->operands[0], node->get_operand_region(0)); - Value rhs = lookupOperand(node->operands[1], node->get_operand_region(1)); - - bool lhsIsScalar = !mlir::isa(lhs.getType()); - bool rhsIsScalar = !mlir::isa(rhs.getType()); - - auto elemType = lhsIsScalar ? lhs.getType() : - mlir::cast(lhs.getType()).getElementType(); - bool isInt = mlir::isa(elemType); - const auto& entry = binIt->second; - const auto& bodyFn = (isInt && entry.int_body) ? entry.int_body : entry.float_body; - - if (lhsIsScalar && rhsIsScalar) { - // Both operands are scalars — emit a plain scalar op - result = bodyFn(*builder, loc, lhs, rhs); - break; - } - - Value arrayOperand = lhsIsScalar ? rhs : lhs; - auto srcType = mlir::cast(arrayOperand.getType()); - int64_t rank = srcType.getRank(); - SmallVector dynShape(rank, ShapedType::kDynamic); - auto outType = MemRefType::get(dynShape, elemType); - - SmallVector dimSizes; - for (int64_t d = 0; d < rank; ++d) - dimSizes.push_back(memref::DimOp::create(*builder, loc, arrayOperand, d).getResult()); - result = memref::AllocOp::create(*builder, loc, outType, ValueRange(dimSizes)); - - emitGenericArrayOp(lhs, rhs, result, bodyFn); - - if (node->is_temp) - tempValues.push_back(result); - break; - } - - // Unary elementwise ops - auto& unReg = unaryOpRegistry(); - auto unIt = unReg.find(node->opcode); - if (unIt != unReg.end()) { - Value input = lookupOperand(node->operands[0], node->get_operand_region(0)); - auto srcType = mlir::cast(input.getType()); - int64_t rank = srcType.getRank(); - auto elemType = srcType.getElementType(); - SmallVector dynShape(rank, ShapedType::kDynamic); - auto outType = MemRefType::get(dynShape, elemType); - - SmallVector dimSizes; - for (int64_t d = 0; d < rank; ++d) - dimSizes.push_back(memref::DimOp::create(*builder, loc, input, d).getResult()); - result = memref::AllocOp::create(*builder, loc, outType, ValueRange(dimSizes)); - - bool isInt = mlir::isa(elemType); - const auto& entry = unIt->second; - const auto& bodyFn = (isInt && entry.int_body) ? entry.int_body : entry.float_body; - emitGenericUnaryOp(input, result, bodyFn); - - if (node->is_temp) - tempValues.push_back(result); - break; - } - - llvm::errs() << "Unknown opcode in JIT emitOp: " << node->opcode << "\n"; - return; - } - } - - nodeValues[node->result_name] = result; -} - -bool MLIRJitCompiler::lowerToCPU() { - mlir::PassManager pm(&context); - // Lower any remaining Linalg ops to Affine loops (no-op if optimizeAndFuse() ran first) - pm.addPass(mlir::createConvertLinalgToAffineLoopsPass()); - // Lower Affine to SCF (structured control flow) - pm.addPass(mlir::createLowerAffinePass()); - // Lower SCF to CF (unstructured control flow) - pm.addPass(mlir::createSCFToControlFlowPass()); - // Decompose memref.subview and similar ops into lower-level memref ops - pm.addPass(mlir::memref::createExpandStridedMetadataPass()); - // Lower CF, Arith, MemRef, Func to LLVM dialect - pm.addPass(mlir::createConvertControlFlowToLLVMPass()); - pm.addPass(mlir::createConvertMathToLLVMPass()); - pm.addPass(mlir::createArithToLLVMConversionPass()); - pm.addPass(mlir::createConvertIndexToLLVMPass()); - pm.addPass(mlir::createFinalizeMemRefToLLVMConversionPass()); - pm.addPass(mlir::createConvertFuncToLLVMPass()); - // Reconcile unrealized conversion casts left by partial lowerings - pm.addPass(mlir::createReconcileUnrealizedCastsPass()); - return mlir::succeeded(pm.run(*module)); -} - -// ---- Elementwise op registries ---- - -const std::unordered_map& MLIRJitCompiler::binaryOpRegistry() { - static const std::unordered_map registry = { - {(int)Opcode::ADD, - {[](mlir::OpBuilder& b, mlir::Location loc, mlir::Value lhs, mlir::Value rhs) { - return mlir::arith::AddFOp::create(b, loc, lhs, rhs).getResult(); - }, - [](mlir::OpBuilder& b, mlir::Location loc, mlir::Value lhs, mlir::Value rhs) { - return mlir::arith::AddIOp::create(b, loc, lhs, rhs).getResult(); - }}}, - {(int)Opcode::SUB, - {[](mlir::OpBuilder& b, mlir::Location loc, mlir::Value lhs, mlir::Value rhs) { - return mlir::arith::SubFOp::create(b, loc, lhs, rhs).getResult(); - }, - [](mlir::OpBuilder& b, mlir::Location loc, mlir::Value lhs, mlir::Value rhs) { - return mlir::arith::SubIOp::create(b, loc, lhs, rhs).getResult(); - }}}, - {(int)Opcode::MUL, - {[](mlir::OpBuilder& b, mlir::Location loc, mlir::Value lhs, mlir::Value rhs) { - return mlir::arith::MulFOp::create(b, loc, lhs, rhs).getResult(); - }, - [](mlir::OpBuilder& b, mlir::Location loc, mlir::Value lhs, mlir::Value rhs) { - return mlir::arith::MulIOp::create(b, loc, lhs, rhs).getResult(); - }}}, - {(int)Opcode::DIV, - {[](mlir::OpBuilder& b, mlir::Location loc, mlir::Value lhs, mlir::Value rhs) { - return mlir::arith::DivFOp::create(b, loc, lhs, rhs).getResult(); - }, - [](mlir::OpBuilder& b, mlir::Location loc, mlir::Value lhs, mlir::Value rhs) { - return mlir::arith::DivSIOp::create(b, loc, lhs, rhs).getResult(); - }}}, - }; - return registry; -} - -const std::unordered_map& MLIRJitCompiler::unaryOpRegistry() { - static const std::unordered_map registry = { - {(int)Opcode::TANH, - {[](mlir::OpBuilder& b, mlir::Location loc, mlir::Value v) { - return mlir::math::TanhOp::create(b, loc, v).getResult(); - }, - nullptr}}, - {(int)Opcode::EXP, - {[](mlir::OpBuilder& b, mlir::Location loc, mlir::Value v) { - return mlir::math::ExpOp::create(b, loc, v).getResult(); - }, - nullptr}}, - }; - return registry; -} - -#ifdef JIT_ENABLE_GPU_BACKEND - -std::string MLIRJitCompiler::generateNVIDIA() { - if (!lowerToGPU()) - return ""; - - // Lower ops inside gpu.module to LLVM + NVVM dialects - mlir::PassManager pm(&context); - auto& gpuPM = pm.nest(); - gpuPM.addPass(mlir::createConvertMathToLLVMPass()); - gpuPM.addPass(mlir::createArithToLLVMConversionPass()); - gpuPM.addPass(mlir::createConvertIndexToLLVMPass()); - gpuPM.addPass(mlir::createFinalizeMemRefToLLVMConversionPass()); - gpuPM.addPass(mlir::createConvertFuncToLLVMPass()); - gpuPM.addPass(mlir::createConvertGpuOpsToNVVMOps()); - gpuPM.addPass(mlir::createReconcileUnrealizedCastsPass()); - if (mlir::failed(pm.run(*module))) - return ""; - - return serializeToPTX(); -} - -std::string MLIRJitCompiler::generateAMD() { - if (!lowerToGPU()) - return ""; - - // Lower ops inside gpu.module to LLVM + ROCDL dialects - mlir::PassManager pm(&context); - auto& gpuPM = pm.nest(); - gpuPM.addPass(mlir::createConvertMathToLLVMPass()); - gpuPM.addPass(mlir::createArithToLLVMConversionPass()); - gpuPM.addPass(mlir::createConvertIndexToLLVMPass()); - gpuPM.addPass(mlir::createFinalizeMemRefToLLVMConversionPass()); - gpuPM.addPass(mlir::createConvertFuncToLLVMPass()); - gpuPM.addPass(mlir::createConvertGpuOpsToROCDLOps()); - gpuPM.addPass(mlir::createReconcileUnrealizedCastsPass()); - if (mlir::failed(pm.run(*module))) - return ""; - - return serializeToGCN(); -} - -std::string MLIRJitCompiler::generateIntel() { - if (!lowerToGPU()) - return ""; - - // Convert GPU module contents to SPIR-V - mlir::PassManager pm(&context); - pm.addPass(mlir::createConvertGPUToSPIRVPass()); - auto& spvPM = pm.nest(); - spvPM.addPass(mlir::spirv::createSPIRVLowerABIAttributesPass()); - spvPM.addPass(mlir::spirv::createSPIRVUpdateVCEPass()); - if (mlir::failed(pm.run(*module))) - return ""; - - return serializeToSPIRV(); -} - -bool MLIRJitCompiler::loadNVIDIA(const std::string& ptx, const char* kernelName, void** moduleOut, - void** functionOut) { -#if defined(USE_NVIDIA) - CUmodule cuModule; - CUfunction cuFunction; - CUresult res; - - res = cuModuleLoadDataEx(&cuModule, ptx.c_str(), 0, 0, 0); - if (res != CUDA_SUCCESS) - return false; - - res = cuModuleGetFunction(&cuFunction, cuModule, kernelName); - if (res != CUDA_SUCCESS) - return false; - - *moduleOut = (void*)cuModule; - *functionOut = (void*)cuFunction; - return true; -#else - return false; -#endif -} - -bool MLIRJitCompiler::loadAMD(const std::string& gcn, const char* kernelName, void** moduleOut, - void** functionOut) { -#if defined(USE_AMD) - hipModule_t hipModule; - hipFunction_t hipFunction; - hipError_t res; - - res = hipModuleLoadData(&hipModule, gcn.data()); - if (res != hipSuccess) - return false; - - res = hipModuleGetFunction(&hipFunction, hipModule, kernelName); - if (res != hipSuccess) - return false; - - *moduleOut = (void*)hipModule; - *functionOut = (void*)hipFunction; - return true; -#else - return false; -#endif -} - -bool MLIRJitCompiler::loadIntel(const std::string& spirv, const char* kernelName, void** moduleOut, - void** functionOut) { -#if defined(USE_INTEL) - // Initialize Level Zero - ze_result_t res = zeInit(ZE_INIT_FLAG_GPU_ONLY); - if (res != ZE_RESULT_SUCCESS) - return false; - - // Get driver - uint32_t driverCount = 1; - ze_driver_handle_t driver; - res = zeDriverGet(&driverCount, &driver); - if (res != ZE_RESULT_SUCCESS) - return false; - - // Get device - uint32_t deviceCount = 1; - ze_device_handle_t device; - res = zeDeviceGet(driver, &deviceCount, &device); - if (res != ZE_RESULT_SUCCESS) - return false; - - // Create context - ze_context_desc_t ctxDesc = {ZE_STRUCTURE_TYPE_CONTEXT_DESC, nullptr, 0}; - ze_context_handle_t zeContext; - res = zeContextCreate(driver, &ctxDesc, &zeContext); - if (res != ZE_RESULT_SUCCESS) - return false; - - // Create module from SPIR-V binary - ze_module_desc_t moduleDesc = {}; - moduleDesc.stype = ZE_STRUCTURE_TYPE_MODULE_DESC; - moduleDesc.format = ZE_MODULE_FORMAT_IL_SPIRV; - moduleDesc.inputSize = spirv.size(); - moduleDesc.pInputModule = reinterpret_cast(spirv.data()); - - ze_module_handle_t zeModule; - ze_module_build_log_handle_t buildLog; - res = zeModuleCreate(zeContext, device, &moduleDesc, &zeModule, &buildLog); - if (res != ZE_RESULT_SUCCESS) - return false; - - // Get kernel - ze_kernel_desc_t kernelDesc = {}; - kernelDesc.stype = ZE_STRUCTURE_TYPE_KERNEL_DESC; - kernelDesc.pKernelName = kernelName; - ze_kernel_handle_t zeKernel; - res = zeKernelCreate(zeModule, &kernelDesc, &zeKernel); - if (res != ZE_RESULT_SUCCESS) - return false; - - *moduleOut = reinterpret_cast(zeModule); - *functionOut = reinterpret_cast(zeKernel); - return true; -#else - return false; -#endif -} - -bool MLIRJitCompiler::lowerToGPU() { - mlir::PassManager pm(&context); - // Convert linalg.generic to scf.parallel (NOT affine.for) - pm.addPass(mlir::createConvertLinalgToParallelLoopsPass()); - // Annotate scf.parallel loops with GPU mapping attributes (block/thread dims) - pm.nest().addPass(mlir::createGpuMapParallelLoopsPass()); - // Convert annotated scf.parallel to gpu.launch - pm.addPass(mlir::createConvertParallelLoopToGpuPass()); - // Outline gpu.launch bodies into gpu.func inside gpu.module - pm.addPass(mlir::createGpuKernelOutliningPass()); - return mlir::succeeded(pm.run(*module)); -} - -std::string MLIRJitCompiler::translateToTarget(llvm::StringRef triple, llvm::StringRef cpu, - llvm::StringRef features, - llvm::CodeGenFileType fileType) { - // 1. Translate MLIR module to LLVM IR - llvm::LLVMContext llvmContext; - auto llvmModule = mlir::translateModuleToLLVMIR(*module, llvmContext); - if (!llvmModule) - return ""; - - llvmModule->setTargetTriple(triple); - - // 2. Look up the LLVM target - std::string error; - const llvm::Target* target = llvm::TargetRegistry::lookupTarget(triple, error); - if (!target) - return ""; - - // 3. Create a TargetMachine - llvm::TargetOptions opts; - std::unique_ptr tm( - target->createTargetMachine(triple, cpu, features, opts, llvm::Reloc::Model::PIC_)); - if (!tm) - return ""; - - llvmModule->setDataLayout(tm->createDataLayout()); - - // 4. Emit to a string (assembly or object depending on fileType) - std::string outStr; - llvm::raw_string_ostream os(outStr); - llvm::legacy::PassManager pass; - if (tm->addPassesToEmitFile(pass, os, nullptr, fileType)) - return ""; - - pass.run(*llvmModule); - os.flush(); - return outStr; -} - -std::string MLIRJitCompiler::serializeToPTX() { - llvm::InitializeAllTargets(); - llvm::InitializeAllTargetMCs(); - llvm::InitializeAllAsmPrinters(); - mlir::registerNVVMDialectTranslation(context); - - return translateToTarget( - /*triple=*/"nvptx64-nvidia-cuda", - /*cpu=*/"sm_70", - /*features=*/"+ptx60", llvm::CodeGenFileType::AssemblyFile); -} - -std::string MLIRJitCompiler::serializeToGCN() { - llvm::InitializeAllTargets(); - llvm::InitializeAllTargetMCs(); - llvm::InitializeAllAsmPrinters(); - mlir::registerROCDLDialectTranslation(context); - - // AMD HSA requires an ELF code object, not assembly text - return translateToTarget( - /*triple=*/"amdgcn-amd-amdhsa", - /*cpu=*/"gfx908", - /*features=*/"", llvm::CodeGenFileType::ObjectFile); -} - -std::string MLIRJitCompiler::serializeToSPIRV() { - std::string result; - - // Walk the module looking for spirv.module ops and serialize each one - module->walk([&](mlir::spirv::ModuleOp spvModule) { - llvm::SmallVector binary; - if (mlir::succeeded(mlir::spirv::serialize(spvModule, binary))) { - // Append raw SPIR-V words to the output - result.append(reinterpret_cast(binary.data()), - binary.size() * sizeof(uint32_t)); - } - }); - - return result; -} - -#endif // JIT_ENABLE_GPU_BACKEND diff --git a/src/jit.hpp b/src/jit.hpp deleted file mode 100644 index b294e03..0000000 --- a/src/jit.hpp +++ /dev/null @@ -1,228 +0,0 @@ -#pragma once - -#include -#include -#include -#include -#include -#include - -// LLVM Backend -#include "llvm/CodeGen/CommandFlags.h" -#include "llvm/IR/LLVMContext.h" -#include "llvm/IR/Module.h" -#include "llvm/MC/TargetRegistry.h" -#include "llvm/Support/MemoryBuffer.h" -#include "llvm/Support/TargetSelect.h" -#include "llvm/Support/raw_ostream.h" -#include "llvm/Target/TargetMachine.h" -#include "llvm/Target/TargetOptions.h" - -// MLIR Execution Engine -#include "mlir/ExecutionEngine/ExecutionEngine.h" -#include "mlir/ExecutionEngine/OptUtils.h" - -// MLIR Translation -#include "mlir/Target/LLVMIR/Dialect/All.h" -#include "mlir/Target/LLVMIR/Dialect/NVVM/NVVMToLLVMIRTranslation.h" -#include "mlir/Target/LLVMIR/Dialect/ROCDL/ROCDLToLLVMIRTranslation.h" -#include "mlir/Target/LLVMIR/Export.h" - -// SPIR-V Serialization -#include "mlir/Dialect/SPIRV/IR/SPIRVOps.h" -#include "mlir/Target/SPIRV/Serialization.h" - -// MLIR Core -#include "mlir/Dialect/Func/IR/FuncOps.h" -#include "mlir/IR/BuiltinAttributes.h" -#include "mlir/IR/Builders.h" -#include "mlir/IR/BuiltinOps.h" -#include "mlir/IR/MLIRContext.h" -#include "mlir/Pass/PassManager.h" - -// Dialects -#include "mlir/Dialect/Affine/IR/AffineOps.h" -#include "mlir/Dialect/Arith/IR/Arith.h" -#include "mlir/Dialect/GPU/IR/GPUDialect.h" -#include "mlir/Dialect/LLVMIR/LLVMDialect.h" -#include "mlir/Dialect/LLVMIR/NVVMDialect.h" -#include "mlir/Dialect/LLVMIR/ROCDLDialect.h" -#include "mlir/Dialect/Linalg/IR/Linalg.h" -#include "mlir/Dialect/Math/IR/Math.h" -#include "mlir/Dialect/MemRef/IR/MemRef.h" -#include "mlir/Dialect/SCF/IR/SCF.h" -#include "mlir/Dialect/SPIRV/IR/SPIRVDialect.h" - -// Passes & Conversions -#include "mlir/Conversion/Passes.h" // umbrella: declares all conversion pass factories -#include "mlir/Dialect/Affine/Transforms/Passes.h" -#include "mlir/Dialect/GPU/Transforms/Passes.h" -#include "mlir/Dialect/Linalg/Passes.h" -#include "mlir/Dialect/MemRef/Transforms/Passes.h" -#include "mlir/Dialect/SPIRV/Transforms/Passes.h" -#include "mlir/Transforms/Passes.h" - -// Application AST / Opcodes -#include "array_region.hpp" -#include "charmtyles/core/dag.hpp" -#include "opcodes.hpp" - -#if defined(USE_NVIDIA) -#include -#elif defined(USE_AMD) -#include -#elif defined(USE_INTEL) -#include -#include -#endif - -/// Registry entry for a binary elementwise op. -struct BinaryOpEntry { - using Fn = std::function; - Fn float_body; - Fn int_body; // nullptr if not supported for integers -}; - -/// Registry entry for a unary elementwise op. -struct UnaryOpEntry { - using Fn = std::function; - Fn float_body; - Fn int_body; // nullptr if not supported for integers -}; - -struct LeafUse { - ASTNode* operand; - Region* region; -}; - -/** - * @brief A JIT Compiler to transform Array ASTs into Fused GPU Kernels. - */ -class MLIRJitCompiler { - public: - MLIRJitCompiler(); - - void buildFromAST(void* astPtr); - bool optimizeAndFuse(); - std::unique_ptr generateCPU(); - bool loadCPU(std::unique_ptr& engine, const char* kernelName, - void** moduleOut, void** functionOut); - void dumpMLIR(); - -#ifdef JIT_ENABLE_GPU_BACKEND - std::string generateNVIDIA(); - std::string generateAMD(); - std::string generateIntel(); - bool loadNVIDIA(const std::string& ptx, const char* kernelName, void** moduleOut, - void** functionOut); - bool loadAMD(const std::string& gcn, const char* kernelName, void** moduleOut, - void** functionOut); - bool loadIntel(const std::string& spirv, const char* kernelName, void** moduleOut, - void** functionOut); -#endif - - /// Op registries: map opcode → body builder lambdas. - static const std::unordered_map& binaryOpRegistry(); - static const std::unordered_map& unaryOpRegistry(); - - private: - mlir::MLIRContext context; - mlir::OwningOpRef module; - std::unique_ptr builder; - - /// Maps AST node result_name → MLIR Value (memref or scalar) produced by that node. - std::unordered_map nodeValues; - /// Distinguish repeated views of the same array by the carried operand region. - std::unordered_map leafValues; - std::vector tempValues; // deferred deallocs for temporaries - - // AST helpers - void collectLeaves(AST* ast, std::vector& memrefLeaves, - std::vector& scalarLeaves, - std::vector& broadcastLeaves); - std::string leafUseKey(ASTNode* operand, Region* region) const; - mlir::Value lookupOperand(ASTNode* operand, Region* region); - void emitOp(ASTNode* node, mlir::Location loc); - - // Lowering helpers - bool lowerToCPU(); -#ifdef JIT_ENABLE_GPU_BACKEND - bool lowerToGPU(); - std::string - translateToTarget(llvm::StringRef triple, llvm::StringRef cpu, llvm::StringRef features, - llvm::CodeGenFileType fileType = llvm::CodeGenFileType::AssemblyFile); - std::string serializeToPTX(); - std::string serializeToGCN(); - std::string serializeToSPIRV(); -#endif - - // Binary element-wise op helper (template must be in header) - template - void emitGenericArrayOp(mlir::Value a, mlir::Value b, mlir::Value out, - BodyBuilderFn bodyBuilder) { - using namespace mlir; - auto loc = builder->getUnknownLoc(); - - // Derive rank from the output memref (always a memref) - auto outType = mlir::cast(out.getType()); - int64_t rank = outType.getRank(); - - bool aIsScalar = !mlir::isa(a.getType()); - bool bIsScalar = !mlir::isa(b.getType()); - - // Identity map for memref operands; empty-result map for 0-d scalar memrefs - auto identityMap = builder->getMultiDimIdentityMap(rank); - auto scalarMap0 = AffineMap::get(rank, /*symbolCount=*/0, /*results=*/{}, &context); - - // linalg.generic requires shaped types (MemRef/Tensor) for all operands. - // Wrap bare scalar values in a 0-d memref so they broadcast correctly. - auto wrapScalar = [&](Value v) -> Value { - auto zeroD = MemRefType::get({}, v.getType()); - auto mem = memref::AllocaOp::create(*builder, loc, zeroD); - memref::StoreOp::create(*builder, loc, v, mem, ValueRange{}); - return mem; - }; - - SmallVector indexingMaps; - indexingMaps.push_back(aIsScalar ? scalarMap0 : identityMap); - indexingMaps.push_back(bIsScalar ? scalarMap0 : identityMap); - indexingMaps.push_back(identityMap); // output is always memref - - SmallVector iterTypes(rank, utils::IteratorType::parallel); - - SmallVector inputs; - SmallVector outputs{out}; - inputs.push_back(aIsScalar ? wrapScalar(a) : a); - inputs.push_back(bIsScalar ? wrapScalar(b) : b); - - linalg::GenericOp::create( - *builder, - loc, TypeRange{}, ValueRange(inputs), ValueRange(outputs), indexingMaps, iterTypes, - [&](OpBuilder& nestedBuilder, Location nestedLoc, ValueRange args) { - Value result = bodyBuilder(nestedBuilder, nestedLoc, args[0], args[1]); - linalg::YieldOp::create(nestedBuilder, nestedLoc, result); - }); - } - - // Unary element-wise op helper - template - void emitGenericUnaryOp(mlir::Value input, mlir::Value out, BodyBuilderFn bodyBuilder) { - using namespace mlir; - auto loc = builder->getUnknownLoc(); - - auto outType = mlir::cast(out.getType()); - int64_t rank = outType.getRank(); - - auto identityMap = builder->getMultiDimIdentityMap(rank); - SmallVector indexingMaps = {identityMap, identityMap}; - SmallVector iterTypes(rank, utils::IteratorType::parallel); - - linalg::GenericOp::create( - *builder, - loc, TypeRange{}, ValueRange{input}, ValueRange{out}, indexingMaps, iterTypes, - [&](OpBuilder& nestedBuilder, Location nestedLoc, ValueRange args) { - Value result = bodyBuilder(nestedBuilder, nestedLoc, args[0]); - linalg::YieldOp::create(nestedBuilder, nestedLoc, result); - }); - } -}; diff --git a/src/opcodes.hpp b/src/opcodes.hpp deleted file mode 100644 index 1c82258..0000000 --- a/src/opcodes.hpp +++ /dev/null @@ -1,49 +0,0 @@ -#pragma once - -enum class Opcode { - COPY = -5, - SET_REGION = -4, - GET_REGION = -3, - CREATE = -2, - NOOP = -1, - ADD = 0, - SUB = 1, - MUL = 2, - DIV = 3, - MATMUL = 4, - TANH = 5, - EXP = 6, - TILE = 7, - REDUCE = 8, - MATMATMUL = 9, - DIAG = 10 -}; - -/// Returns true for binary elementwise ops. -inline bool is_binary_elementwise(Opcode op) { - switch (op) { - case Opcode::ADD: - case Opcode::SUB: - case Opcode::MUL: - case Opcode::DIV: - return true; - default: - return false; - } -} - -/// Returns true for unary elementwise ops. -inline bool is_unary_elementwise(Opcode op) { - switch (op) { - case Opcode::TANH: - case Opcode::EXP: - return true; - default: - return false; - } -} - -/// Returns true for any elementwise op (binary or unary). -inline bool is_elementwise(Opcode op) { - return is_binary_elementwise(op) || is_unary_elementwise(op); -} diff --git a/src/partition_comm.cpp b/src/partition_comm.cpp deleted file mode 100644 index fe1cbb6..0000000 --- a/src/partition_comm.cpp +++ /dev/null @@ -1,136 +0,0 @@ -#include "backend_internal.hpp" -#include "dispatch.hpp" - -#include - -#ifdef USE_KOKKOS -template -void PartitionImpl::receive_data(int& node_id, int& input_index, int& name, int& ndims, - int*& region_data, int64_t& size, char*& data, - CkDeviceBufferPost* devicePost) { - // Post entry method: allocate device buffer for incoming GPU data - data = static_cast(Kokkos::kokkos_malloc(size)); -} - -template -void PartitionImpl::send_complete(int send_id) { - auto it = pending_sends.find(send_id); - if (it != pending_sends.end()) { - Kokkos::kokkos_free(it->second); - pending_sends.erase(it); - } -} -#endif - -template -void PartitionImpl::receive_data(int node_id, int input_index, int name, int ndims, - int* region_data, int64_t size, char* data) { - PendingComm& comm = executor->pending[node_id]; - -#ifdef USE_KOKKOS - // data is already a device pointer (allocated in post method, filled by GPU messaging) - char* buf = data; -#else - char* buf = new char[size]; - memcpy(buf, data, size); -#endif - - // Unpack region from region_data: [start0, stop0, step0, start1, stop1, step1, ...] - std::array start, stop, step; - for (int d = 0; d < N; ++d) { - start[d] = region_data[d * 3 + 0]; - stop[d] = region_data[d * 3 + 1]; - step[d] = region_data[d * 3 + 2]; - } - ArrayRegion region(start, stop, step); - - comm.remote_buffers[input_index].push_back({buf, region, size}); - - if (comm.incremental) { - // Incremental mode: compute with this panel immediately - executor->on_matmatmul_receive(node_id, input_index); - } - - if (comm.expected_msgs > 0) { - comm.expected_msgs--; - if (comm.expected_msgs == 0) - executor->on_comm_done(node_id); - } else { - comm.expected_msgs--; - } -} - -template -void PartitionImpl::comm_done(int node_id) { - executor->on_comm_done(node_id); -} - -template -void PartitionImpl::reduce_result(CkReductionMsg* msg) { - if constexpr (N == 1) { - // Extract metadata and summed value from the custom reducer's output - ReduceContrib result; - memcpy(&result, msg->getData(), sizeof(ReduceContrib)); - int node_id = result.node_id; - int result_name = result.result_name; - DType dt = static_cast(result.dtype_int); - - // Create the scalar result array (size 1) on chare 0 - if (arrays.find(result_name) == arrays.end()) { - auto* dag_group = static_cast(dag_proxy.ckLocalBranch()); - ArrayDecomp<1> out_decomp = dag_group->array_meta[result_name].decomp<1>(); - std::array out_start = {0}; - std::array out_stop = {1}; - std::array out_step = {1}; - std::array out_gs = {1}; - ArrayRegion<1> out_region(out_start, out_stop, out_step); - arrays[result_name] = - allocate_or_reuse(out_region, out_gs, result_name, dt, out_decomp); - } - - // Store the reduced value - int elem_size = dtype_size(dt); -#ifdef USE_KOKKOS - Kokkos::deep_copy( - Kokkos::View>( - static_cast(arrays[result_name]->device_data_ptr()), elem_size), - Kokkos::View>( - result.value, elem_size)); -#else - memcpy(arrays[result_name]->data_ptr(), result.value, elem_size); -#endif - - DBG_PRINT("[1D Chare %d] reduce_result: node=%d result_name=%d stored\n", - index[0], node_id, result_name); - - delete msg; - executor->node_finished(node_id); - } else { - delete msg; - } -} - -#ifdef USE_KOKKOS -template void PartitionImpl<1>::receive_data(int&, int&, int&, int&, int*&, int64_t&, char*&, - CkDeviceBufferPost*); -template void PartitionImpl<2>::receive_data(int&, int&, int&, int&, int*&, int64_t&, char*&, - CkDeviceBufferPost*); -template void PartitionImpl<3>::receive_data(int&, int&, int&, int&, int*&, int64_t&, char*&, - CkDeviceBufferPost*); - -template void PartitionImpl<1>::send_complete(int); -template void PartitionImpl<2>::send_complete(int); -template void PartitionImpl<3>::send_complete(int); -#endif - -template void PartitionImpl<1>::receive_data(int, int, int, int, int*, int64_t, char*); -template void PartitionImpl<2>::receive_data(int, int, int, int, int*, int64_t, char*); -template void PartitionImpl<3>::receive_data(int, int, int, int, int*, int64_t, char*); - -template void PartitionImpl<1>::comm_done(int); -template void PartitionImpl<2>::comm_done(int); -template void PartitionImpl<3>::comm_done(int); - -template void PartitionImpl<1>::reduce_result(CkReductionMsg*); -template void PartitionImpl<2>::reduce_result(CkReductionMsg*); -template void PartitionImpl<3>::reduce_result(CkReductionMsg*); diff --git a/src/partition_get.cpp b/src/partition_get.cpp deleted file mode 100644 index f878b66..0000000 --- a/src/partition_get.cpp +++ /dev/null @@ -1,71 +0,0 @@ -#include "backend_internal.hpp" - -template -void PartitionImpl::process_get(int epoch) { - - auto nd_idx = this->nd_index(); - - DBG_PRINT("Partition<%d> processing get for epoch %d\n", N, epoch); - auto* req = executor->group->get_get_request(N, epoch); - if (!req) { - CkAbort("No get request found for epoch %d", epoch); - return; - } - int name = req->name; - if (arrays.find(name) == arrays.end()) { - // This partition doesn't own any portion of the array (e.g. the array - // is smaller than the chare grid). Nothing to contribute. - DBG_PRINT("Partition<%d>: array %d not local, skipping get for epoch %d\n", N, name, - epoch); - return; - } - - arrays[name]->copyToHost(); - - { - CTArrayBase* arr = arrays[name]; - int esz = arr->elem_size(); - auto global_region = arr->decomp.chare_region_global(nd_idx); - - std::array global_strides; - global_strides[N - 1] = 1; - for (int d = N - 2; d >= 0; --d) - global_strides[d] = global_strides[d + 1] * (int64_t)arr->global_shape[d + 1]; - - std::array local_strides; - local_strides[N - 1] = 1; - for (int d = N - 2; d >= 0; --d) - local_strides[d] = local_strides[d + 1] * (int64_t)arr->region.size(d + 1); - - int64_t inner_size = arr->region.size(N - 1); - - std::array idx = {}; - while (true) { - int64_t global_offset = 0; - int64_t local_offset = 0; - for (int d = 0; d < N; ++d) { - global_offset += ((int64_t)global_region.start[d] + idx[d]) * global_strides[d]; - local_offset += (int64_t)idx[d] * local_strides[d]; - } - // Send as bytes: offset and size in bytes - int64_t byte_offset = global_offset * esz; - int64_t byte_size = inner_size * esz; - char* base = static_cast(arr->data_ptr()); - dag_proxy[0].gather(epoch, name, byte_offset, byte_size, base + local_offset * esz); - - int d = N - 2; - while (d >= 0) { - if (++idx[d] < arr->region.size(d)) - break; - idx[d] = 0; - --d; - } - if (d < 0) - break; - } - } -} - -template void PartitionImpl<1>::process_get(int); -template void PartitionImpl<2>::process_get(int); -template void PartitionImpl<3>::process_get(int); diff --git a/src/partition_lifecycle.cpp b/src/partition_lifecycle.cpp deleted file mode 100644 index 2b50d19..0000000 --- a/src/partition_lifecycle.cpp +++ /dev/null @@ -1,149 +0,0 @@ -#include "backend_internal.hpp" -#include "dispatch.hpp" - -template -void PartitionImpl::init(typename PartitionTraits::ProxyType proxy, std::array idx, - CProxy_ArrayDAGGroup dag_proxy_) { - thisProxy = proxy; - index = idx; - dag_proxy = dag_proxy_; -#ifdef USE_NVIDIA - cudaStreamCreate(&compute_stream_raw); - cudaStreamCreate(&comm_stream_raw); - compute_exec = Kokkos::Cuda(compute_stream_raw); - comm_exec = Kokkos::Cuda(comm_stream_raw); -#endif - auto* dag_group = static_cast(dag_proxy_.ckLocalBranch()); - executor = new ArrayDAGExecutorND(dag_group, this); - // Skip to the epoch where this partition was first created. - // start_epoch - 1 so that next_epoch() advances to start_epoch. - int se = dag_group->partition_grid[N].start_epoch; - if (se > 0) - executor->epoch = se - 1; - ChareIndex ci; - for (int d = 0; d < N; ++d) - ci.idx[d] = index[d]; - proxy_at(thisProxy, ci).run(); -} - -template -PartitionImpl::~PartitionImpl() { - delete executor; - for (auto& [name, arr] : arrays) - delete arr; - for (auto& [key, bucket] : free_arrays) - for (auto* arr : bucket) - delete arr; -#ifdef USE_NVIDIA - cudaStreamDestroy(compute_stream_raw); - cudaStreamDestroy(comm_stream_raw); -#endif -} - -template -int PartitionImpl::create(ArrayRegion* region, int name, DType dtype, - const ArrayDecomp& decomp) { - - auto nd_idx = this->nd_index(); - auto* dag_group = static_cast(dag_proxy.ckLocalBranch()); - - std::array global_shape{}; - std::array local_sizes{}; - std::array zeros{}; - std::array steps{}; - zeros.fill(0); - - for (int d = 0; d < N; ++d) { - global_shape[d] = region->size(d); - steps[d] = region->step[d]; - } - - ArrayDAGGroup::ArrayMetadata live_meta = {}; - live_meta.ndims = N; - live_meta.global_shape = {0, 0, 0}; - live_meta.offset = {0, 0, 0}; - live_meta.tile = decomp.tile; - for (int d = 0; d < N; ++d) { - live_meta.global_shape[d] = global_shape[d]; - live_meta.offset[d] = decomp.offset[d]; - } - dag_group->live_array_meta[name] = live_meta; - - auto local_region = decomp.chare_region_local(nd_idx); - for (int d = 0; d < N; ++d) { - local_sizes[d] = local_region.size(d); - if (local_sizes[d] <= 0) - return name; - } - - ArrayRegion chare_region(zeros, local_sizes, steps); - auto old_it = arrays.find(name); - if (old_it != arrays.end()) { - delete old_it->second; - arrays.erase(old_it); - } - - arrays[name] = allocate_or_reuse(chare_region, global_shape, name, dtype, decomp); - return name; -} - -template -void PartitionImpl::run() { - // If a DAG is currently executing (not all nodes done), ignore this - // stale run() from the NONE polling loop. The DAG completion callback - // will re-invoke run() when the DAG finishes. - if (executor->dag != nullptr && - executor->dag->num_nodes_done < executor->dag->num_nodes) { - return; - } - -#ifndef NDEBUG - // Print cumulative communication volume at the end of each DAG epoch - if (executor->dag != nullptr && - executor->dag->num_nodes_done >= executor->dag->num_nodes) { - DBG_PRINT("[PE %d] Partition<%d> chare (%d", - CkMyPe(), N, index[0]); - for (int d = 1; d < N; ++d) - CkPrintf(",%d", index[d]); - CkPrintf("): epoch %d done, cumulative comm_bytes_sent=%lld\n", - executor->epoch, (long long)comm_bytes_sent); - } -#endif - - ChareIndex ci; - for (int d = 0; d < N; ++d) - ci.idx[d] = index[d]; - - EpochType type = executor->next_epoch(); - DBG_PRINT("[PE %d] Partition<%d> chare %d: run() epoch=%d type=%s\n", - CkMyPe(), N, index[0], executor->epoch, - type == EpochType::DAG ? "DAG" : type == EpochType::GET ? "GET" : "NONE"); - switch (type) { - case EpochType::DAG: - executor->execute_dag( - CkCallback(PartitionTraits::CkIndexType::run(), proxy_at(thisProxy, ci))); - break; - case EpochType::GET: - process_get(executor->epoch); - proxy_at(thisProxy, ci).run(); - break; - case EpochType::NONE: - break; - } -} - -template void PartitionImpl<1>::init(CProxy_Partition1D, std::array, CProxy_ArrayDAGGroup); -template void PartitionImpl<2>::init(CProxy_Partition2D, std::array, CProxy_ArrayDAGGroup); -template void PartitionImpl<3>::init(CProxy_Partition3D, std::array, CProxy_ArrayDAGGroup); - -template PartitionImpl<1>::~PartitionImpl(); -template PartitionImpl<2>::~PartitionImpl(); -template PartitionImpl<3>::~PartitionImpl(); - -template int PartitionImpl<1>::create(ArrayRegion<1>*, int, DType, const ArrayDecomp<1>&); -template int PartitionImpl<2>::create(ArrayRegion<2>*, int, DType, const ArrayDecomp<2>&); -template int PartitionImpl<3>::create(ArrayRegion<3>*, int, DType, const ArrayDecomp<3>&); - -template void PartitionImpl<1>::run(); -template void PartitionImpl<2>::run(); -template void PartitionImpl<3>::run(); diff --git a/src/server.ci b/src/server.ci new file mode 100644 index 0000000..9f065a5 --- /dev/null +++ b/src/server.ci @@ -0,0 +1,9 @@ +mainmodule server +{ + extern module libaum; + + mainchare Main + { + entry Main(CkArgMsg*); + }; +} diff --git a/src/server.cpp b/src/server.cpp index f02cc87..0f290b6 100644 --- a/src/server.cpp +++ b/src/server.cpp @@ -1,23 +1,33 @@ -#include +#include +#include "server.hpp" +#include "converse.h" +#include "conv-ccs.h" -#include +#include "server.decl.h" -#include "backend.hpp" -CkGroupID create_dag_group() { return CProxy_ArrayDAGGroup::ckNew(); } - -class Main : public CBase_Main { - public: - Main(CkArgMsg* m) { - create_dag_group(); - // Server::initialize is called from ArrayDAGGroup::proxies_ready - // after all PEs have received partition proxies via reduction. +class Main : public CBase_Main +{ +public: + Main(CkArgMsg* msg) + { + Server::initialize(); + register_handlers(); +#ifndef NDEBUG + CkPrintf("Initialization done\n"); +#endif } - Main(CkMigrateMessage* m) : CBase_Main(m) { create_dag_group(); } - - void pup(PUP::er& p) {} + void register_handlers() + { + CcsRegisterHandler("aum_connect", (CmiHandler) Server::connection_handler); + CcsRegisterHandler("aum_disconnect", (CmiHandler) Server::disconnection_handler); + CcsRegisterHandler("aum_operation", (CmiHandler) Server::operation_handler); + CcsRegisterHandler("aum_creation", (CmiHandler) Server::creation_handler); + CcsRegisterHandler("aum_fetch", (CmiHandler) Server::fetch_handler); + CcsRegisterHandler("aum_delete", (CmiHandler) Server::delete_handler); + CcsRegisterHandler("aum_exit", (CmiHandler) Server::exit_server); + } }; -#include "charmtyles.def.h" -#include "charmnumeric.def.h" +#include "server.def.h" diff --git a/src/server.hpp b/src/server.hpp new file mode 100644 index 0000000..1cf6f1c --- /dev/null +++ b/src/server.hpp @@ -0,0 +1,480 @@ +#include +#include +#include +#include +#include +#include +#include "ast.hpp" + + +using aum_name_t = uint64_t; +using aum_array_t = std::variant; +std::unordered_map symbol_table; +std::stack client_ids; + + +aum_array_t calculate(astnode* node, std::vector &metadata); + + +class Server +{ +public: + static void initialize() + { + for (int16_t i = 255; i >= 0; i--) + client_ids.push((uint8_t) i); + } + + inline static void insert(aum_name_t name, aum_array_t arr) + { +#ifndef NDEBUG + CkPrintf("Created array %" PRIu64 " on server\n", name); +#endif + symbol_table.emplace(name, arr); + } + + inline static void remove(aum_name_t name) + { + symbol_table.erase(name); +#ifndef NDEBUG + CkPrintf("Deleted array %" PRIu64 " on server\n", name); +#endif + } + + inline static uint8_t get_client_id() + { + if (client_ids.empty()) + CmiAbort("Too many clients connected to the server"); + uint8_t client_id = client_ids.top(); + client_ids.pop(); + return client_id; + } + + static aum_array_t& lookup(aum_name_t name) + { + auto find = symbol_table.find(name); + if (find == std::end(symbol_table)) + { +#ifndef NDEBUG + CkPrintf("Active symbols: "); + for (auto it: symbol_table) + CkPrintf("%" PRIu64 ", ", it.first); + CkPrintf("\n"); +#endif + CmiAbort("Symbol %i not found", name); + } + return find->second; + } + + static void operation_handler(char* msg) + { + char* cmd = msg + CmiMsgHeaderSizeBytes; + astnode* head = decode(cmd); + std::vector metadata; + calculate(head, metadata); + delete_ast(head); + CcsSendReply(sizeof(uint64_t) * metadata.size(), (void*) &metadata[0]); + } + + static void creation_handler(char* msg) + { + /* First 32 bits are number of dimensions + * Each of the next 64 bits is the size of the array + * in each dimension + */ + // FIXME need a field for dtype + try + { + char* cmd = msg + CmiMsgHeaderSizeBytes; + aum_name_t res_name = extract(cmd); + uint32_t ndim = extract(cmd); + bool has_buf = extract(cmd); + bool has_init = extract(cmd); + switch(ndim) + { + case 0: { + // create scalar + CmiAbort("Not implemented"); + } + case 1: { + // create vector + uint64_t size = extract(cmd); + aum_array_t res; + if (has_buf) + { + double* init_buf = (double*) cmd; + res = aum::vector(size, init_buf); + } + else if (has_init) + { + double init_value = extract(cmd); + res = aum::vector(size, init_value); + } + else + { + res = aum::vector(size, aum::random{}); + } + insert(res_name, res); + break; + } + case 2: { + // create matrix + uint64_t size1 = extract(cmd); + uint64_t size2 = extract(cmd); + aum_array_t res; + if (has_buf) + { + double* init_buf = (double*) cmd; + res = aum::matrix(size1, size2, init_buf); + } + else if (has_init) + { + double init_value = extract(cmd); + res = aum::matrix(size1, size2, init_value); + } + else + { + res = aum::matrix(size1, size2, aum::random{}); + } + insert(res_name, res); + break; + } + default: { + // FIXME is this correctly caught? + CmiAbort("Greater than 2 dimensions not supported"); + } + } + + bool status = true; + CcsSendReply(1, (void*) &status); + } + catch(...) + { + bool status = false; + CcsSendReply(1, (void*) &status); + } + } + + static void fetch_handler(char* msg) + { + char* cmd = msg + CmiMsgHeaderSizeBytes; + aum_name_t name = extract(cmd); + aum_array_t& arr = lookup(name); + void* reply = nullptr; + int reply_size = 0; + std::visit( + [&](auto& x) { + using T = std::decay_t; + if constexpr(std::is_same_v) + { + double value = x.get(); + reply = (void*) &value; + reply_size += 8; + CcsSendReply(reply_size, reply); + } + else + CmiAbort("Operation not implemented"); + }, arr); + } + + static void delete_handler(char* msg) + { + char* cmd = msg + CmiMsgHeaderSizeBytes; + aum_name_t name = extract(cmd); + remove(name); + } + + static void connection_handler(char* msg) + { + uint8_t client_id = get_client_id(); + CcsSendReply(1, (void*) &client_id); + } + + static void disconnection_handler(char* msg) + { + char* cmd = msg + CmiMsgHeaderSizeBytes; + uint8_t client_id = extract(cmd); + client_ids.push(client_id); +#ifndef NDEBUG + CkPrintf("Disconnected %" PRIu8 " from server\n", client_id); +#endif + } + + inline static void exit_server(char* msg) + { + CkExit(); + } +}; + +void extract_metadata(aum_name_t name, aum_array_t &res, std::vector &metadata) +{ + metadata.push_back(name); + std::visit( + [&](auto& x) { + using T = std::decay_t; + if constexpr (std::is_same_v) + { + metadata.push_back(x.rows()); + metadata.push_back(x.cols()); + } + else if constexpr (std::is_same_v) + metadata.push_back(x.size()); + else if constexpr (!std::is_same_v) + CmiAbort("Array type not recognized"); + }, res); +} + +aum_array_t calculate(astnode* node, std::vector &metadata) +{ + switch (node->oper) + { + case operation::noop: { + if (node->is_scalar) + return *reinterpret_cast(&(node->name)); + else + return Server::lookup(node->name); + } + case operation::add: { + aum_array_t s1 = calculate(node->operands[0], metadata); + aum_array_t s2 = calculate(node->operands[1], metadata); + aum_array_t res; + + std::visit( + [&](auto& x, auto& y) { + using T = std::decay_t; + using V = std::decay_t; + if constexpr(std::is_same_v && !std::is_same_v) + { + if (!node->operands[0]->store) + res = x.add_inplace(y); + else if (!node->operands[1]->store) + res = y.add_inplace(x); + else + res = x + y; + } + else + CmiAbort("Operation not permitted"); + }, s1, s2); + + if (node->store) + { + Server::insert(node->name, res); + extract_metadata(node->name, res, metadata); + } + return res; + } + case operation::sub: { + aum_array_t s1 = calculate(node->operands[0], metadata); + aum_array_t s2 = calculate(node->operands[1], metadata); + aum_array_t res; + + std::visit( + [&](auto& x, auto& y) { + using T = std::decay_t; + using V = std::decay_t; + if constexpr(std::is_same_v && !std::is_same_v) + { + if (!node->operands[0]->store) + res = x.sub_inplace_1(y); + else if (!node->operands[1]->store) + res = y.sub_inplace_2(x); + else + res = x - y; + } + else + CmiAbort("Operation not permitted"); + }, s1, s2); + + if (node->store) + { + Server::insert(node->name, res); + extract_metadata(node->name, res, metadata); + } + return res; + } + case operation::mul: { + aum_array_t s1 = calculate(node->operands[0], metadata); + aum_array_t s2 = calculate(node->operands[1], metadata); + aum_array_t res; + + std::visit( + [&](auto& x, auto& y) { + using T = std::decay_t; + using V = std::decay_t; + if constexpr(std::is_same_v || + std::is_same_v) + { + // FIXME implement inplace multiply + // if (!node->operands[1]->store) + // res = y.mul_inplace(x); + // else + res = x * y; + } + else if constexpr(std::is_same_v || + std::is_same_v) + { + // if (!node->operands[0]->store) + // res = x.mul_inplace(y); + // else + res = y * x; + } + else + CmiAbort("Operation not permitted"); + }, s1, s2); + + if (node->store) + { + Server::insert(node->name, res); + extract_metadata(node->name, res, metadata); + } + return res; + } + case operation::div: { + aum_array_t s1 = calculate(node->operands[0], metadata); + aum_array_t s2 = calculate(node->operands[1], metadata); + aum_array_t res; + + std::visit( + [&](auto& x, auto& y) { + using T = std::decay_t; + using V = std::decay_t; + if constexpr((std::is_same_v || + std::is_same_v) && + (std::is_same_v || + std::is_same_v)) + { + res = x / y; + } + else + CmiAbort("Operation not permitted"); + }, s1, s2); + + if (node->store) + { + Server::insert(node->name, res); + extract_metadata(node->name, res, metadata); + } + return res; + } + case operation::matmul: { + aum_array_t s1 = calculate(node->operands[0], metadata); + aum_array_t s2 = calculate(node->operands[1], metadata); + aum_array_t res; + + std::visit( + [&](auto& x, auto& y) { + using T = std::decay_t; + using V = std::decay_t; + if constexpr (std::is_same_v && + std::is_same_v) + CmiAbort("Matrix multiplication not yet implemented"); + else if constexpr (std::is_same_v && + std::is_same_v) + { + if (!node->operands[0]->store) + res = aum::inplace_dot_1(x, y); + else if (!node->operands[1]->store) + res = aum::inplace_dot_2(x, y); + else + res = aum::dot(x, y); + } + else if constexpr ((std::is_same_v || + std::is_same_v) && + std::is_same_v) + { + // TODO implement inplace dot + res = aum::dot(x, y); + } + else + CmiAbort("Operation not permitted"); + }, s1, s2); + + if (node->store) + { + Server::insert(node->name, res); + extract_metadata(node->name, res, metadata); + } + return res; + } + case operation::copy: { + aum_array_t s1 = calculate(node->operands[0], metadata); + aum_array_t res; + + std::visit( + [&](auto& x) { + using T = std::decay_t; + if constexpr(std::is_same_v || + std::is_same_v) + res = aum::copy(x); + else + CmiAbort("Matrix copy not implemented"); + }, s1); + + if (node->store) + { + Server::insert(node->name, res); + extract_metadata(node->name, res, metadata); + } + return res; + } + case operation::axpy: { + aum_array_t s1 = calculate(node->operands[0], metadata); + aum_array_t s2 = calculate(node->operands[1], metadata); + aum_array_t s3 = calculate(node->operands[2], metadata); + aum_array_t res; + + std::visit( + [&](auto& a, auto& x, auto& y) { + using S = std::decay_t; + using T = std::decay_t; + using V = std::decay_t; + if constexpr(std::is_same_v && + std::is_same_v && + std::is_same_v) + res = aum::blas::axpy(a, x, y); + else + CmiAbort("Operation not permitted"); + }, s1, s2, s3); + + if (node->store) + { + Server::insert(node->name, res); + extract_metadata(node->name, res, metadata); + } + return res; + } + case operation::axpy_multiplier: { + aum_array_t s1 = calculate(node->operands[0], metadata); + aum_array_t s2 = calculate(node->operands[1], metadata); + aum_array_t s3 = calculate(node->operands[2], metadata); + aum_array_t multiplier = calculate(node->operands[3], metadata); + aum_array_t res; + + std::visit( + [&](auto& multiplier, auto& a, auto& x, auto& y) { + using S = std::decay_t; + using T = std::decay_t; + using V = std::decay_t; + using M = std::decay_t; + if constexpr(std::is_same_v && + std::is_same_v && + std::is_same_v && + std::is_same_v) + res = aum::blas::axpy(multiplier, a, x, y); + else + CmiAbort("Operation not permitted"); + }, multiplier, s1, s2, s3); + + if (node->store) + { + Server::insert(node->name, res); + extract_metadata(node->name, res, metadata); + } + return res; + } + default: { + CmiAbort("Operation not implemented"); + } + } +} + diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt deleted file mode 100644 index 93e0a0e..0000000 --- a/tests/CMakeLists.txt +++ /dev/null @@ -1,68 +0,0 @@ -cmake_minimum_required(VERSION 3.20) -project(test_jit LANGUAGES CXX) - -set(CMAKE_CXX_STANDARD 17) -set(CMAKE_CXX_STANDARD_REQUIRED ON) - -if(DEFINED ENV{CHARM_HOME} AND NOT DEFINED CHARM_HOME) - set(CHARM_HOME "$ENV{CHARM_HOME}") -endif() -if(NOT DEFINED CHARM_HOME) - message(FATAL_ERROR "CHARM_HOME is not set. Point it to your Charm++ installation.") -endif() - -# ---------- Find MLIR (and LLVM transitively) ---------- -find_package(MLIR REQUIRED CONFIG) -find_package(LLVM REQUIRED CONFIG) - -message(STATUS "Using MLIRConfig.cmake in: ${MLIR_DIR}") -message(STATUS "Using LLVMConfig.cmake in: ${LLVM_DIR}") - -list(APPEND CMAKE_MODULE_PATH "${MLIR_CMAKE_DIR}") -list(APPEND CMAKE_MODULE_PATH "${LLVM_CMAKE_DIR}") - -include(TableGen) -include(AddLLVM) -include(AddMLIR) - -include_directories(${LLVM_INCLUDE_DIRS}) -include_directories(${MLIR_INCLUDE_DIRS}) -include_directories(${CHARM_HOME}/include) - -# ---------- Charmtyles / example includes ---------- -set(CHARMTYLES_HOME "${CMAKE_CURRENT_SOURCE_DIR}/../../..") -include_directories(${CHARMTYLES_HOME}/include) -include_directories(${CHARMTYLES_HOME}/example/charmnumeric/src) - -# ---------- Build the test ---------- -add_executable(test_jit test_jit.cpp) -add_executable(test_compute_decompositions test_compute_decompositions.cpp) -target_compile_definitions(test_compute_decompositions PRIVATE - CT_MIN_TILE_1D=64 - CT_MIN_TILE_2D=64 - CT_MIN_TILE_3D=64 -) - -# Use MLIR's cmake helper to resolve ALL transitive deps automatically -get_property(mlir_dialect_libs GLOBAL PROPERTY MLIR_DIALECT_LIBS) -get_property(mlir_conversion_libs GLOBAL PROPERTY MLIR_CONVERSION_LIBS) - -target_link_libraries(test_jit - PRIVATE - ${mlir_dialect_libs} - ${mlir_conversion_libs} - MLIRPass - MLIRTransforms - MLIRTargetLLVMIRExport - MLIRLLVMToLLVMIRTranslation - MLIRNVVMToLLVMIRTranslation - MLIRROCDLToLLVMIRTranslation - MLIRIR - MLIRSupport - MLIRParser - LLVMCore - LLVMSupport -) - -enable_testing() -add_test(NAME test_compute_decompositions COMMAND test_compute_decompositions) diff --git a/tests/Makefile b/tests/Makefile deleted file mode 100644 index 455191e..0000000 --- a/tests/Makefile +++ /dev/null @@ -1,61 +0,0 @@ -# Makefile for standalone JIT test -# Requires: CHARMTYLES_HOME and MLIR/LLVM installed -# -# Usage: -# make test_jit -# ./test_jit - -CXX ?= clang++ -CXXFLAGS = -std=c++17 -O2 -g -DNDEBUG - -# Include paths: charmtyles headers + example src headers (jit.hpp, opcodes.hpp, array_region.hpp) -# Auto-derive CHARMTYLES_HOME from this Makefile's location -CHARMTYLES_HOME ?= $(realpath $(dir $(lastword $(MAKEFILE_LIST)))../../..) - -# Include paths: charmtyles headers + example src headers (jit.hpp, opcodes.hpp, array_region.hpp) -INCLUDES = -I$(CHARMTYLES_HOME)/include \ - -I$(CHARMTYLES_HOME)/example/charmnumeric/src - -# MLIR/LLVM flags -MLIR_CXXFLAGS := $(shell llvm-config --cxxflags 2>/dev/null) -MLIR_LDFLAGS := $(shell llvm-config --ldflags --libs all --system-libs 2>/dev/null) - -# MLIR libraries (adjust to your MLIR build) -MLIR_LIBS = -lMLIRLinalgTransforms -lMLIRLinalgDialect -lMLIRLinalgUtils \ - -lMLIRAffineTransforms -lMLIRAffineDialect -lMLIRAffineAnalysis -lMLIRAffineUtils \ - -lMLIRGPUDialect -lMLIRGPUTransforms \ - -lMLIRArithDialect -lMLIRArithUtils \ - -lMLIRMemRefDialect -lMLIRMemRefTransforms -lMLIRMemRefUtils \ - -lMLIRFuncDialect -lMLIRFuncTransforms \ - -lMLIRTensorDialect -lMLIRTensorUtils -lMLIRTensorTransforms \ - -lMLIRMathDialect \ - -lMLIRSCFDialect -lMLIRSCFTransforms -lMLIRSCFUtils \ - -lMLIRBufferizationDialect -lMLIRBufferizationTransforms \ - -lMLIRVectorDialect -lMLIRVectorTransforms -lMLIRVectorUtils \ - -lMLIRPass -lMLIRTransforms -lMLIRTransformUtils \ - -lMLIRSPIRVDialect -lMLIRSPIRVSerialization \ - -lMLIRLLVMDialect -lMLIRNVVMDialect -lMLIRROCDLDialect \ - -lMLIRLLVMToLLVMIRTranslation -lMLIRNVVMToLLVMIRTranslation \ - -lMLIRROCDLToLLVMIRTranslation \ - -lMLIRTargetLLVMIRExport \ - -lMLIRAnalysis -lMLIRSideEffectInterfaces -lMLIRControlFlowInterfaces \ - -lMLIRLoopLikeInterface -lMLIRViewLikeInterface \ - -lMLIRDestinationStyleOpInterface -lMLIRInferTypeOpInterface \ - -lMLIRDialectUtils -lMLIRRewrite \ - -lMLIRIR -lMLIRSupport -lMLIRParser - -all: test_jit - -test_jit: test_jit.cpp ../src/jit.hpp ../src/opcodes.hpp ../src/array_region.hpp - $(CXX) $(CXXFLAGS) $(INCLUDES) $(MLIR_CXXFLAGS) \ - test_jit.cpp \ - $(MLIR_LDFLAGS) $(MLIR_LIBS) \ - -o $@ - -run: test_jit - ./test_jit - -clean: - rm -f test_jit - -.PHONY: all run clean diff --git a/tests/conftest.py b/tests/conftest.py deleted file mode 100644 index 9ef28e6..0000000 --- a/tests/conftest.py +++ /dev/null @@ -1,186 +0,0 @@ -"""Shared pytest fixtures for backend-driven charmnumeric tests.""" - -from __future__ import annotations - -from pathlib import Path -import gc -import os -import sys - -import pytest - - -HERE = Path(__file__).resolve().parent -CHARMNUMERIC_ROOT = HERE.parent -REPO_ROOT = CHARMNUMERIC_ROOT.parent.parent -DEFAULT_SERVER_BINARY = CHARMNUMERIC_ROOT / "src" / "build" / "server.out" - - -for build_dir in ( - REPO_ROOT / "build" / "lib", - CHARMNUMERIC_ROOT / "build" / "lib", -): - build_dir_str = str(build_dir) - while build_dir_str in sys.path: - sys.path.remove(build_dir_str) - -repo_root_str = str(REPO_ROOT) -while repo_root_str in sys.path: - sys.path.remove(repo_root_str) -sys.path.insert(0, repo_root_str) - -charmnumeric_root_str = str(CHARMNUMERIC_ROOT) -while charmnumeric_root_str in sys.path: - sys.path.remove(charmnumeric_root_str) -sys.path.insert(1, charmnumeric_root_str) - - -def pytest_addoption(parser): - group = parser.getgroup("charmnumeric") - group.addoption( - "--charmnumeric-server", - action="store", - default=None, - help="Path to a built charmnumeric server.out binary.", - ) - group.addoption( - "--charmnumeric-host", - action="store", - default=None, - help="Reuse an already running charmnumeric backend at this host.", - ) - group.addoption( - "--charmnumeric-port", - action="store", - type=int, - default=None, - help="Port for the charmnumeric backend connection.", - ) - group.addoption( - "--charmnumeric-odf", - action="store", - type=int, - default=None, - help="Objects-per-PE factor used during backend connect.", - ) - group.addoption( - "--charmnumeric-pes", - action="store", - type=int, - default=None, - help="Number of PEs to launch for the local charmnumeric backend.", - ) - group.addoption( - "--charmnumeric-charmrun", - action="store", - default=None, - help="Optional path to charmrun for launching the local backend.", - ) - group.addoption( - "--charmnumeric-startup-timeout", - action="store", - type=float, - default=None, - help="Seconds to wait for a launched charmnumeric backend to accept connections.", - ) - - -def _get_str_option(pytestconfig, option_name, env_name): - value = pytestconfig.getoption(option_name) - if value: - return value - return os.environ.get(env_name) - - -def _get_int_option(pytestconfig, option_name, env_name, default): - value = pytestconfig.getoption(option_name) - if value is not None: - return value - env_value = os.environ.get(env_name) - if env_value is not None: - return int(env_value) - return default - - -def _get_float_option(pytestconfig, option_name, env_name, default): - value = pytestconfig.getoption(option_name) - if value is not None: - return value - env_value = os.environ.get(env_name) - if env_value is not None: - return float(env_value) - return default - - -def _resolve_server_binary(pytestconfig): - value = _get_str_option(pytestconfig, "charmnumeric_server", "CHARMNUMERIC_SERVER") - if value: - path = Path(value).expanduser().resolve() - if not path.is_file(): - pytest.exit(f"charmnumeric server not found: {path}") - return path - - if DEFAULT_SERVER_BINARY.is_file(): - return DEFAULT_SERVER_BINARY - return None - - -@pytest.fixture(autouse=True) -def _collect_garbage(): - gc.collect() - yield - gc.collect() - - -@pytest.fixture(autouse=True) -def _reset_frontend_state(): - from charmtyles.core import reset_frontend_state - - reset_frontend_state() - yield - reset_frontend_state() - - -@pytest.fixture(scope="session") -def interface(pytestconfig, tmp_path_factory): - from charmnumeric.interface import LocalCluster, CharmNumericInterface - - server_binary = _resolve_server_binary(pytestconfig) - host = _get_str_option(pytestconfig, "charmnumeric_host", "CHARMNUMERIC_HOST") - port = _get_int_option(pytestconfig, "charmnumeric_port", "CHARMNUMERIC_PORT", 1234) - odf = _get_int_option(pytestconfig, "charmnumeric_odf", "CHARMNUMERIC_ODF", 4) - pes = _get_int_option(pytestconfig, "charmnumeric_pes", "CHARMNUMERIC_PES", 1) - charmrun = _get_str_option(pytestconfig, "charmnumeric_charmrun", "CHARMNUMERIC_CHARMRUN") - startup_timeout = _get_float_option( - pytestconfig, - "charmnumeric_startup_timeout", - "CHARMNUMERIC_STARTUP_TIMEOUT", - 20.0, - ) - - if server_binary is not None: - cluster = LocalCluster( - server_binary=str(server_binary), - server_port=port, - odf=odf, - max_pes=pes, - charmrun=charmrun, - workdir=tmp_path_factory.mktemp("charmnumeric-backend"), - startup_timeout=startup_timeout, - ) - try: - yield cluster - finally: - cluster.close() - return - - if host: - iface = CharmNumericInterface() - iface.connect(host, port, odf) - yield iface - return - - pytest.skip( - "Build example/charmnumeric/src/build/server.out first, or pass " - "--charmnumeric-server / CHARMNUMERIC_SERVER." - ) diff --git a/tests/test_backend_integration.py b/tests/test_backend_integration.py deleted file mode 100644 index 032d728..0000000 --- a/tests/test_backend_integration.py +++ /dev/null @@ -1,126 +0,0 @@ -from __future__ import annotations - -import numpy as np -import pytest - -from charmnumeric.operations import diag -from charmnumeric.charmnumeric import arange, create_array - - -pytestmark = pytest.mark.integration - - -def test_elementwise_expression_roundtrip(interface): - a = create_array((32,), dtype=np.float32) - b = create_array((32,), dtype=np.float32) - - a[:] = 1.5 - b[:] = -0.5 - - got = (2.0 * (a + b + 3.0)).get(interface) - want = np.full((32,), 8.0, dtype=np.float32) - - np.testing.assert_allclose(got, want) - - -@pytest.mark.parametrize("shape", [(8,), (4, 4), (4, 4, 4)]) -def test_elementwise_add_roundtrip_all_ranks(interface, shape): - lhs = create_array(shape, dtype=np.float32) - rhs = create_array(shape, dtype=np.float32) - - lhs[...] = 1.0 - rhs[...] = 1.0 - - got = (lhs + rhs).get(interface) - want = np.full(shape, 2.0, dtype=np.float32) - - np.testing.assert_allclose(got, want) - - -def test_shifted_slice_expression_get_matches_numpy(interface): - src = arange(512, dtype=np.float32) - - got = (src[100:300] + src[150:350]).get(interface) - host = np.arange(512, dtype=np.float32) - want = host[100:300] + host[150:350] - - np.testing.assert_allclose(got, want) - - -def test_strided_slice_assignment_roundtrip(interface): - n = 17 - coarse_n = (n - 1) // 2 + 1 - - coarse = create_array((coarse_n, coarse_n), dtype=np.float64) - fine = create_array((n, n), dtype=np.float64) - - coarse[:, :] = 0.0 - fine[:, :] = 1.0 - coarse[1:-1, 1:-1] = fine[2:-2:2, 2:-2:2] - - got = coarse.get(interface) - want = np.zeros((coarse_n, coarse_n), dtype=np.float64) - want[1:-1, 1:-1] = np.ones((n, n), dtype=np.float64)[2:-2:2, 2:-2:2] - - np.testing.assert_array_equal(got, want) - - -def test_jacobi3d_single_step_matches_numpy(interface): - u = create_array((4, 4, 4), dtype=np.float32) - - u[0, :, :] = 1.0 - u[-1, :, :] = 1.0 - u[:, 0, :] = 1.0 - u[:, -1, :] = 1.0 - u[:, :, 0] = 1.0 - u[:, :, -1] = 1.0 - - u[1:-1, 1:-1, 1:-1] = (1.0 / 6.0) * ( - u[:-2, 1:-1, 1:-1] - + u[2:, 1:-1, 1:-1] - + u[1:-1, :-2, 1:-1] - + u[1:-1, 2:, 1:-1] - + u[1:-1, 1:-1, :-2] - + u[1:-1, 1:-1, 2:] - ) - - got = u.get(interface) - - want = np.zeros((4, 4, 4), dtype=np.float32) - want[0, :, :] = 1.0 - want[-1, :, :] = 1.0 - want[:, 0, :] = 1.0 - want[:, -1, :] = 1.0 - want[:, :, 0] = 1.0 - want[:, :, -1] = 1.0 - want[1:-1, 1:-1, 1:-1] = (1.0 / 6.0) * ( - want[:-2, 1:-1, 1:-1] - + want[2:, 1:-1, 1:-1] - + want[1:-1, :-2, 1:-1] - + want[1:-1, 2:, 1:-1] - + want[1:-1, 1:-1, :-2] - + want[1:-1, 1:-1, 2:] - ) - - np.testing.assert_allclose(got, want) - - -@pytest.mark.parametrize("k", [0, 1, -1]) -def test_diag_vector_matches_numpy(interface, k): - vector = create_array((8,), dtype=np.float64) - vector[:] = 2.0 - - got = diag(vector, k=k).get(interface) - want = np.diag(np.full(8, 2.0, dtype=np.float64), k=k) - - np.testing.assert_allclose(got, want) - - -def test_diag_matrix_extract_matches_numpy(interface): - matrix = create_array((6, 10), dtype=np.float64) - matrix[:, :] = 3.0 - - got = diag(matrix, k=2).get(interface) - want = np.diag(np.full((6, 10), 3.0, dtype=np.float64), k=2) - - np.testing.assert_allclose(got, want) diff --git a/tests/test_compute_decompositions.cpp b/tests/test_compute_decompositions.cpp deleted file mode 100644 index 445133c..0000000 --- a/tests/test_compute_decompositions.cpp +++ /dev/null @@ -1,416 +0,0 @@ -#include -#include -#include -#include -#include -#include -#include -#include - -#define DBG_PRINT(...) -#include "decomposition_solver.hpp" - -namespace { - -extern "C" void CmiAbort(const char*, ...) { - std::abort(); -} - -struct TestArrayMetadata { - int ndims = 0; - std::array global_shape = {0, 0, 0}; - std::array offset = {0, 0, 0}; - int tile = 0; - bool decomp_final = false; -}; - -using MetaMap = std::unordered_map; - -[[noreturn]] void fail(const std::string& message) { - throw std::runtime_error(message); -} - -void expect_eq(const std::string& label, int got, int want) { - if (got == want) - return; - std::ostringstream oss; - oss << label << ": expected " << want << ", got " << got; - fail(oss.str()); -} - -void expect_true(const std::string& label, bool value) { - if (!value) - fail(label); -} - -void expect_offset(const std::string& label, const std::array& got, - const std::array& want) { - if (got == want) - return; - std::ostringstream oss; - oss << label << ": expected (" << want[0] << ", " << want[1] << ", " << want[2] - << "), got (" << got[0] << ", " << got[1] << ", " << got[2] << ")"; - fail(oss.str()); -} - -ASTNode* make_leaf(int name, int ndims = 1) { - return new ASTNode(name, static_cast(Opcode::NOOP), - /*is_temp=*/false, /*is_scalar=*/false, - /*scalar=*/0.0, ndims); -} - -ASTNode* make_copy(int result_name, ASTNode* src, Region* src_region = nullptr, - bool is_temp = false, int ndims = 1) { - auto* node = new ASTNode(result_name, static_cast(Opcode::COPY), - is_temp, false, 0.0, ndims); - node->add_operand(src); - node->operand_regions.push_back(src_region); - return node; -} - -ASTNode* make_add(int result_name, ASTNode* lhs, Region* lhs_region, - ASTNode* rhs, Region* rhs_region, bool is_temp = false, int ndims = 1) { - auto* node = new ASTNode(result_name, static_cast(Opcode::ADD), - is_temp, false, 0.0, ndims); - node->add_operand(lhs); - node->operand_regions.push_back(lhs_region); - node->add_operand(rhs); - node->operand_regions.push_back(rhs_region); - return node; -} - -ASTNode* make_set_region(int result_name, ASTNode* dst, Region* dst_region, ASTNode* src, - int ndims = 1) { - auto* node = new ASTNode(result_name, static_cast(Opcode::SET_REGION), - /*is_temp=*/false, /*is_scalar=*/false, - /*scalar=*/0.0, ndims); - node->add_operand(dst); - node->add_operand(src); - node->region = dst_region; - return node; -} - -template -ArrayRegion* make_region(std::array start, std::array stop, - std::array step) { - return new ArrayRegion(start, stop, step); -} - -DAG* make_single_node_dag(std::vector roots) { - auto* ast = new AST(std::move(roots)); - auto* node = new DAGNode(0, ast); - auto* dag = new DAG({node}); - dag->num_nodes = 1; - return dag; -} - -TestArrayMetadata make_meta(int ndims, int size0, int offset0, int tile, bool final) { - TestArrayMetadata meta; - meta.ndims = ndims; - meta.global_shape = {size0, 0, 0}; - meta.offset = {offset0, 0, 0}; - meta.tile = tile; - meta.decomp_final = final; - return meta; -} - -TestArrayMetadata make_meta(int ndims, std::array shape, std::array offset, - int tile, bool final) { - TestArrayMetadata meta; - meta.ndims = ndims; - meta.global_shape = shape; - meta.offset = offset; - meta.tile = tile; - meta.decomp_final = final; - return meta; -} - -void test_shifted_copy_uses_absolute_offset() { - std::cout << "=== test_shifted_copy_uses_absolute_offset ===" << std::endl; - - constexpr int A = 1; - constexpr int B = 2; - - MetaMap meta; - meta[A] = make_meta(/*ndims=*/1, /*size0=*/512, /*offset0=*/0, /*tile=*/64, /*final=*/true); - meta[B] = make_meta(/*ndims=*/1, /*size0=*/200, /*offset0=*/0, /*tile=*/0, /*final=*/false); - - ASTNode* copy = make_copy(B, make_leaf(A), - make_region<1>({150}, {350}, {1}), - /*is_temp=*/false, /*ndims=*/1); - DAG* dag = make_single_node_dag({copy}); - - decomposition_solver::compute_decompositions(meta, dag, /*odf=*/4, /*num_pes=*/1); - - expect_eq("B.tile", meta.at(B).tile, 64); - expect_eq("B.offset[0]", meta.at(B).offset[0], 150); - expect_eq("B.phase", meta.at(B).offset[0] % meta.at(B).tile, 22); - expect_true("B.decomp_final", meta.at(B).decomp_final); - - delete dag; -} - -void test_shifted_copy_from_nonzero_source_offset() { - std::cout << "=== test_shifted_copy_from_nonzero_source_offset ===" << std::endl; - - constexpr int A = 1; - constexpr int B = 2; - - MetaMap meta; - meta[A] = make_meta(/*ndims=*/1, /*size0=*/512, /*offset0=*/132, /*tile=*/64, /*final=*/true); - meta[B] = make_meta(/*ndims=*/1, /*size0=*/100, /*offset0=*/0, /*tile=*/0, /*final=*/false); - - ASTNode* copy = make_copy(B, make_leaf(A), - make_region<1>({70}, {170}, {1}), - /*is_temp=*/false, /*ndims=*/1); - DAG* dag = make_single_node_dag({copy}); - - decomposition_solver::compute_decompositions(meta, dag, /*odf=*/4, /*num_pes=*/1); - - expect_eq("B.tile", meta.at(B).tile, 64); - expect_eq("B.offset[0]", meta.at(B).offset[0], 202); - expect_eq("B.phase", meta.at(B).offset[0] % meta.at(B).tile, 10); - - delete dag; -} - -void test_shifted_expression_temp_uses_absolute_representative() { - std::cout << "=== test_shifted_expression_temp_uses_absolute_representative ===" - << std::endl; - - constexpr int A = 1; - constexpr int TMP = 10; - constexpr int SET = 11; - - MetaMap meta; - meta[A] = make_meta(/*ndims=*/1, /*size0=*/512, /*offset0=*/0, /*tile=*/64, /*final=*/true); - meta[TMP] = make_meta(/*ndims=*/1, /*size0=*/200, /*offset0=*/0, /*tile=*/0, /*final=*/false); - - ASTNode* tmp = make_add(TMP, - make_leaf(A), make_region<1>({100}, {300}, {1}), - make_leaf(A), make_region<1>({150}, {350}, {1}), - /*is_temp=*/true, /*ndims=*/1); - ASTNode* set_region = - make_set_region(SET, make_leaf(A), make_region<1>({200}, {400}, {1}), make_leaf(TMP)); - DAG* dag = make_single_node_dag({tmp, set_region}); - - decomposition_solver::compute_decompositions(meta, dag, /*odf=*/4, /*num_pes=*/1); - - expect_eq("TMP.tile", meta.at(TMP).tile, 64); - expect_eq("TMP.offset[0]", meta.at(TMP).offset[0], 150); - expect_eq("TMP.phase", meta.at(TMP).offset[0] % meta.at(TMP).tile, 22); - expect_eq("A.offset[0]", meta.at(A).offset[0], 0); - - delete dag; -} - -void test_jacobi1d_temp_offset() { - std::cout << "=== test_jacobi1d_temp_offset ===" << std::endl; - - constexpr int U = 1; - constexpr int F = 2; - constexpr int TMP_STENCIL = 10; - constexpr int TMP_WEIGHTED = 11; - constexpr int SET = 12; - - MetaMap meta; - meta[U] = make_meta(/*ndims=*/1, /*size0=*/129, /*offset0=*/0, /*tile=*/64, /*final=*/true); - meta[F] = make_meta(/*ndims=*/1, /*size0=*/129, /*offset0=*/0, /*tile=*/64, /*final=*/true); - meta[TMP_STENCIL] = - make_meta(/*ndims=*/1, /*size0=*/127, /*offset0=*/0, /*tile=*/64, /*final=*/false); - meta[TMP_WEIGHTED] = - make_meta(/*ndims=*/1, /*size0=*/127, /*offset0=*/0, /*tile=*/64, /*final=*/false); - - ASTNode* tmp_stencil = - new ASTNode(TMP_STENCIL, static_cast(Opcode::ADD), - /*is_temp=*/true, /*is_scalar=*/false, - /*scalar=*/0.0, /*ndims=*/1); - tmp_stencil->add_operand(make_leaf(U)); - tmp_stencil->operand_regions.push_back(make_region<1>({0}, {127}, {1})); - tmp_stencil->add_operand(make_leaf(U)); - tmp_stencil->operand_regions.push_back(make_region<1>({2}, {129}, {1})); - tmp_stencil->add_operand(make_leaf(F)); - tmp_stencil->operand_regions.push_back(make_region<1>({1}, {128}, {1})); - - ASTNode* tmp_weighted = - new ASTNode(TMP_WEIGHTED, static_cast(Opcode::ADD), - /*is_temp=*/true, /*is_scalar=*/false, - /*scalar=*/0.0, /*ndims=*/1); - tmp_weighted->add_operand(make_leaf(U)); - tmp_weighted->operand_regions.push_back(make_region<1>({1}, {128}, {1})); - tmp_weighted->add_operand(make_leaf(TMP_STENCIL)); - tmp_weighted->operand_regions.push_back(nullptr); - - ASTNode* set_region = - make_set_region(SET, make_leaf(U), make_region<1>({1}, {128}, {1}), - make_leaf(TMP_WEIGHTED)); - DAG* dag = make_single_node_dag({tmp_stencil, tmp_weighted, set_region}); - - decomposition_solver::compute_decompositions(meta, dag, /*odf=*/4, /*num_pes=*/1); - - expect_eq("TMP_STENCIL.tile", meta.at(TMP_STENCIL).tile, 64); - expect_eq("TMP_WEIGHTED.tile", meta.at(TMP_WEIGHTED).tile, 64); - expect_eq("TMP_STENCIL.offset[0]", meta.at(TMP_STENCIL).offset[0], 1); - expect_eq("TMP_WEIGHTED.offset[0]", meta.at(TMP_WEIGHTED).offset[0], 1); - - delete dag; -} - -void test_jacobi2d_temp_offset() { - std::cout << "=== test_jacobi2d_temp_offset ===" << std::endl; - - constexpr int U = 1; - constexpr int F = 2; - constexpr int TMP_STENCIL = 10; - constexpr int TMP_WEIGHTED = 11; - constexpr int SET = 12; - - MetaMap meta; - meta[U] = make_meta(/*ndims=*/2, /*shape=*/{129, 129, 0}, /*offset=*/{0, 0, 0}, - /*tile=*/64, /*final=*/true); - meta[F] = make_meta(/*ndims=*/2, /*shape=*/{129, 129, 0}, /*offset=*/{0, 0, 0}, - /*tile=*/64, /*final=*/true); - meta[TMP_STENCIL] = make_meta(/*ndims=*/2, /*shape=*/{127, 127, 0}, /*offset=*/{0, 0, 0}, - /*tile=*/64, /*final=*/false); - meta[TMP_WEIGHTED] = make_meta(/*ndims=*/2, /*shape=*/{127, 127, 0}, /*offset=*/{0, 0, 0}, - /*tile=*/64, /*final=*/false); - - ASTNode* tmp_stencil = - new ASTNode(TMP_STENCIL, static_cast(Opcode::ADD), - /*is_temp=*/true, /*is_scalar=*/false, - /*scalar=*/0.0, /*ndims=*/2); - tmp_stencil->add_operand(make_leaf(U, 2)); - tmp_stencil->operand_regions.push_back(make_region<2>({0, 1}, {127, 128}, {1, 1})); - tmp_stencil->add_operand(make_leaf(U, 2)); - tmp_stencil->operand_regions.push_back(make_region<2>({2, 1}, {129, 128}, {1, 1})); - tmp_stencil->add_operand(make_leaf(U, 2)); - tmp_stencil->operand_regions.push_back(make_region<2>({1, 0}, {128, 127}, {1, 1})); - tmp_stencil->add_operand(make_leaf(U, 2)); - tmp_stencil->operand_regions.push_back(make_region<2>({1, 2}, {128, 129}, {1, 1})); - tmp_stencil->add_operand(make_leaf(F, 2)); - tmp_stencil->operand_regions.push_back(make_region<2>({1, 1}, {128, 128}, {1, 1})); - - ASTNode* tmp_weighted = - new ASTNode(TMP_WEIGHTED, static_cast(Opcode::ADD), - /*is_temp=*/true, /*is_scalar=*/false, - /*scalar=*/0.0, /*ndims=*/2); - tmp_weighted->add_operand(make_leaf(U, 2)); - tmp_weighted->operand_regions.push_back(make_region<2>({1, 1}, {128, 128}, {1, 1})); - tmp_weighted->add_operand(make_leaf(TMP_STENCIL, 2)); - tmp_weighted->operand_regions.push_back(nullptr); - - ASTNode* set_region = - make_set_region(SET, make_leaf(U, 2), make_region<2>({1, 1}, {128, 128}, {1, 1}), - make_leaf(TMP_WEIGHTED, 2), /*ndims=*/2); - DAG* dag = make_single_node_dag({tmp_stencil, tmp_weighted, set_region}); - - decomposition_solver::compute_decompositions(meta, dag, /*odf=*/4, /*num_pes=*/1); - - expect_eq("TMP_STENCIL.tile", meta.at(TMP_STENCIL).tile, 64); - expect_eq("TMP_WEIGHTED.tile", meta.at(TMP_WEIGHTED).tile, 64); - expect_offset("TMP_STENCIL.offset", meta.at(TMP_STENCIL).offset, {1, 1, 0}); - expect_offset("TMP_WEIGHTED.offset", meta.at(TMP_WEIGHTED).offset, {1, 1, 0}); - - delete dag; -} - -void test_jacobi3d_temp_offset() { - std::cout << "=== test_jacobi3d_temp_offset ===" << std::endl; - - constexpr int U = 1; - constexpr int TMP_STENCIL = 10; - constexpr int TMP_WEIGHTED = 11; - constexpr int SET = 12; - - MetaMap meta; - meta[U] = make_meta(/*ndims=*/3, /*shape=*/{65, 65, 65}, /*offset=*/{0, 0, 0}, - /*tile=*/64, /*final=*/true); - meta[TMP_STENCIL] = make_meta(/*ndims=*/3, /*shape=*/{63, 63, 63}, /*offset=*/{0, 0, 0}, - /*tile=*/64, /*final=*/false); - meta[TMP_WEIGHTED] = make_meta(/*ndims=*/3, /*shape=*/{63, 63, 63}, /*offset=*/{0, 0, 0}, - /*tile=*/64, /*final=*/false); - - ASTNode* tmp_stencil = - new ASTNode(TMP_STENCIL, static_cast(Opcode::ADD), - /*is_temp=*/true, /*is_scalar=*/false, - /*scalar=*/0.0, /*ndims=*/3); - tmp_stencil->add_operand(make_leaf(U, 3)); - tmp_stencil->operand_regions.push_back(make_region<3>({0, 1, 1}, {63, 64, 64}, {1, 1, 1})); - tmp_stencil->add_operand(make_leaf(U, 3)); - tmp_stencil->operand_regions.push_back(make_region<3>({2, 1, 1}, {65, 64, 64}, {1, 1, 1})); - tmp_stencil->add_operand(make_leaf(U, 3)); - tmp_stencil->operand_regions.push_back(make_region<3>({1, 0, 1}, {64, 63, 64}, {1, 1, 1})); - tmp_stencil->add_operand(make_leaf(U, 3)); - tmp_stencil->operand_regions.push_back(make_region<3>({1, 2, 1}, {64, 65, 64}, {1, 1, 1})); - tmp_stencil->add_operand(make_leaf(U, 3)); - tmp_stencil->operand_regions.push_back(make_region<3>({1, 1, 0}, {64, 64, 63}, {1, 1, 1})); - tmp_stencil->add_operand(make_leaf(U, 3)); - tmp_stencil->operand_regions.push_back(make_region<3>({1, 1, 2}, {64, 64, 65}, {1, 1, 1})); - - ASTNode* tmp_weighted = - new ASTNode(TMP_WEIGHTED, static_cast(Opcode::ADD), - /*is_temp=*/true, /*is_scalar=*/false, - /*scalar=*/0.0, /*ndims=*/3); - tmp_weighted->add_operand(make_leaf(U, 3)); - tmp_weighted->operand_regions.push_back(make_region<3>({1, 1, 1}, {64, 64, 64}, {1, 1, 1})); - tmp_weighted->add_operand(make_leaf(TMP_STENCIL, 3)); - tmp_weighted->operand_regions.push_back(nullptr); - - ASTNode* set_region = - make_set_region(SET, make_leaf(U, 3), make_region<3>({1, 1, 1}, {64, 64, 64}, {1, 1, 1}), - make_leaf(TMP_WEIGHTED, 3), /*ndims=*/3); - DAG* dag = make_single_node_dag({tmp_stencil, tmp_weighted, set_region}); - - decomposition_solver::compute_decompositions(meta, dag, /*odf=*/4, /*num_pes=*/1); - - expect_eq("TMP_STENCIL.tile", meta.at(TMP_STENCIL).tile, 64); - expect_eq("TMP_WEIGHTED.tile", meta.at(TMP_WEIGHTED).tile, 64); - expect_offset("TMP_STENCIL.offset", meta.at(TMP_STENCIL).offset, {1, 1, 1}); - expect_offset("TMP_WEIGHTED.offset", meta.at(TMP_WEIGHTED).offset, {1, 1, 1}); - - delete dag; -} - -void test_strided_copy_reduces_tile_and_sets_offset() { - std::cout << "=== test_strided_copy_reduces_tile_and_sets_offset ===" << std::endl; - - constexpr int A = 1; - constexpr int B = 2; - - MetaMap meta; - meta[A] = make_meta(/*ndims=*/1, /*size0=*/512, /*offset0=*/0, /*tile=*/128, /*final=*/true); - meta[B] = make_meta(/*ndims=*/1, /*size0=*/100, /*offset0=*/0, /*tile=*/128, /*final=*/false); - - ASTNode* copy = make_copy(B, make_leaf(A), - make_region<1>({2}, {202}, {2}), - /*is_temp=*/false, /*ndims=*/1); - DAG* dag = make_single_node_dag({copy}); - - decomposition_solver::compute_decompositions(meta, dag, /*odf=*/4, /*num_pes=*/1); - - expect_eq("B.tile", meta.at(B).tile, 64); - expect_eq("B.offset[0]", meta.at(B).offset[0], 1); - expect_eq("B.phase", meta.at(B).offset[0] % meta.at(B).tile, 1); - - delete dag; -} - -} // namespace - -int main() { - try { - test_shifted_copy_uses_absolute_offset(); - test_shifted_copy_from_nonzero_source_offset(); - test_shifted_expression_temp_uses_absolute_representative(); - test_jacobi1d_temp_offset(); - test_jacobi2d_temp_offset(); - test_jacobi3d_temp_offset(); - test_strided_copy_reduces_tile_and_sets_offset(); - } catch (const std::exception& ex) { - std::cerr << "test_compute_decompositions failed: " << ex.what() << std::endl; - return 1; - } - - std::cout << "test_compute_decompositions passed" << std::endl; - return 0; -} diff --git a/tests/test_jit.cpp b/tests/test_jit.cpp deleted file mode 100644 index 3cae520..0000000 --- a/tests/test_jit.cpp +++ /dev/null @@ -1,345 +0,0 @@ -/** - * @file test_jit.cpp - * @brief Standalone test program for the MLIRJitCompiler. - * - * Builds AST trees by hand and exercises buildFromAST + optimizeAndFuse. - * Prints MLIR IR before and after fusion so you can visually verify - * the generated code. - * - * Build (adjust paths to your MLIR/LLVM install): - * - * clang++ -std=c++17 -O2 \ - * -I${CHARMTYLES_HOME}/include \ - * -I${CHARMTYLES_HOME}/example/charmnumeric/src \ - * $(llvm-config --cxxflags) \ - * test_jit.cpp \ - * $(llvm-config --ldflags --libs all) \ - * -lMLIR -lMLIRLinalgDialect -lMLIRAffineDialect -lMLIRGPUDialect \ - * -lMLIRArithDialect -lMLIRMemRefDialect -lMLIRFuncDialect \ - * -lMLIRPass -lMLIRTransforms \ - * -o test_jit - */ - -#include -#include -#include -#include - -#include "jit.hpp" - -// ------------------------------------------------------------------ // -// Helpers to build AST nodes -// ------------------------------------------------------------------ // - -/// Create a leaf (NOOP) node representing an existing array. -static ASTNode* makeLeaf(int name, int ndims = 1) { - return new ASTNode(name, static_cast(Opcode::NOOP), - /*is_temp=*/false, /*is_scalar=*/false, - /*scalar=*/0.0f, /*ndims=*/ndims); -} - -/// Create a scalar leaf node. -static ASTNode* makeScalar(float value) { - static int scalar_counter = -1; - return new ASTNode(scalar_counter--, static_cast(Opcode::NOOP), - /*is_temp=*/false, /*is_scalar=*/true, value); -} - -/// Create a binary op node (ADD, SUB, MUL, DIV). -static ASTNode* makeBinOp(int result_name, Opcode op, ASTNode* lhs, ASTNode* rhs, - bool is_temp = false) { - auto* node = new ASTNode(result_name, static_cast(op), is_temp); - node->add_operand(lhs); - node->add_operand(rhs); - return node; -} - -/// Create a GET_REGION node with an n-dimensional ArrayRegion. -static ASTNode* makeGetRegion(int result_name, ASTNode* src, std::vector start, - std::vector stop, std::vector step) { - auto* node = new ASTNode(result_name, static_cast(Opcode::GET_REGION)); - node->add_operand(src); - node->region = new ArrayRegion(start, stop, step); - return node; -} - -// ------------------------------------------------------------------ // -// Test cases -// ------------------------------------------------------------------ // - -/// Test 1: Simple C = A + B -void test_simple_add() { - std::cout << "=== Test 1: C = A + B ===" << std::endl; - - ASTNode* a = makeLeaf(0); // array A (name=0) - ASTNode* b = makeLeaf(1); // array B (name=1) - ASTNode* c = makeBinOp(2, Opcode::ADD, a, b); // C = A + B - - // Build a flat AST: leaves are implicit, roots = [c] - AST ast({c}); - - MLIRJitCompiler jit; - jit.buildFromAST(&ast); - - std::cout << "--- Before fusion ---" << std::endl; - jit.dumpMLIR(); - - bool ok = jit.optimizeAndFuse(); - std::cout << "--- After fusion (ok=" << ok << ") ---" << std::endl; - jit.dumpMLIR(); - std::cout << std::endl; -} - -/// Test 2: Chained ops D = (A + B) * C — should fuse into one loop -void test_fused_add_mul() { - std::cout << "=== Test 2: D = (A + B) * C ===" << std::endl; - - ASTNode* a = makeLeaf(0); - ASTNode* b = makeLeaf(1); - ASTNode* c = makeLeaf(2); - - // tmp = A + B (temporary, should be fused away) - ASTNode* tmp = makeBinOp(100, Opcode::ADD, a, b, /*is_temp=*/true); - // D = tmp * C - ASTNode* d = makeBinOp(3, Opcode::MUL, tmp, c); - - AST ast({tmp, d}); // flat order: tmp first, then d - - MLIRJitCompiler jit; - jit.buildFromAST(&ast); - - std::cout << "--- Before fusion ---" << std::endl; - jit.dumpMLIR(); - - bool ok = jit.optimizeAndFuse(); - std::cout << "--- After fusion (ok=" << ok << ") ---" << std::endl; - jit.dumpMLIR(); - std::cout << std::endl; -} - -/// Test 3: Scalar broadcast C = 2.0 * A -void test_scalar_mul() { - std::cout << "=== Test 3: C = 2.0 * A ===" << std::endl; - - ASTNode* a = makeLeaf(0); - ASTNode* two = makeScalar(2.0f); - - ASTNode* c = makeBinOp(1, Opcode::MUL, two, a); - - AST ast({c}); - - MLIRJitCompiler jit; - jit.buildFromAST(&ast); - - std::cout << "--- Before fusion ---" << std::endl; - jit.dumpMLIR(); - - bool ok = jit.optimizeAndFuse(); - std::cout << "--- After fusion (ok=" << ok << ") ---" << std::endl; - jit.dumpMLIR(); - std::cout << std::endl; -} - -/// Test 4: Subview B = A[10:50:1] + A[10:50:1] -void test_subview() { - std::cout << "=== Test 4: subview A[10:50] ===" << std::endl; - - ASTNode* a = makeLeaf(0); - - // slice = A[10:50:1] - ASTNode* slice = makeGetRegion(100, a, {10}, {50}, {1}); - - // B = slice + slice - ASTNode* b = makeBinOp(1, Opcode::ADD, slice, slice); - - AST ast({slice, b}); - - MLIRJitCompiler jit; - jit.buildFromAST(&ast); - - std::cout << "--- Before fusion ---" << std::endl; - jit.dumpMLIR(); - - bool ok = jit.optimizeAndFuse(); - std::cout << "--- After fusion (ok=" << ok << ") ---" << std::endl; - jit.dumpMLIR(); - std::cout << std::endl; -} - -/// Test 5: All four binary ops chained E = (A + B) - (C * D) -void test_all_binops() { - std::cout << "=== Test 5: E = (A + B) - (C * D) ===" << std::endl; - - ASTNode* a = makeLeaf(0); - ASTNode* b = makeLeaf(1); - ASTNode* c = makeLeaf(2); - ASTNode* d = makeLeaf(3); - - ASTNode* ab = makeBinOp(100, Opcode::ADD, a, b, /*is_temp=*/true); - ASTNode* cd = makeBinOp(101, Opcode::MUL, c, d, /*is_temp=*/true); - ASTNode* e = makeBinOp(4, Opcode::SUB, ab, cd); - - AST ast({ab, cd, e}); - - MLIRJitCompiler jit; - jit.buildFromAST(&ast); - - std::cout << "--- Before fusion ---" << std::endl; - jit.dumpMLIR(); - - bool ok = jit.optimizeAndFuse(); - std::cout << "--- After fusion (ok=" << ok << ") ---" << std::endl; - jit.dumpMLIR(); - std::cout << std::endl; -} - -/// Test 6: 2D add C = A + B where A, B are 2D arrays -void test_2d_add() { - std::cout << "=== Test 6: 2D C = A + B ===" << std::endl; - - ASTNode* a = makeLeaf(0, /*ndims=*/2); // memref - ASTNode* b = makeLeaf(1, /*ndims=*/2); - ASTNode* c = makeBinOp(2, Opcode::ADD, a, b); - - AST ast({c}); - - MLIRJitCompiler jit; - jit.buildFromAST(&ast); - - std::cout << "--- Before fusion ---" << std::endl; - jit.dumpMLIR(); - - bool ok = jit.optimizeAndFuse(); - std::cout << "--- After fusion (ok=" << ok << ") ---" << std::endl; - jit.dumpMLIR(); - std::cout << std::endl; -} - -/// Test 7: 2D subview B = A[2:8, 0:10] -void test_2d_subview() { - std::cout << "=== Test 7: 2D subview A[2:8, 0:10] ===" << std::endl; - - ASTNode* a = makeLeaf(0, /*ndims=*/2); - - ASTNode* slice = makeGetRegion(100, a, {2, 0}, {8, 10}, {1, 1}); - - // B = slice + slice - ASTNode* b = makeBinOp(1, Opcode::ADD, slice, slice); - - AST ast({slice, b}); - - MLIRJitCompiler jit; - jit.buildFromAST(&ast); - - std::cout << "--- Before fusion ---" << std::endl; - jit.dumpMLIR(); - - bool ok = jit.optimizeAndFuse(); - std::cout << "--- After fusion (ok=" << ok << ") ---" << std::endl; - jit.dumpMLIR(); - std::cout << std::endl; -} - -/// Test 8: aligned packed fragments keep dense strides for stepped regions. -void test_align_fragments_step_strides() { - std::cout << "=== Test 8: align stepped packed fragments ===" << std::endl; - - ArrayRegion<1> parent({0}, {64}, {2}); - ArrayRegion<1> left({0}, {32}, {2}); - ArrayRegion<1> right({32}, {64}, {2}); - - std::vector packed_full(32); - std::vector packed_left(16); - std::vector packed_right(16); - for (int i = 0; i < 32; ++i) - packed_full[i] = static_cast(i); - for (int i = 0; i < 16; ++i) { - packed_left[i] = static_cast(100 + i); - packed_right[i] = static_cast(200 + i); - } - - std::vector>> inputs(2); - inputs[0].push_back({parent, packed_full.data()}); - inputs[1].push_back({left, packed_left.data()}); - inputs[1].push_back({right, packed_right.data()}); - - auto aligned = align_fragments<1, float>(inputs, {parent, parent}); - assert(aligned.size() == 2); - assert(aligned[0].size() == 2); - assert(aligned[1].size() == 2); - - assert(aligned[0][0].memref.sizes[0] == 16); - assert(aligned[0][1].memref.sizes[0] == 16); - assert(aligned[0][0].memref.strides[0] == 1); - assert(aligned[0][1].memref.strides[0] == 1); - assert(aligned[0][0].memref.aligned == packed_full.data()); - assert(aligned[0][1].memref.aligned == packed_full.data() + 16); - assert(aligned[1][0].memref.aligned == packed_left.data()); - assert(aligned[1][1].memref.aligned == packed_right.data()); - - std::vector dense_base(64); - for (int i = 0; i < 64; ++i) - dense_base[i] = static_cast(300 + i); - FragmentData<1, float> local_dense{parent, dense_base.data()}; - local_dense.src_strides[0] = 1; - - std::vector>> local_inputs(2); - local_inputs[0].push_back(local_dense); - local_inputs[1].push_back({left, packed_left.data()}); - local_inputs[1].push_back({right, packed_right.data()}); - - auto local_aligned = align_fragments<1, float>(local_inputs, {parent, parent}); - assert(local_aligned[0].size() == 2); - assert(local_aligned[0][0].memref.strides[0] == 2); - assert(local_aligned[0][1].memref.strides[0] == 2); - assert(local_aligned[0][0].memref.aligned == dense_base.data()); - assert(local_aligned[0][1].memref.aligned == dense_base.data() + 32); - - std::cout << "step alignment checks passed" << std::endl; - std::cout << std::endl; -} - -/// Test 9: mapping preserves the full logical extent of stepped regions. -void test_map_preserves_strided_extent() { - std::cout << "=== Test 9: map preserves stepped extent across tile boundaries ===" - << std::endl; - - ArrayRegion<1> parent_input({2}, {127}, {2}); - ArrayRegion<1> tile_input({64}, {127}, {2}); - ArrayRegion<1> parent_output({1}, {64}, {1}); - - auto mapped = map(tile_input, parent_input, parent_output); - assert(mapped.start[0] == 32); - assert(mapped.stop[0] == 64); - assert(mapped.step[0] == 1); - assert(mapped.size(0) == tile_input.size(0)); - - ArrayRegion<1> output_fragment({32}, {64}, {1}); - auto mapped_back = map(output_fragment, parent_output, parent_input); - assert(mapped_back.start[0] == 64); - assert(mapped_back.stop[0] == 128); - assert(mapped_back.step[0] == 2); - assert(mapped_back.size(0) == output_fragment.size(0)); - - std::cout << "strided map checks passed" << std::endl; - std::cout << std::endl; -} - -// ------------------------------------------------------------------ // -// Main -// ------------------------------------------------------------------ // - -int main() { - test_simple_add(); - test_fused_add_mul(); - test_scalar_mul(); - test_subview(); - test_all_binops(); - test_2d_add(); - test_2d_subview(); - test_align_fragments_step_strides(); - test_map_preserves_strided_extent(); - - std::cout << "All JIT tests passed." << std::endl; - return 0; -}