22 Commits

Author SHA1 Message Date
ilgeco 05a04b09a5 test pimcomp adversarial memory scheduling
Validate Operations / validate-operations (push) Has been cancelled
2026-08-21 17:07:37 +02:00
ilgeco a9559abec3 Merge branch 'TestRottoConDeadLock' of chef.heaplab.deib.polimi.it:nnicolosi/Raptor into TestRottoConDeadLock 2026-08-21 16:48:31 +02:00
NiccoloN d634484df2 update pimsim submodule
Validate Operations / validate-operations (push) Has been cancelled
2026-08-21 16:28:21 +02:00
ilgeco 2d001bafb6 Update Readme conflict 2026-08-21 15:29:38 +02:00
ilgeco 558faaf74e Update README 2026-08-21 15:25:25 +02:00
ilgeco 4e7fe721f8 pim simulator adversary mode 2026-08-21 15:22:16 +02:00
NiccoloN 6d08686d32 fix ablation study
Validate Operations / validate-operations (push) Has been cancelled
2026-08-21 15:20:45 +02:00
NiccoloN b009e1ff08 add ablation study
Validate Operations / validate-operations (push) Has been cancelled
normalize names and artifact paths
2026-08-20 17:58:02 +02:00
NiccoloN add20e56eb minor fix
Validate Operations / validate-operations (push) Has been cancelled
2026-08-19 16:33:04 +02:00
NiccoloN db8d1c1707 better throughput in pipeline mode
Validate Operations / validate-operations (push) Has been cancelled
2026-08-19 16:03:52 +02:00
NiccoloN 4a2487d095 update submodule
Validate Operations / validate-operations (push) Has been cancelled
2026-08-11 11:55:58 +02:00
NiccoloN 45072ca743 add pipeline stages synchronization
Validate Operations / validate-operations (push) Has been cancelled
full ops throughput validation now passes
2026-08-11 11:34:31 +02:00
NiccoloN c55d9f3dad add throughput mode to validation scripts
make raptor also emit input sizes
2026-08-11 10:34:50 +02:00
NiccoloN 910701dfaf add throughput mode to pim-simulator 2026-08-11 10:28:28 +02:00
NiccoloN c69bec6636 rename ops validation onnx better
update related operations readme
2026-08-08 11:03:51 +02:00
NiccoloN 1b7d22b87e better comparison scripts
Validate Operations / validate-operations (push) Has been cancelled
2026-08-07 13:19:03 +02:00
NiccoloN ac84040e16 fix timeouts and pimcomp artifacts dir
Validate Operations / validate-operations (push) Has been cancelled
2026-08-07 11:16:36 +02:00
NiccoloN 4ce2ec8171 avoid raptor automatic build from comparison script
Validate Operations / validate-operations (push) Has been cancelled
2026-08-06 22:04:43 +02:00
NiccoloN 1c07faace9 minor fix
Validate Operations / validate-operations (push) Has been cancelled
2026-08-06 21:57:59 +02:00
NiccoloN 2e76164aed more complete pimcomp comparison scripts
Validate Operations / validate-operations (push) Has been cancelled
update pimsim-nn submodule
2026-08-06 21:49:54 +02:00
NiccoloN 4acd3b0c81 restore unwanted changes
Validate Operations / validate-operations (push) Has been cancelled
2026-08-06 15:01:30 +02:00
ilgeco 42c236b6a5 Raptor ggraph explorer main
Validate Operations / validate-operations (push) Has been cancelled
2026-08-06 14:48:09 +02:00
146 changed files with 9700 additions and 3793 deletions
@@ -0,0 +1,22 @@
# pimsim-nn Oracle Invariant
`backend-simulators/pim/pimsim-nn` is the performance oracle. Its simulation
behavior defines the hardware model used for Raptor/PIMCOMP comparisons.
## Required invariant
- Do not add instruction or operator support to pimsim-nn.
- Changes to pimsim-nn must preserve simulation behavior exactly. Acceptable
changes are limited to behavior-neutral maintenance proven not to alter
simulated timing, scheduling, power, energy, or supported input programs.
- Unsupported pimsim-nn operations must remain unsupported; do not approximate
their timing or map them onto another operation.
- Adapt compiler inputs to the oracle instead. For YOLO, use
`validation/networks/pimcomp_models/yolo11n/yolo11n-pimsim-nn.onnx`, the
dedicated pimsim-ready performance artifact with Softmax operations removed.
Use `validation/networks/yolo11n/depth_51/yolo11n_depth_51.onnx` for YOLO
functional validation; removing Softmax changes the model's numerical
behavior, so the pimsim-ready artifact is not a correctness reference.
Any proposed pimsim-nn behavior change requires explicit user authorization and
must not be introduced as part of a compiler optimization.
@@ -0,0 +1,53 @@
# Pipeline Scheduling Invariant
## Scope
This invariant applies to pipeline stage partitioning, physical-core
assignment, scheduled materialization, deferred transfers, and pipeline
synchronization.
## Invariant
A scheduled compute operation and all of its lanes belong to exactly one
pipeline stage. An operation may consume results produced in its own stage or
the immediately preceding stage only. Therefore every compute-graph edge from
stage `S` targets stage `S` or `S + 1`; backward edges and dependencies that
skip a stage are invalid.
Dynamic function inputs are stage-zero sources. Any operation that directly
consumes one must belong to stage 0. A later stage may consume that data only
through an explicit result forwarded by the preceding stage.
Each logical core belongs to exactly one stage capacity range before physical
placement. Those ranges cover every core but may have different sizes when the
initial partitioner predicts a lower maximum stage interval. Physical placement
may map a stage to arbitrary core IDs using the injected target topology.
Synchronization and deferred transfers consume the explicit stage identity;
they must not infer it from a physical core number after placement.
## Ownership
Logical PEFT remains pipeline-agnostic. Stage partitioning is the first phase
of pipeline scheduling and owns this invariant. It must construct a valid
operation-level partition before physical-core packing. Operations split for
physical capacity retain one shared stage identity. Repacking may move work
only within its assigned stage. Deferred-transfer planning and
synchronization consume the verified stage assignment; they must not repair
or reinterpret it.
## Verification
Before scheduled materialization, verify that:
- every compute instance has one valid physical core and stage;
- all lanes of one compute operation have the same stage;
- every direct dynamic-function-input consumer belongs to stage 0;
- every compute-graph edge stays within a stage or advances exactly one stage;
- every stage-local resident-weight set fits its assigned physical core; and
- stage capacities cover all logical cores exactly once; and
- physical placement is a permutation of all target cores.
Pipeline scheduling tests must include an uneven physical-core layout and a
graph with a long-lived dependency that would cross multiple naive stage
cuts. End-to-end validation must preserve functional results and exercise the
existing synchronization lowering without simulator changes.
+2
View File
@@ -6,6 +6,8 @@ Before modifying the relevant subsystem, read:
* `.agents/invariants/GRAPH_COMPUTE_BATCH_INVARIANT.md` * `.agents/invariants/GRAPH_COMPUTE_BATCH_INVARIANT.md`
* `.agents/invariants/PERFORMANCE_OPTIMIZATION_INVARIANT.md` * `.agents/invariants/PERFORMANCE_OPTIMIZATION_INVARIANT.md`
* `.agents/invariants/PIMSIM_NN_ORACLE_INVARIANT.md`
* `.agents/invariants/PIPELINE_SCHEDULING_INVARIANT.md`
* `.agents/invariants/SPATIAL_TARGET_GENERALITY_INVARIANT.md` * `.agents/invariants/SPATIAL_TARGET_GENERALITY_INVARIANT.md`
* Build commands: * Build commands:
* `cmake --build ./build_release` * `cmake --build ./build_release`
+92 -34
View File
@@ -5,7 +5,7 @@ targeting in-memory computing / processing-in-memory (PIM) architectures. It
extends ONNX-MLIR with a PIM accelerator and progressively lowers ONNX-MLIR extends ONNX-MLIR with a PIM accelerator and progressively lowers ONNX-MLIR
through custom MLIR dialects to simulator artifacts. through custom MLIR dialects to simulator artifacts.
The current target is the PIM simulator stack under `backend-simulators/pim`. The current target is the Pim simulator stack under `backend-simulators/pim`.
Raptor emits binary per-core `.pim` instruction files by default, plus Raptor emits binary per-core `.pim` instruction files by default, plus
`memory.bin`, `config.json`, and weight binaries. It can also emit per-core JSON `memory.bin`, `config.json`, and weight binaries. It can also emit per-core JSON
instruction files with `--pim-emit-json`. instruction files with `--pim-emit-json`.
@@ -29,9 +29,9 @@ lowering, scheduling, memory layout, and code-generation optimizations.
- `backend-simulators/pim/pim-simulator` is the in-tree Rust functional - `backend-simulators/pim/pim-simulator` is the in-tree Rust functional
simulator used by validation. It reads Raptor's `pim/` artifact directory and simulator used by validation. It reads Raptor's `pim/` artifact directory and
compares simulator output against native ONNX-MLIR execution. compares simulator output against native ONNX-MLIR execution.
- `backend-simulators/pim/pimsim-nn` is the non-functional simulator submodule - `backend-simulators/pim/pimsim-nn` contains the non-functional Pimsim
used internally by validation for latency, power, and energy. simulator used internally by validation for latency, power, and energy.
The helper scripts in `pimcomp_utils/` are for comparison with PIMCOMP-NN and The helper scripts in `pimcomp_utils/` are for comparison with Pimcomp and
contain local paths; treat them as local utilities, not portable workflows. contain local paths; treat them as local utilities, not portable workflows.
## Compilation pipeline ## Compilation pipeline
@@ -43,7 +43,7 @@ them to ONNX-MLIR through generated shim directories under
High-level lowering flow: High-level lowering flow:
``` ```
ONNX-MLIR -> Spatial -> Pim (tensor) -> Pim (bufferized) -> PIM artifacts ONNX-MLIR -> Spatial -> Pim (tensor) -> Pim (bufferized) -> Pim artifacts
``` ```
1. **ONNX -> Spatial** (`src/PIM/Conversion/ONNXToSpatial`). 1. **ONNX -> Spatial** (`src/PIM/Conversion/ONNXToSpatial`).
@@ -81,20 +81,20 @@ ONNX-MLIR -> Spatial -> Pim (tensor) -> Pim (bufferized) -> PIM artifacts
addressable accesses, and `PimBufferizationVerification` checks tensor addressable accesses, and `PimBufferizationVerification` checks tensor
absence, contiguity, and copy address spaces. absence, contiguity, and copy address spaces.
5. **PIM local-memory planning** 5. **Pim local-memory planning**
(`src/PIM/Dialect/Pim/Passes/Transforms/LocalMemoryPlanning`). (`src/PIM/Dialect/Pim/Passes/Transforms/LocalMemoryPlanning`).
Computes whole-core lifetimes, reuses addresses for non-overlapping Computes whole-core lifetimes, reuses addresses for non-overlapping
allocations, and records the explicit plan in PIM IR. Reusable lifetime allocations, and records the explicit plan in Pim IR. Reusable lifetime
analysis lives under `src/PIM/Dialect/Pim/Passes/Analyses`. analysis lives under `src/PIM/Dialect/Pim/Passes/Analyses`.
6. **PIM verification and code generation** (`src/PIM/Passes/PimCodegen` and 6. **Pim verification and code generation** (`src/PIM/Passes/PimCodegen` and
`src/PIM/Compiler`). `src/PIM/Compiler`).
Verifies the memory plan and other PIM invariants, then emits `.pim` core Verifies the memory plan and other Pim invariants, then emits `.pim` core
files, weights, and `memory.bin` / `config.json` without rerunning liveness. files, weights, and `memory.bin` / `config.json` without rerunning liveness.
Supporting pieces: Supporting pieces:
- `src/PIM/Common` - shared IR, filesystem, diagnostics, reports, and utility - `src/PIM/Common` - shared IR, filesystem, diagnostics, reports, and utility
helpers. helpers.
- `src/PIM/Compiler` - PIM compiler options, planned-address materialization, binary - `src/PIM/Compiler` - Pim compiler options, planned-address materialization, binary
instruction format, artifact writing, weight emission, and codegen entry instruction format, artifact writing, weight emission, and codegen entry
points. points.
- `src/PIM/Conversion/SpatialToGraphviz` - optional Spatial graphviz conversion - `src/PIM/Conversion/SpatialToGraphviz` - optional Spatial graphviz conversion
@@ -102,40 +102,97 @@ Supporting pieces:
- `src/PIM/Passes` - pass registration and auxiliary passes. - `src/PIM/Passes` - pass registration and auxiliary passes.
- `src/PIM/PimAccelerator.{cpp,hpp}` - ONNX-MLIR accelerator entry point. - `src/PIM/PimAccelerator.{cpp,hpp}` - ONNX-MLIR accelerator entry point.
## PIM compiler options ## Pim compiler options
Pass these to `onnx-mlir` when compiling for PIM. These are all Raptor/PIM-specific Pass these to `onnx-mlir` when compiling for Pim. These are all Raptor/Pim-specific
options; `onnx-mlir --help` lists the inherited ONNX-MLIR options. options; `onnx-mlir --help` lists the inherited ONNX-MLIR options.
- `--maccel=PIM` - select the PIM accelerator. - `--maccel=PIM` - select the Pim accelerator. Default: no Pim accelerator.
- `--EmitSpatial`, `--EmitPim`, `--EmitPimBufferized`, - `--EmitSpatial`, `--EmitPim`, `--EmitPimBufferized`,
`--EmitPimCodegen` - stop the PIM pipeline at the requested stage. The PIM `--EmitPimCodegen` - stop the Pim pipeline at the requested stage. Default:
default is `--EmitPimCodegen`. `--EmitPimCodegen` for Pim compilation.
- `--core-count=<N>` - required positive core count for PIM compilation. - `--core-count=<N>` - required positive core count for Pim compilation.
- `--crossbar-size=<N>` - crossbar width/height. Default in code is `128`. Default: none; this option is required.
- `--crossbar-count=<N>` - crossbars per core. Default in code is `64`. - `--crossbar-size=<N>` - required positive crossbar width/height for Pim
- `--pim-target-config=<PATH>` - optional PIM target configuration used by the compilation. Default: none; this option is required.
- `--crossbar-count=<N>` - required positive crossbar count per core for Pim
compilation. Default: none; this option is required.
- `--pipeline=<N>` - number of throughput pipeline stages; `1` preserves
latency scheduling. Default: `1`.
- `--pim-target-config=<PATH>` - optional Pim target configuration used by the
target adapter to construct the target-neutral Spatial scheduling cost and target adapter to construct the target-neutral Spatial scheduling cost and
topology model. Resource values must match the explicit core/crossbar flags. topology model. Resource values must match the explicit core/crossbar flags.
Default: empty; use the built-in target model.
- `--pim-memory-report=<summary|none>` - emit the concise combined memory report - `--pim-memory-report=<summary|none>` - emit the concise combined memory report
under `reports/memory_report.txt`, or disable it. Default is `summary`. under `reports/memory_report.txt`, or disable it. Default: `summary`.
- `--pim-only-codegen` - assume input is already bufferized PIM IR and only run - `--pim-only-codegen` - assume input is already bufferized Pim IR and only run
the codegen tail. the codegen tail. Default: off.
- `--pim-disable-synchronization` - omit generated `wait` and `sync`
instructions for performance ablation. Default: off.
- `--pim-disable-spatial-planning` - select the first, trivial DenseNCHW layout
alternative for every Spatial plan operation, disabling cost-based layout
planning while leaving ONNX rewrites and graph-compute merging enabled.
Default: off.
### Spatial layout plan variants
Spatial plan operations advertise alternatives as an exact combination of
operand physical layouts and one result physical layout. Every plan operation
has the default `DenseNCHW -> DenseNCHW` alternative. The planner can select
the following additional variants when the operation, tensor shapes, and
target resources make them legal:
| Physical layout or plan | Meaning and current use |
|---|---|
| `DenseNCHW` | Ordinary dense NCHW storage. This is the first alternative and the one selected by `--pim-disable-spatial-planning`. |
| `NHWCRowStrip` | Row-strip storage for NCHW logical tensors: spatial rows are processed as channel vectors. This enables row-strip lowering through compatible chains. |
| `Fragmented` | Fragmented physical input accepted by `Flatten`, which reassembles it to dense NCHW. It is not currently selected as a plan result. |
| `NCHWRowStrip` | A Spatial IR layout enum value reserved for NCHW-oriented row strips; current layout-capability implementations do not advertise it as a plan alternative. |
The operation-specific non-trivial alternatives are:
| Plan operation | Additional alternatives beyond dense NCHW |
|---|---|
| `Conv2D` | Dense input to row-strip output, or row-strip input to row-strip output when the target-dependent Conv lowering supports it. |
| `Flatten` | Fragmented input to dense output, or row-strip input to dense output when legal. |
| `Relu` | Row-strip input to row-strip output. |
| `SiLU` | Row-strip input to row-strip output, with a stronger intrinsic cost preference than the generic row-strip variant. |
| `ResizeNearest` | Row-strip input to row-strip output when its lowering is legal. |
| `MaxPool2D` | Dense input to row-strip output, or row-strip input to row-strip output. |
| `GlobalAveragePool` | Dense input to row-strip output, or row-strip input to row-strip output. |
| `BiasAdd` | Row-strip data input plus a dense bias input to row-strip output when the bias shape is supported. |
| `Add` | All data inputs row-strip to row-strip output. |
| `Concat` | All inputs row-strip to row-strip output. |
Cost-based planning scores intrinsic alternative cost, operand layout
mismatches, and downstream incompatibility, then iterates in alternating
forward and reverse operation order until the bounded analysis converges.
Function results are required to remain `DenseNCHW`; explicit materialization
operations reconcile layout mismatches at boundaries. With
`--pim-disable-spatial-planning`, the pass still runs and records a valid plan,
but chooses the first dense alternative for every plan operation. Later graph
compute merging is unchanged, so elementwise operations such as `Relu` remain
separate from neighboring parallel operations and can create fan-out/fan-in
diamonds.
- `--pim-emit-json` - also emit `core_*.json` instruction files alongside - `--pim-emit-json` - also emit `core_*.json` instruction files alongside
`core_*.pim`. `core_*.pim`. Default: off.
- `--pim-export-spatial-dataflow=<none|spatial1|spatial2|spatial3|spatial4|all>` - - `--pim-export-spatial-dataflow=<none|spatial1|spatial2|spatial3|spatial4|all>` -
control Spatial dataflow CSV reports for the graph, trivially merged graph, control Spatial dataflow CSV reports for the graph, trivially merged graph,
scheduled, and realized snapshots under `reports/`. Default is `none`. scheduled, and realized snapshots under `reports/`. Default: `none`.
- `--pim-conv-lowering=<auto|legacy|depthwise|packed-im2col|streamed-patch|streamed-packed|output-channel-tiled|input-k-tiled|tiled-2d>` - - `--pim-conv-lowering=<auto|legacy|depthwise|packed-im2col|streamed-patch|streamed-packed|output-channel-tiled|input-k-tiled|tiled-2d>` -
select the convolution lowering strategy. Default is `auto`. select the convolution lowering strategy. Default: `auto`.
- `--pim-conv-im2col-max-elements=<N>` - maximum globally materialized im2col - `--pim-conv-im2col-max-elements=<N>` - maximum globally materialized im2col
elements per convolution before streaming. Default is `1048576`. elements per convolution before streaming. Default: `1048576`.
- `--pim-conv-stream-chunk-positions=<N>` - maximum output positions per - `--pim-conv-stream-chunk-positions=<N>` - maximum output positions per
streamed convolution chunk. Default is `1024`. streamed convolution chunk. Default: `1024`.
- `--pim-report-conv-lowering=<true|false>` - emit a bounded convolution
lowering report. Default: `true`.
- `--pim-detect-communication-deadlock` - statically simulate expanded - `--pim-detect-communication-deadlock` - statically simulate expanded
send/receive ordering and reject blocking deadlocks. Default is off. send/receive ordering and reject blocking deadlocks. Default: off.
- `--pim-verify-bufferization-copy-freedom` - run the expensive official Pim
tensor-copy freedom proof before bufferization. Default: off.
## Standard PIM hardware profile ## Standard Pim hardware profile
Raptor's standard development and YOLO validation profile is: Raptor's standard development and YOLO validation profile is:
@@ -149,7 +206,8 @@ Canonical compiler flags:
`--crossbar-count=64 --crossbar-size=128 --core-count=144` `--crossbar-count=64 --crossbar-size=128 --core-count=144`
`--core-count` remains mandatory and must be passed explicitly to the compiler. `--crossbar-size`, `--crossbar-count`, and `--core-count` remain mandatory and
must be passed explicitly to the compiler.
Example: Example:
@@ -159,11 +217,11 @@ Example:
--crossbar-count=64 --crossbar-size=128 --core-count=144 --crossbar-count=64 --crossbar-size=128 --core-count=144
``` ```
This writes PIM artifacts under `/tmp/raptor/pim/`. This writes Pim artifacts under `/tmp/raptor/pim/`.
## Validation ## Validation
Functional validation compiles ONNX models, compares native ONNX-MLIR and PIM Functional validation compiles ONNX models, compares native ONNX-MLIR and Pim
simulator outputs, and optionally reports latency, power, and energy. See simulator outputs, and optionally reports latency, power, and energy. See
[`validation/README.md`](validation/README.md) for prerequisites, usage, [`validation/README.md`](validation/README.md) for prerequisites, usage,
options, artifacts, and results. options, artifacts, and results.
@@ -276,7 +334,7 @@ cd backend-simulators/pim/pim-simulator
cargo test cargo test
``` ```
## Repository Layout ## Repository layout
- `src/PIM/` - PIM accelerator implementation. - `src/PIM/` - PIM accelerator implementation.
- `test/PIM/` - PIM C++ unit tests. - `test/PIM/` - PIM C++ unit tests.
@@ -284,6 +342,6 @@ cargo test
slices, and pimsim config generation. slices, and pimsim config generation.
- `backend-simulators/pim/pim-simulator/` - in-tree Rust functional simulator. - `backend-simulators/pim/pim-simulator/` - in-tree Rust functional simulator.
- `backend-simulators/pim/pimsim-nn/` - non-functional simulator submodule. - `backend-simulators/pim/pimsim-nn/` - non-functional simulator submodule.
- `pimcomp_utils/` - local comparison helpers for PIMCOMP-NN. - `pimcomp_utils/` - local comparison helpers for Pimcomp.
- `.github/actions/` and `.github/workflows/validate_operations.yml` - CI setup - `.github/actions/` and `.github/workflows/validate_operations.yml` - CI setup
for MLIR/Protobuf caching, building Raptor, and validation. for MLIR/Protobuf caching, building Raptor, and validation.
@@ -4,17 +4,18 @@ use mimalloc::MiMalloc;
static GLOBAL: MiMalloc = MiMalloc; static GLOBAL: MiMalloc = MiMalloc;
use anyhow::{Context, Result, bail}; use anyhow::{Context, Result, bail};
use clap::Parser; use clap::{Parser, ValueEnum};
use glob::glob; use glob::glob;
use pimcore::binary_to_instruction::binary_to_executor; use pimcore::binary_to_instruction::binary_to_executor;
use pimcore::cpu::crossbar::Crossbar; use pimcore::cpu::crossbar::Crossbar;
use pimcore::json_to_instruction::json_to_executor; use pimcore::json_to_instruction::json_to_executor;
use pimcore::memory_manager::CoreMemory; use pimcore::memory_manager::CoreMemory;
use pimcore::tracing::TRACER; use pimcore::tracing::TRACER;
use pimcore::{DiagnosticSchedulePolicy, DiagnosticScheduleTarget};
use serde_json::Value; use serde_json::Value;
use std::collections::HashMap; use std::collections::HashMap;
use std::fs::{self, File}; use std::fs::{self, File};
use std::io::{BufReader, Write}; use std::io::BufReader;
use std::path::PathBuf; use std::path::PathBuf;
/// Program to simulate core execution configuration /// Program to simulate core execution configuration
@@ -44,14 +45,79 @@ struct Args {
/// Comma separated list of (address,size) for memory output dump /// Comma separated list of (address,size) for memory output dump
#[arg(short, long, value_delimiter = ',', num_args = 1.., value_name = "ADDR,SIZE")] #[arg(short, long, value_delimiter = ',', num_args = 1.., value_name = "ADDR,SIZE")]
dump: Vec<usize>, dump: Vec<usize>,
/// Simulator execution mode
#[arg(long, value_enum, default_value_t = ExecutionMode::Latency)]
mode: ExecutionMode,
/// Number of inputs to execute (required in throughput mode)
#[arg(long)]
batch_size: Option<u32>,
/// Directory containing input_*.bin files, one per batch entry
#[arg(long = "input-dir")]
input_dir: PathBuf,
/// Optional directory for per-iteration output dumps
#[arg(long)]
batch_output_dir: Option<PathBuf>,
/// Optional JSONL shadow provenance trace
#[arg(long)]
provenance_trace: Option<PathBuf>,
/// Diagnostic-only barrier between global throughput iterations
#[arg(long)]
provenance_global_barrier: bool,
/// Diagnostic-only ready-core delay, formatted as CORE:CYCLES
#[arg(long, value_name = "CORE:CYCLES")]
provenance_core_stall: Option<String>,
/// Diagnostic scheduler policy; greedy is the unchanged default
#[arg(long, value_enum, default_value_t = DiagnosticSchedulePolicyArg::Greedy)]
diagnostic_schedule_policy: DiagnosticSchedulePolicyArg,
/// Deterministic seed for the randomized diagnostic scheduler
#[arg(long, default_value_t = 0)]
diagnostic_schedule_seed: u64,
/// Adversarial target: WCORE:WPC:RCORE:RPC:BEGIN:END[:READER_ITER:WRITER_MIN_ITER]
#[arg(long, value_name = "WCORE:WPC:RCORE:RPC:BEGIN:END[:RITER:WMIN]")]
diagnostic_schedule_target: Option<String>,
/// Maximum number of target-consumer deferrals
#[arg(long, default_value_t = 10_000)]
diagnostic_schedule_deferral_budget: u64,
/// Diagnostic delay applied when a target reader reaches its PC, formatted as CORE:PC:CYCLES or CORE:PC:ITERATION:CYCLES
#[arg(long, value_name = "CORE:PC[:ITERATION]:CYCLES")]
diagnostic_target_stall: Option<String>,
}
#[derive(Clone, Debug, ValueEnum)]
enum ExecutionMode {
Latency,
Throughput,
}
#[derive(Clone, Copy, Debug, ValueEnum)]
enum DiagnosticSchedulePolicyArg {
Greedy,
Randomized,
Adversarial,
} }
fn main() -> Result<()> { fn main() -> Result<()> {
let args = Args::parse(); let args = Args::parse();
let config_json = retrive_config(&args)?; let config_json = retrieve_config(&args)?;
let mut core_inputs = retrive_cores(&args)?; let batch_size = batch_size(&args)?;
let memory = retrive_memory(&args)?; let input_regions = input_regions(&config_json)?;
let input_data = retrieve_inputs(&args, batch_size)?;
let inputs: Vec<&[u8]> = input_data.iter().map(Vec::as_slice).collect();
let mut core_inputs = retrieve_cores(&args)?;
let memory = retrieve_memory(&args)?;
let global_crossbars = get_crossbars(&config_json, &args).unwrap(); let global_crossbars = get_crossbars(&config_json, &args).unwrap();
let crossbars = map_crossbars_to_cores(&config_json, &args, &global_crossbars); let crossbars = map_crossbars_to_cores(&config_json, &args, &global_crossbars);
let mut executor = match &mut core_inputs { let mut executor = match &mut core_inputs {
@@ -63,15 +129,161 @@ fn main() -> Result<()> {
} }
}; };
set_memory(&mut executor, memory); set_memory(&mut executor, memory);
if let Some(path) = &args.provenance_trace {
executor.enable_provenance(path)?;
}
executor.set_provenance_global_barrier(args.provenance_global_barrier);
if let Some(spec) = args.provenance_core_stall.as_deref() {
let (core, cycles) = parse_core_stall(spec)?;
executor.set_provenance_core_stall(core, cycles);
}
executor.set_diagnostic_schedule_policy(match args.diagnostic_schedule_policy {
DiagnosticSchedulePolicyArg::Greedy => DiagnosticSchedulePolicy::Greedy,
DiagnosticSchedulePolicyArg::Randomized => DiagnosticSchedulePolicy::Randomized,
DiagnosticSchedulePolicyArg::Adversarial => DiagnosticSchedulePolicy::Adversarial,
});
executor.set_diagnostic_schedule_seed(args.diagnostic_schedule_seed);
executor.set_diagnostic_schedule_deferral_budget(args.diagnostic_schedule_deferral_budget);
if let Some(spec) = args.diagnostic_target_stall.as_deref() {
let (core, pc, iteration, cycles) = parse_target_stall(spec)?;
executor.set_diagnostic_target_stall(core, pc, iteration, cycles);
}
if let Some(spec) = args.diagnostic_schedule_target.as_deref() {
executor.set_diagnostic_schedule_target(parse_schedule_target(spec)?);
} else if matches!(
args.diagnostic_schedule_policy,
DiagnosticSchedulePolicyArg::Adversarial
) {
bail!("adversarial scheduling requires --diagnostic-schedule-target");
}
TRACER TRACER
.lock() .lock()
.unwrap() .unwrap()
.init(executor.cpu().num_core(), args.output.clone()); .init(executor.cpu().num_core(), args.output.clone());
executor.execute()?; let dumps = dump_ranges(&args.dump)?;
dump_memory(executor, &args)?; let batch_outputs = executor.execute_batch(&inputs, &input_regions, &dumps)?;
fs::write(
&args.output,
batch_outputs
.last()
.context("simulation produced no output")?,
)?;
if let Some(batch_output_dir) = args.batch_output_dir {
write_batch_outputs(batch_output_dir, batch_outputs)?;
}
Ok(()) Ok(())
} }
fn parse_core_stall(spec: &str) -> Result<(usize, u64)> {
let (core, cycles) = spec
.split_once(':')
.context("--provenance-core-stall must be CORE:CYCLES")?;
let core = core.parse().context("invalid stalled core")?;
let cycles = cycles.parse().context("invalid stall cycle count")?;
if cycles == 0 {
bail!("--provenance-core-stall cycles must be positive");
}
Ok((core, cycles))
}
fn parse_schedule_target(spec: &str) -> Result<DiagnosticScheduleTarget> {
let values: Vec<usize> = spec
.split(':')
.map(|value| {
value
.parse()
.with_context(|| format!("invalid schedule target field: {value}"))
})
.collect::<Result<_>>()?;
if values.len() != 6 && values.len() != 8 {
bail!(
"--diagnostic-schedule-target requires WCORE:WPC:RCORE:RPC:BEGIN:END with optional RITER:WMIN"
);
}
if values[5] <= values[4] {
bail!("schedule target address end must be greater than begin");
}
Ok(DiagnosticScheduleTarget {
writer_core: values[0],
writer_pc: values[1],
reader_core: values[2],
reader_pc: values[3],
address_begin: values[4],
address_end: values[5],
reader_iteration: (values.len() == 8).then_some(values[6] as u32),
writer_min_iteration: (values.len() == 8).then_some(values[7] as u32),
})
}
fn parse_target_stall(spec: &str) -> Result<(usize, usize, Option<u32>, u64)> {
let values: Vec<&str> = spec.split(':').collect();
if values.len() != 3 && values.len() != 4 {
bail!("--diagnostic-target-stall must be CORE:PC:CYCLES or CORE:PC:ITERATION:CYCLES");
}
let core = values[0].parse().context("invalid target-stall core")?;
let pc = values[1].parse().context("invalid target-stall PC")?;
let (iteration, cycle_field) = if values.len() == 4 {
(
Some(
values[2]
.parse()
.context("invalid target-stall iteration")?,
),
values[3],
)
} else {
(None, values[2])
};
let cycles = cycle_field
.parse()
.context("invalid target-stall cycle count")?;
if cycles == 0 {
bail!("--diagnostic-target-stall cycles must be positive");
}
Ok((core, pc, iteration, cycles))
}
fn batch_size(args: &Args) -> Result<u32> {
match (&args.mode, args.batch_size) {
(ExecutionMode::Latency, None | Some(1)) => Ok(1),
(ExecutionMode::Latency, Some(_)) => bail!("latency mode requires batch size 1"),
(ExecutionMode::Throughput, Some(0)) => bail!("batch size must be positive"),
(ExecutionMode::Throughput, Some(batch_size)) => Ok(batch_size),
(ExecutionMode::Throughput, None) => bail!("throughput mode requires --batch-size"),
}
}
fn input_regions(config: &Value) -> Result<Vec<(usize, usize)>> {
let addresses = config
.get("inputs_addresses")
.and_then(Value::as_array)
.context("config.json has no inputs_addresses array")?;
let sizes = config
.get("inputs_sizes")
.and_then(Value::as_array)
.context("config.json has no inputs_sizes array")?;
if addresses.len() != sizes.len() {
bail!("config.json input address/size count mismatch");
}
addresses
.iter()
.zip(sizes)
.map(|(address, size)| {
Ok((
usize::try_from(address.as_u64().context("invalid input address")?)?,
usize::try_from(size.as_u64().context("invalid input size")?)?,
))
})
.collect()
}
fn retrieve_inputs(args: &Args, batch_size: u32) -> Result<Vec<Vec<u8>>> {
(0..batch_size)
.map(|index| args.input_dir.join(format!("input_{index}.bin")))
.map(|path| fs::read(&path).with_context(|| format!("Failed to read input file: {path:?}")))
.collect()
}
fn map_crossbars_to_cores<'c>( fn map_crossbars_to_cores<'c>(
config: &Value, config: &Value,
args: &Args, args: &Args,
@@ -114,7 +326,7 @@ fn map_crossbars_to_cores<'c>(
let path_as_str = real_path.to_str().unwrap(); let path_as_str = real_path.to_str().unwrap();
assert!( assert!(
global_crossbars.contains_key(path_as_str), global_crossbars.contains_key(path_as_str),
"symlink point to {:?}\n a not stored crossbar", "symlink points to {:?}\n a crossbar that was not stored",
real_path real_path
); );
@@ -131,7 +343,7 @@ fn map_crossbars_to_cores<'c>(
fn get_crossbars(config: &Value, args: &Args) -> anyhow::Result<HashMap<String, Crossbar>> { fn get_crossbars(config: &Value, args: &Args) -> anyhow::Result<HashMap<String, Crossbar>> {
let xbar_size = config.get("xbar_size").unwrap().as_array().unwrap(); let xbar_size = config.get("xbar_size").unwrap().as_array().unwrap();
let rows_crossbar = xbar_size[0].as_i64().unwrap() as usize; let rows_crossbar = xbar_size[0].as_i64().unwrap() as usize;
let column_corssbar = xbar_size[1].as_i64().unwrap() as usize; let column_crossbar = xbar_size[1].as_i64().unwrap() as usize;
let mut res = HashMap::new(); let mut res = HashMap::new();
if let Some(folder) = args.folder.as_ref() { if let Some(folder) = args.folder.as_ref() {
@@ -154,7 +366,7 @@ fn get_crossbars(config: &Value, args: &Args) -> anyhow::Result<HashMap<String,
let bytes = std::fs::read(weight_file.path()).expect("Failed to read binary file"); let bytes = std::fs::read(weight_file.path()).expect("Failed to read binary file");
let stored_row_bytes = bytes.len() / rows_crossbar; let stored_row_bytes = bytes.len() / rows_crossbar;
let mut crossbar = Crossbar::new( let mut crossbar = Crossbar::new(
std::cmp::max(column_corssbar * 4, stored_row_bytes), std::cmp::max(column_crossbar * 4, stored_row_bytes),
rows_crossbar, rows_crossbar,
CoreMemory::new(), CoreMemory::new(),
); );
@@ -174,21 +386,22 @@ fn get_crossbars(config: &Value, args: &Args) -> anyhow::Result<HashMap<String,
Ok(res) Ok(res)
} }
fn dump_memory(mut executor: pimcore::Executable, args: &Args) -> Result<()> { fn dump_ranges(values: &[usize]) -> Result<Vec<(usize, usize)>> {
let dumps: Vec<(usize, usize)> = args if !values.len().is_multiple_of(2) {
.dump bail!("memory dump requires address,size pairs");
}
Ok(values
.chunks_exact(2) .chunks_exact(2)
.map(|chunk| (chunk[0], chunk[1])) .map(|chunk| (chunk[0], chunk[1]))
.collect(); .collect())
let mut out_file = fs::OpenOptions::new() }
.create(true)
.write(true)
.truncate(true)
.open(&args.output)
.with_context(|| format!("cannot open file {:?} for writing", args.output))?;
for (address, size) in dumps { fn write_batch_outputs(output_dir: PathBuf, outputs: Vec<Vec<u8>>) -> Result<()> {
out_file.write_all(executor.cpu_mut().host().load::<u8>(address, size).unwrap()[0])?; fs::create_dir_all(&output_dir)
.with_context(|| format!("cannot create batch output directory {output_dir:?}"))?;
for (iteration, output) in outputs.into_iter().enumerate() {
let path = output_dir.join(format!("output_{iteration:06}.bin"));
fs::write(&path, output).with_context(|| format!("cannot write batch output {path:?}"))?;
} }
Ok(()) Ok(())
} }
@@ -197,7 +410,7 @@ fn set_memory(executor: &mut pimcore::Executable, memory: Vec<u8>) {
executor.cpu_mut().host().execute_store(0, &memory).unwrap(); executor.cpu_mut().host().execute_store(0, &memory).unwrap();
} }
fn retrive_memory(args: &Args) -> Result<Vec<u8>> { fn retrieve_memory(args: &Args) -> Result<Vec<u8>> {
let memory_path = if let Some(mem_override) = &args.memory { let memory_path = if let Some(mem_override) = &args.memory {
mem_override.clone() mem_override.clone()
} else if let Some(folder) = &args.folder.as_ref() { } else if let Some(folder) = &args.folder.as_ref() {
@@ -237,7 +450,7 @@ enum CoreInputs {
Binary(Vec<Vec<u8>>), Binary(Vec<Vec<u8>>),
} }
fn retrive_cores(args: &Args) -> Result<CoreInputs, anyhow::Error> { fn retrieve_cores(args: &Args) -> Result<CoreInputs, anyhow::Error> {
if let Some(cores_override) = &args.cores { if let Some(cores_override) = &args.cores {
let first_extension = cores_override let first_extension = cores_override
.first() .first()
@@ -310,7 +523,7 @@ fn core_sort_key(path: &PathBuf) -> i32 {
stem.parse::<i32>().unwrap() stem.parse::<i32>().unwrap()
} }
fn retrive_config(args: &Args) -> Result<Value, anyhow::Error> { fn retrieve_config(args: &Args) -> Result<Value, anyhow::Error> {
let config_path: PathBuf = { let config_path: PathBuf = {
let override_path = args.config.as_ref(); let override_path = args.config.as_ref();
let folder = args.folder.as_ref(); let folder = args.folder.as_ref();
@@ -80,19 +80,19 @@ fn read_i32_le(bytes: &[u8], offset: usize) -> i32 {
fn parse_binary_records(bytes: &[u8]) -> Result<Vec<InstructionRecord>> { fn parse_binary_records(bytes: &[u8]) -> Result<Vec<InstructionRecord>> {
ensure!(bytes.len() >= HEADER_SIZE, "binary core file too small"); ensure!(bytes.len() >= HEADER_SIZE, "binary core file too small");
ensure!(&bytes[0..4] == MAGIC, "invalid PIM binary magic"); ensure!(&bytes[0..4] == MAGIC, "invalid Pim binary magic");
let version = read_u32_le(bytes, 4); let version = read_u32_le(bytes, 4);
ensure!( ensure!(
version == VERSION, version == VERSION,
"unsupported PIM binary version {version}" "unsupported Pim binary version {version}"
); );
let instruction_count = read_u32_le(bytes, 8) as usize; let instruction_count = read_u32_le(bytes, 8) as usize;
let expected_len = HEADER_SIZE + instruction_count * RECORD_SIZE; let expected_len = HEADER_SIZE + instruction_count * RECORD_SIZE;
ensure!( ensure!(
bytes.len() == expected_len, bytes.len() == expected_len,
"PIM binary size mismatch: expected {expected_len} bytes, got {}", "Pim binary size mismatch: expected {expected_len} bytes, got {}",
bytes.len() bytes.len()
); );
@@ -335,7 +335,7 @@ fn append_record(
.set_offset_select_value(generic1, 0); .set_offset_select_value(generic1, 0);
inst_builder.make_inst(sync, inst_data_builder.build()); inst_builder.make_inst(sync, inst_data_builder.build());
} }
_ => bail!("unsupported PIM binary opcode {opcode}"), _ => bail!("unsupported Pim binary opcode {opcode}"),
} }
Ok(()) Ok(())
} }
@@ -2,17 +2,48 @@ use crate::utility::AddressArg;
use anyhow::{Context, Result, ensure}; use anyhow::{Context, Result, ensure};
use std::{collections::HashMap, fmt::Debug}; use std::{collections::HashMap, fmt::Debug};
use super::{DiagnosticSchedulePolicy, DiagnosticScheduleTarget};
use crate::{ use crate::{
cpu::crossbar::Crossbar, cpu::crossbar::Crossbar,
instruction_set::Instructions, instruction_set::Instructions,
memory_manager::{CoreMemory, MemoryStorable, type_traits::TryToUsize}, memory_manager::{CoreMemory, MemoryStorable, type_traits::TryToUsize},
provenance::ProvenanceTracker,
}; };
use serde_json::json;
pub mod crossbar; pub mod crossbar;
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct CPU<'a> { pub struct CPU<'a> {
cores: Box<[Core<'a>]>, cores: Box<[Core<'a>]>,
batch_outputs: Option<BatchOutputs>,
provenance: Option<ProvenanceTracker>,
}
#[derive(Debug, Clone)]
struct BatchOutputs {
iteration: usize,
ranges: Vec<(usize, usize)>,
outputs: Vec<Vec<u8>>,
}
impl BatchOutputs {
fn record(&mut self, address: usize, bytes: &[u8]) {
let output = &mut self.outputs[self.iteration];
let store_end = address + bytes.len();
let mut output_offset = 0;
for &(range_address, range_size) in &self.ranges {
let start = address.max(range_address);
let end = store_end.min(range_address + range_size);
if start < end {
let size = end - start;
output[output_offset + start - range_address
..output_offset + start - range_address + size]
.copy_from_slice(&bytes[start - address..start - address + size]);
}
output_offset += range_size;
}
}
} }
impl<'a> CPU<'a> { impl<'a> CPU<'a> {
@@ -25,9 +56,320 @@ impl<'a> CPU<'a> {
} }
Self { Self {
cores: cores.into(), cores: cores.into(),
batch_outputs: None,
provenance: None,
} }
} }
pub(crate) fn enable_provenance(
&mut self,
path: impl AsRef<std::path::Path>,
) -> std::io::Result<()> {
self.provenance = Some(ProvenanceTracker::new(self.cores.len(), path)?);
Ok(())
}
pub(crate) fn begin_provenance_batch(&mut self, batch_size: usize) {
if let Some(provenance) = &mut self.provenance {
provenance.begin_batch(batch_size);
}
}
pub(crate) fn set_execution_context(
&mut self,
cycle: u64,
core: usize,
pc: usize,
iteration: u32,
) {
if let Some(provenance) = &mut self.provenance {
provenance.set_context(cycle, core, pc, iteration);
}
}
pub(crate) fn provenance_input_store(&mut self, address: usize, size: usize, sample: u32) {
if let Some(provenance) = &mut self.provenance {
provenance.input_store(address, size, sample);
}
}
pub(crate) fn provenance_global_store_from_local(
&mut self,
core: usize,
global_address: usize,
local_address: usize,
size: usize,
) {
if let Some(provenance) = &mut self.provenance {
provenance.global_store_from_local(core, global_address, local_address, size);
}
}
pub(crate) fn provenance_global_load_to_local(
&mut self,
core: usize,
global_address: usize,
local_address: usize,
size: usize,
) {
if let Some(provenance) = &mut self.provenance {
provenance.global_load_to_local(core, global_address, local_address, size);
}
}
pub(crate) fn provenance_local_copy(
&mut self,
core: usize,
destination: usize,
source: usize,
size: usize,
operation: &'static str,
) {
if let Some(provenance) = &mut self.provenance {
provenance.local_copy(core, destination, source, size, operation);
}
}
pub(crate) fn provenance_local_strided_copy(
&mut self,
core: usize,
destination: usize,
source: usize,
element_size: usize,
stride: usize,
element_count: usize,
operation: &'static str,
) {
if let Some(provenance) = &mut self.provenance {
provenance.local_strided_copy(
core,
destination,
source,
element_size,
stride,
element_count,
operation,
);
}
}
pub(crate) fn provenance_local_transform(
&mut self,
core: usize,
destination: usize,
sources: &[(usize, usize)],
output_size: usize,
operation: &'static str,
) {
if let Some(provenance) = &mut self.provenance {
provenance.local_transform(core, destination, sources, output_size, operation);
}
}
pub(crate) fn provenance_local_broadcast_transform(
&mut self,
core: usize,
destination: usize,
source: usize,
source_size: usize,
output_size: usize,
operation: &'static str,
) {
if let Some(provenance) = &mut self.provenance {
provenance.local_broadcast_transform(
core,
destination,
source,
source_size,
output_size,
operation,
);
}
}
pub(crate) fn provenance_local_mvm_transform(
&mut self,
core: usize,
destination: usize,
source: usize,
element_size: usize,
output_size: usize,
used_rows: &[bool],
operation: &'static str,
) {
if let Some(provenance) = &mut self.provenance {
provenance.local_mvm_transform(
core,
destination,
source,
element_size,
output_size,
used_rows,
operation,
);
}
}
pub(crate) fn provenance_send_transfer(
&mut self,
sender: usize,
receiver: usize,
source: usize,
destination: usize,
size: usize,
) {
if let Some(provenance) = &mut self.provenance {
provenance.send_transfer(sender, receiver, source, destination, size);
}
}
pub(crate) fn finish_provenance(&mut self) {
if let Some(provenance) = &mut self.provenance {
provenance.flush();
}
}
pub(crate) fn provenance_schedule_stall(&mut self, core: usize, remaining_cycles: u64) {
if let Some(provenance) = &self.provenance {
provenance.schedule_stall(core, remaining_cycles);
}
}
pub(crate) fn provenance_schedule_target_stall(
&mut self,
core: usize,
pc: usize,
remaining_cycles: u64,
) {
if let Some(provenance) = &self.provenance {
provenance.schedule_target_stall(core, pc, remaining_cycles);
}
}
pub(crate) fn provenance_schedule_config(&mut self, config: super::DiagnosticScheduleConfig) {
if let Some(provenance) = &self.provenance {
provenance.schedule_config(
match config.policy {
DiagnosticSchedulePolicy::Greedy => "greedy",
DiagnosticSchedulePolicy::Randomized => "randomized",
DiagnosticSchedulePolicy::Adversarial => "adversarial",
},
config.seed,
config.target.map(|target| {
json!({
"writer_core": target.writer_core,
"writer_pc": target.writer_pc,
"reader_core": target.reader_core,
"reader_pc": target.reader_pc,
"address_begin": target.address_begin,
"address_end": target.address_end,
"reader_iteration": target.reader_iteration,
"writer_min_iteration": target.writer_min_iteration,
})
}),
config.deferral_budget,
config.fixed_stall,
config.fixed_target_stall,
);
}
}
pub(crate) fn provenance_scheduler_event(
&mut self,
event: &'static str,
core: usize,
pc: usize,
iteration: u32,
reason: &'static str,
selected_core: Option<usize>,
target: Option<DiagnosticScheduleTarget>,
deferrals: u64,
) {
if let Some(provenance) = &self.provenance {
let target = target.map(|target| {
json!({
"writer_core": target.writer_core,
"writer_pc": target.writer_pc,
"reader_core": target.reader_core,
"reader_pc": target.reader_pc,
"address_begin": target.address_begin,
"address_end": target.address_end,
"reader_iteration": target.reader_iteration,
"writer_min_iteration": target.writer_min_iteration,
})
});
provenance.scheduler_event(json!({
"event": event,
"cycle": self.provenance_cycle(),
"core": core,
"pc": pc,
"core_iteration": iteration,
"reason": reason,
"selected_core": selected_core,
"target": target,
"deferrals": deferrals,
}));
}
}
fn provenance_cycle(&self) -> u64 {
self.provenance.as_ref().map_or(0, ProvenanceTracker::cycle)
}
pub(crate) fn set_current_iteration(&mut self, iteration: u32) {
if let Some(batch_outputs) = &mut self.batch_outputs {
batch_outputs.iteration = iteration as usize;
}
}
pub(crate) fn begin_host_store_recording(
&mut self,
batch_size: usize,
dump_ranges: &[(usize, usize)],
) -> Result<()> {
let mut initial = Vec::new();
for &(address, size) in dump_ranges {
initial.extend_from_slice(self.host().load::<u8>(address, size)?[0]);
}
self.batch_outputs = Some(BatchOutputs {
iteration: 0,
ranges: dump_ranges.to_vec(),
outputs: vec![initial; batch_size],
});
Ok(())
}
pub(crate) fn store_to_host(
&mut self,
core: impl TryToUsize,
host_address: impl AddressArg,
core_address: impl AddressArg,
size: impl TryToUsize,
) -> Result<()> {
let core = core.try_into().expect("core can not be negative");
let host_address = host_address.to_address_usize()?;
let core_address = core_address.to_address_usize()?;
let size = size.try_into().context("size can not be negative")?;
let Self {
cores,
batch_outputs,
..
} = self;
let (host, cores) = cores.split_at_mut(1);
let bytes = cores[core - 1].load::<u8>(core_address, size)?[0];
host[0].execute_store(host_address, bytes)?;
if let Some(batch_outputs) = batch_outputs {
batch_outputs.record(host_address, bytes);
}
Ok(())
}
pub(crate) fn finish_host_store_recording(&mut self) -> Vec<Vec<u8>> {
self.batch_outputs
.take()
.map_or_else(Vec::new, |batch_outputs| batch_outputs.outputs)
}
pub fn host<'b>(&'b mut self) -> &'b mut Core<'a> pub fn host<'b>(&'b mut self) -> &'b mut Core<'a>
where where
'a: 'b, 'a: 'b,
@@ -1,7 +1,7 @@
use crate::{ use crate::{
cpu::{CPU, crossbar}, cpu::{CPU, crossbar},
instruction_set::{ instruction_set::{
Instruction, InstructionData, InstructionStatus, InstructionType, VectorBitWith, Instruction, InstructionData, InstructionStatus, InstructionType, VectorBitWidth,
helper::add_all, helper::add_all,
}, },
memory_manager::{ memory_manager::{
@@ -200,20 +200,20 @@ pub fn isa_simd(functor: InstructionType) -> bool {
pub fn dispatch_simd( pub fn dispatch_simd(
functor: InstructionType, functor: InstructionType,
vector_bit_with: VectorBitWith, vector_bit_width: VectorBitWidth,
) -> Result<InstructionType> { ) -> Result<InstructionType> {
let VectorBitWith { let VectorBitWidth {
vector_input_bitwith, vector_input_bitwidth,
vector_output_bitwith, vector_output_bitwidth,
} = vector_bit_with; } = vector_bit_width;
let res = SIMD let res = SIMD
.get(&(functor as usize)) .get(&(functor as usize))
.context("Request a non present simd")? .context("Request a non present simd")?
.get(&(vector_input_bitwith, vector_output_bitwith)) .get(&(vector_input_bitwidth, vector_output_bitwidth))
.with_context(|| { .with_context(|| {
format!( format!(
"Function not found for the requested size input:{} output:{}", "Function not found for the requested size input:{} output:{}",
vector_input_bitwith, vector_output_bitwith vector_input_bitwidth, vector_output_bitwidth
) )
})?; })?;
Ok(*res) Ok(*res)
@@ -285,6 +285,10 @@ where
let load = loads[0]; let load = loads[0];
let vec: Cow<[M]> = load.up(); let vec: Cow<[M]> = load.up();
let matrix = crossbar.load::<M>(crossbar_stored_bytes)?[0]; let matrix = crossbar.load::<M>(crossbar_stored_bytes)?[0];
let used_rows: Vec<bool> = matrix
.chunks_exact(crossbar_elem_width)
.map(|row| row.iter().any(|value| *value != M::from_f32(0.0)))
.collect();
// --- FAER IMPLEMENTATION --- // --- FAER IMPLEMENTATION ---
@@ -323,6 +327,16 @@ where
let res_up: Cow<[T]> = res.as_slice().up(); let res_up: Cow<[T]> = res.as_slice().up();
core.execute_store(rd_val, res_up.as_ref()); core.execute_store(rd_val, res_up.as_ref());
let _ = core;
cores.provenance_local_mvm_transform(
core_indx as usize,
rd_val as usize,
r1_val as usize,
size_of::<F>(),
res_up.len() * size_of::<T>(),
&used_rows,
"mvmul",
);
TRACER.lock().unwrap().post_mvm::<F, M, T>(cores, data); TRACER.lock().unwrap().post_mvm::<F, M, T>(cores, data);
Ok(InstructionStatus::Completed) Ok(InstructionStatus::Completed)
@@ -389,6 +403,14 @@ where
); );
let res_up: Cow<[T]> = res.as_slice().up(); let res_up: Cow<[T]> = res.as_slice().up();
core.execute_store(rd_val, res_up.as_ref()); core.execute_store(rd_val, res_up.as_ref());
let _ = core;
cores.provenance_local_transform(
core_indx as usize,
rd_val,
&[(r1_val, byte_len), (r2_val, byte_len)],
byte_len,
"vvadd",
);
TRACER.lock().unwrap().post_vvadd::<F, T>(cores, data); TRACER.lock().unwrap().post_vvadd::<F, T>(cores, data);
Ok(InstructionStatus::Completed) Ok(InstructionStatus::Completed)
} }
@@ -474,6 +496,13 @@ where
); );
let res_up: Cow<[T]> = res.as_slice().up(); let res_up: Cow<[T]> = res.as_slice().up();
core.execute_store(rd_val, res_up.as_ref()); core.execute_store(rd_val, res_up.as_ref());
cores.provenance_local_transform(
core_indx as usize,
rd_val,
&[(r1_val, byte_len), (r2_val, byte_len)],
byte_len,
"vvmul",
);
Ok(InstructionStatus::Completed) Ok(InstructionStatus::Completed)
} }
@@ -780,6 +809,15 @@ where
); );
} }
core.execute_store(destination, &result)?; core.execute_store(destination, &result)?;
cores.provenance_local_strided_copy(
core_indx as usize,
destination as usize,
source,
size_of::<F>(),
stride,
element_count,
"vmv",
);
Ok(InstructionStatus::Completed) Ok(InstructionStatus::Completed)
} }
@@ -799,16 +837,23 @@ pub fn vrsl(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
#[inline(never)] #[inline(never)]
pub fn ld(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> { pub fn ld(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
TRACER.lock().unwrap().pre_ld(cores, data); TRACER.lock().unwrap().pre_ld(cores, data);
let (core, rd, r1, _, imm_len, offset_select, offset_value) = let (core_index, rd, r1, _, imm_len, offset_select, offset_value) =
data.get_core_rd_r1_r2_immlen_offset(); data.get_core_rd_r1_r2_immlen_offset();
ensure!(core != 0, "LD cannot be used to move from host to host"); ensure!(
let (host, core) = cores.host_and_cores(core); core_index != 0,
let r1_val = core.register(r1); "LD cannot be used to move from host to host"
let rd_val = core.register(rd); );
let r1_val = add_offset_r1(r1_val, offset_select, offset_value); let (r1_val, rd_val) = {
let rd_val = add_offset_rd(rd_val, offset_select, offset_value); let (host, core) = cores.host_and_cores(core_index);
let global_memory = host.load::<u8>(r1_val, imm_len)?; let r1_val = core.register(r1);
core.execute_store(rd_val, global_memory[0])?; let rd_val = core.register(rd);
let r1_val = add_offset_r1(r1_val, offset_select, offset_value);
let rd_val = add_offset_rd(rd_val, offset_select, offset_value);
let global_memory = host.load::<u8>(r1_val, imm_len)?;
core.execute_store(rd_val, global_memory[0])?;
(r1_val, rd_val)
};
cores.provenance_global_load_to_local(core_index as usize, r1_val, rd_val, imm_len as usize);
TRACER.lock().unwrap().post_ld(cores, data); TRACER.lock().unwrap().post_ld(cores, data);
Ok(InstructionStatus::Completed) Ok(InstructionStatus::Completed)
} }
@@ -819,13 +864,16 @@ pub fn st(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
let (core, rd, r1, _, imm_len, offset_select, offset_value) = let (core, rd, r1, _, imm_len, offset_select, offset_value) =
data.get_core_rd_r1_r2_immlen_offset(); data.get_core_rd_r1_r2_immlen_offset();
ensure!(core != 0, "ST cannot be used to move from host to host"); ensure!(core != 0, "ST cannot be used to move from host to host");
let (host, core) = cores.host_and_cores(core); let (rd_val, r1_val) = {
let r1_val = core.register(r1); let core = cores.core(core);
let rd_val = core.register(rd); let r1_val = core.register(r1);
let r1_val = add_offset_r1(r1_val, offset_select, offset_value); let rd_val = core.register(rd);
let rd_val = add_offset_rd(rd_val, offset_select, offset_value); let r1_val = add_offset_r1(r1_val, offset_select, offset_value);
let local_memory = core.load::<u8>(r1_val, imm_len)?; let rd_val = add_offset_rd(rd_val, offset_select, offset_value);
host.execute_store(rd_val, local_memory[0]); (rd_val, r1_val)
};
cores.store_to_host(core, rd_val, r1_val, imm_len)?;
cores.provenance_global_store_from_local(core as usize, rd_val, r1_val, imm_len as usize);
TRACER.lock().unwrap().post_st(cores, data); TRACER.lock().unwrap().post_st(cores, data);
Ok(InstructionStatus::Completed) Ok(InstructionStatus::Completed)
} }
@@ -850,9 +898,9 @@ pub fn lldi(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
#[inline(never)] #[inline(never)]
pub fn lmv(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> { pub fn lmv(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
TRACER.lock().unwrap().pre_lmv(cores, data); TRACER.lock().unwrap().pre_lmv(cores, data);
let (core, rd, r1, _, imm_len, offset_select, offset_value) = let (core_index, rd, r1, _, imm_len, offset_select, offset_value) =
data.get_core_rd_r1_r2_immlen_offset(); data.get_core_rd_r1_r2_immlen_offset();
let core = cores.core(core); let core = cores.core(core_index);
let r1_val = core.register(r1); let r1_val = core.register(r1);
let rd_val = core.register(rd); let rd_val = core.register(rd);
let r1_val = add_offset_r1(r1_val, offset_select, offset_value); let r1_val = add_offset_r1(r1_val, offset_select, offset_value);
@@ -860,6 +908,8 @@ pub fn lmv(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
let local_memory = core.load::<u8>(r1_val, imm_len)?; let local_memory = core.load::<u8>(r1_val, imm_len)?;
let tmp = local_memory[0].to_vec(); let tmp = local_memory[0].to_vec();
core.execute_store(rd_val, tmp.as_slice()); core.execute_store(rd_val, tmp.as_slice());
let _ = core;
cores.provenance_local_copy(core_index as usize, rd_val, r1_val, imm_len as usize, "lmv");
TRACER.lock().unwrap().post_lmv(cores, data); TRACER.lock().unwrap().post_lmv(cores, data);
Ok(InstructionStatus::Completed) Ok(InstructionStatus::Completed)
} }
@@ -881,7 +931,7 @@ pub fn isa_recv(functor: usize) -> bool {
#[inline(never)] #[inline(never)]
pub fn recv(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> { pub fn recv(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
Ok(InstructionStatus::Reciving(data)) Ok(InstructionStatus::Receiving(data))
} }
#[inline(never)] #[inline(never)]
@@ -22,7 +22,7 @@ pub enum InstructionStatus {
Completed, Completed,
Waiting(InstructionData), Waiting(InstructionData),
Sending(InstructionData), Sending(InstructionData),
Reciving(InstructionData), Receiving(InstructionData),
Sync(InstructionData), Sync(InstructionData),
#[default] #[default]
NotExecuted, NotExecuted,
@@ -59,21 +59,21 @@ pub type Instructions = Vec<Instruction>;
pub type InstructionType = fn(&mut CPU, InstructionData) -> Result<InstructionStatus>; pub type InstructionType = fn(&mut CPU, InstructionData) -> Result<InstructionStatus>;
#[derive(Debug, Clone, Copy, Default)] #[derive(Debug, Clone, Copy, Default)]
pub struct VectorBitWith { pub struct VectorBitWidth {
pub vector_input_bitwith: usize, pub vector_input_bitwidth: usize,
pub vector_output_bitwith: usize, pub vector_output_bitwidth: usize,
} }
/// Support for the /// Support for the
/// setbw ibiw, obiw /// setbw ibiw, obiw
/// Set the bit-widths of each element for input vectors and output vectors. Related vector instructions /// Set the bit-widths of each element for input vectors and output vectors. Related vector instructions
/// use the configured bit-widths. Once setbw is caled, all subsequent related vector instructions will /// use the configured bit-widths. Once setbw is called, all subsequent related vector instructions will
/// use the configured bit-widths, until a new setbw is called. Once ibiw and obiw are set, ibyw and /// use the configured bit-widths, until a new setbw is called. Once ibiw and obiw are set, ibyw and
/// obyw are also set accordingly by the hardware. /// obyw are also set accordingly by the hardware.
/// If the hardware does not support variable bit-width, this instruction is invalid and the matrix/vector /// If the hardware does not support variable bit-width, this instruction is invalid and the matrix/vector
/// instructions use the fixed bit-width of the hardware. /// instructions use the fixed bit-width of the hardware.
pub struct InstructionsBuilder { pub struct InstructionsBuilder {
vector_bit_with: VectorBitWith, vector_bit_width: VectorBitWidth,
instructions: Instructions, instructions: Instructions,
} }
@@ -86,9 +86,9 @@ impl Default for InstructionsBuilder {
impl InstructionsBuilder { impl InstructionsBuilder {
pub fn new() -> Self { pub fn new() -> Self {
Self { Self {
vector_bit_with: VectorBitWith { vector_bit_width: VectorBitWidth {
vector_input_bitwith: 32, vector_input_bitwidth: 32,
vector_output_bitwith: 32, vector_output_bitwidth: 32,
}, },
instructions: Instructions::new(), instructions: Instructions::new(),
} }
@@ -97,9 +97,9 @@ impl InstructionsBuilder {
pub fn make_inst(&mut self, functor: InstructionType, data: InstructionData) { pub fn make_inst(&mut self, functor: InstructionType, data: InstructionData) {
if is_setbw(functor) { if is_setbw(functor) {
let (ibiw, obiw) = data.get_ibiw_obiw(); let (ibiw, obiw) = data.get_ibiw_obiw();
self.vector_bit_with.vector_input_bitwith = self.vector_bit_width.vector_input_bitwidth =
ibiw.try_into().expect("ibiw can not be negative"); ibiw.try_into().expect("ibiw can not be negative");
self.vector_bit_with.vector_output_bitwith = self.vector_bit_width.vector_output_bitwidth =
obiw.try_into().expect("obiw can not be negative"); obiw.try_into().expect("obiw can not be negative");
return; return;
} }
@@ -107,7 +107,7 @@ impl InstructionsBuilder {
if (isa_simd(functor)) { if (isa_simd(functor)) {
self.instructions.push(Instruction::new( self.instructions.push(Instruction::new(
data, data,
dispatch_simd(functor, self.vector_bit_with).unwrap(), dispatch_simd(functor, self.vector_bit_width).unwrap(),
)) ))
} else { } else {
self.instructions.push(Instruction::new(data, functor)) self.instructions.push(Instruction::new(data, functor))
@@ -1,8 +1,14 @@
#![allow(unused)] #![allow(unused)]
use anyhow::{Result, bail}; use anyhow::{Context, Result, bail};
use serde_json::json;
use std::{ use std::{
collections::{HashMap, HashSet}, collections::{HashMap, HashSet},
path::Path,
sync::{
Mutex,
atomic::{AtomicU32, Ordering},
},
time::{Duration, SystemTime}, time::{Duration, SystemTime},
}; };
@@ -21,10 +27,14 @@ pub mod cpu;
pub mod instruction_set; pub mod instruction_set;
pub mod json_to_instruction; pub mod json_to_instruction;
pub mod memory_manager; pub mod memory_manager;
pub mod provenance;
pub mod send_recv; pub mod send_recv;
pub mod tracing; pub mod tracing;
pub mod utility; pub mod utility;
static GLOBAL_ITERATION: AtomicU32 = AtomicU32::new(0);
static EXECUTION_LOCK: Mutex<()> = Mutex::new(());
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct CoreInstructionsBuilder { pub struct CoreInstructionsBuilder {
core_instructions: Vec<CoreInstructions>, core_instructions: Vec<CoreInstructions>,
@@ -54,6 +64,7 @@ impl CoreInstructionsBuilder {
pub struct CoreInstructions { pub struct CoreInstructions {
instructions: Instructions, instructions: Instructions,
program_counter: usize, program_counter: usize,
current_iteration: u32,
} }
impl CoreInstructions { impl CoreInstructions {
@@ -61,6 +72,7 @@ impl CoreInstructions {
Self { Self {
instructions, instructions,
program_counter, program_counter,
current_iteration: 0,
} }
} }
@@ -68,6 +80,7 @@ impl CoreInstructions {
Self { Self {
instructions: Vec::new(), instructions: Vec::new(),
program_counter: 0, program_counter: 0,
current_iteration: 0,
} }
} }
} }
@@ -77,15 +90,320 @@ impl From<Instructions> for CoreInstructions {
CoreInstructions { CoreInstructions {
instructions: value, instructions: value,
program_counter: 0, program_counter: 0,
current_iteration: 0,
} }
} }
} }
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub enum DiagnosticSchedulePolicy {
#[default]
Greedy,
Randomized,
Adversarial,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct DiagnosticScheduleTarget {
pub writer_core: usize,
pub writer_pc: usize,
pub reader_core: usize,
pub reader_pc: usize,
pub address_begin: usize,
pub address_end: usize,
pub reader_iteration: Option<u32>,
pub writer_min_iteration: Option<u32>,
}
#[derive(Clone, Copy, Debug)]
struct DiagnosticScheduleConfig {
policy: DiagnosticSchedulePolicy,
seed: u64,
target: Option<DiagnosticScheduleTarget>,
deferral_budget: u64,
fixed_stall: Option<(usize, u64)>,
fixed_target_stall: Option<(usize, usize, Option<u32>, u64)>,
}
impl Default for DiagnosticScheduleConfig {
fn default() -> Self {
Self {
policy: DiagnosticSchedulePolicy::Greedy,
seed: 0,
target: None,
deferral_budget: 10_000,
fixed_stall: None,
fixed_target_stall: None,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct ScheduleCoreState {
program_counter: usize,
instruction_count: usize,
current_iteration: u32,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum ScheduleAction {
Execute,
Defer {
next_core: usize,
reason: &'static str,
},
Stall {
remaining: u64,
target: bool,
},
Force {
reason: &'static str,
},
}
#[derive(Clone, Copy, Debug)]
struct ScheduleChoice {
core: usize,
reason: &'static str,
}
#[derive(Debug)]
struct DiagnosticScheduler {
config: DiagnosticScheduleConfig,
random_state: u64,
deferrals: u64,
writer_store_iteration: Option<u32>,
last_writer_state: Option<ScheduleCoreState>,
writer_stagnation: u64,
}
impl DiagnosticScheduler {
fn new(config: DiagnosticScheduleConfig) -> Self {
let random_state = if config.seed == 0 {
0x9e37_79b9_7f4a_7c15
} else {
config.seed
};
Self {
config,
random_state,
deferrals: 0,
writer_store_iteration: None,
last_writer_state: None,
writer_stagnation: 0,
}
}
fn state(cores: &[CoreInstructions]) -> Vec<ScheduleCoreState> {
cores
.iter()
.map(|core| ScheduleCoreState {
program_counter: core.program_counter,
instruction_count: core.instructions.len(),
current_iteration: core.current_iteration,
})
.collect()
}
fn has_work(core: ScheduleCoreState, batch_size: u32) -> bool {
!((core.instruction_count == 0)
|| (core.program_counter == core.instruction_count
&& core.current_iteration + 1 >= batch_size))
}
fn candidates(states: &[ScheduleCoreState], current: usize, batch_size: u32) -> Vec<usize> {
states
.iter()
.enumerate()
.filter_map(|(index, &core)| {
(index != current && Self::has_work(core, batch_size)).then_some(index)
})
.collect()
}
fn next_random(&mut self) -> u64 {
let mut value = self.random_state;
value ^= value << 13;
value ^= value >> 7;
value ^= value << 17;
self.random_state = value;
value
}
fn random_choice(&mut self, candidates: &[usize]) -> Option<ScheduleChoice> {
(!candidates.is_empty()).then(|| ScheduleChoice {
core: candidates[(self.next_random() as usize) % candidates.len()],
reason: "randomized_ready_core",
})
}
fn target_reader(&self, core: usize, state: ScheduleCoreState) -> bool {
self.config.target.is_some_and(|target| {
target.reader_core == core
&& target.reader_pc == state.program_counter
&& target
.reader_iteration
.is_none_or(|iteration| iteration == state.current_iteration)
})
}
fn target_writer_has_work(&self, states: &[ScheduleCoreState], batch_size: u32) -> bool {
self.config.target.is_some_and(|target| {
states
.get(target.writer_core)
.is_some_and(|&state| Self::has_work(state, batch_size))
})
}
fn reader_may_execute(&self, reader_iteration: u32) -> bool {
let minimum = self
.config
.target
.and_then(|target| target.writer_min_iteration)
.unwrap_or(reader_iteration + 1);
self.writer_store_iteration
.is_some_and(|iteration| iteration >= minimum)
}
fn before_instruction(
&mut self,
current: usize,
states: &[ScheduleCoreState],
batch_size: u32,
) -> ScheduleAction {
if let Some((core, pc, iteration, remaining)) = self.config.fixed_target_stall
&& core == current
&& states.get(current).is_some_and(|state| {
state.program_counter == pc
&& iteration.is_none_or(|iteration| state.current_iteration == iteration)
})
&& remaining > 0
{
self.config.fixed_target_stall = Some((core, pc, iteration, remaining - 1));
return ScheduleAction::Stall {
remaining: remaining - 1,
target: true,
};
}
if let Some((core, remaining)) = self.config.fixed_stall
&& core == current
&& remaining > 0
{
self.config.fixed_stall = Some((core, remaining - 1));
return ScheduleAction::Stall {
remaining: remaining - 1,
target: false,
};
}
let Some(state) = states.get(current).copied() else {
return ScheduleAction::Execute;
};
let candidates = Self::candidates(states, current, batch_size);
match self.config.policy {
DiagnosticSchedulePolicy::Greedy => ScheduleAction::Execute,
// Random choices are made after a blocking/ready boundary in
// `after_block`. Deferring here would allow two ready cores to
// defer each other forever without executing an instruction.
DiagnosticSchedulePolicy::Randomized => ScheduleAction::Execute,
DiagnosticSchedulePolicy::Adversarial => {
let Some(target) = self.config.target else {
return ScheduleAction::Execute;
};
if !self.target_reader(current, state)
|| self.reader_may_execute(state.current_iteration)
{
return ScheduleAction::Execute;
}
if self.deferrals >= self.config.deferral_budget {
return ScheduleAction::Force {
reason: "DEFERRAL_LIMIT_REACHED",
};
}
let writer_state = states.get(target.writer_core).copied();
let next_core = if Self::has_work(writer_state.unwrap_or(state), batch_size) {
if self.last_writer_state == writer_state {
self.writer_stagnation += 1;
} else {
self.last_writer_state = writer_state;
self.writer_stagnation = 0;
}
if self.writer_stagnation >= 256 {
return ScheduleAction::Force {
reason: "PRODUCER_BLOCKED_BY_REAL_DEPENDENCY",
};
}
target.writer_core
} else {
candidates.first().copied().unwrap_or(current)
};
if next_core == current {
ScheduleAction::Force {
reason: "NO_ALTERNATIVE_READY_EVENT",
}
} else {
self.deferrals += 1;
ScheduleAction::Defer {
next_core,
reason: "target_consumer",
}
}
}
}
}
fn after_block(
&mut self,
current: usize,
states: &[ScheduleCoreState],
batch_size: u32,
) -> Option<ScheduleChoice> {
let candidates = Self::candidates(states, current, batch_size);
match self.config.policy {
DiagnosticSchedulePolicy::Greedy => None,
DiagnosticSchedulePolicy::Randomized => self.random_choice(&candidates),
DiagnosticSchedulePolicy::Adversarial => {
let target = self.config.target?;
if self.target_writer_has_work(states, batch_size)
&& target.writer_core != current
&& candidates.contains(&target.writer_core)
{
Some(ScheduleChoice {
core: target.writer_core,
reason: "target_producer",
})
} else {
candidates.first().copied().map(|core| ScheduleChoice {
core,
reason: "adversarial_ready_core",
})
}
}
}
}
fn note_completed(&mut self, core: usize, pc: usize, iteration: u32) {
if self
.config
.target
.is_some_and(|target| target.writer_core == core && target.writer_pc == pc)
{
self.writer_store_iteration = Some(iteration);
}
}
fn config(&self) -> DiagnosticScheduleConfig {
self.config
}
}
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct Executable<'a> { pub struct Executable<'a> {
cpu: CPU<'a>, cpu: CPU<'a>,
core_instructions: Vec<CoreInstructions>, core_instructions: Vec<CoreInstructions>,
send_recv: SendRecv, send_recv: SendRecv,
provenance_global_barrier: bool,
diagnostic_schedule: DiagnosticScheduleConfig,
} }
struct DeadlockInfo { struct DeadlockInfo {
@@ -123,63 +441,321 @@ impl<'a> Executable<'a> {
cpu, cpu,
core_instructions, core_instructions,
send_recv, send_recv,
provenance_global_barrier: false,
diagnostic_schedule: DiagnosticScheduleConfig::default(),
} }
} }
pub fn enable_provenance(&mut self, path: impl AsRef<Path>) -> Result<()> {
self.cpu
.enable_provenance(path)
.context("cannot enable provenance tracing")?;
Ok(())
}
pub fn set_provenance_global_barrier(&mut self, enabled: bool) {
self.provenance_global_barrier = enabled;
}
pub fn set_provenance_core_stall(&mut self, core: usize, cycles: u64) {
self.diagnostic_schedule.fixed_stall = Some((core, cycles));
}
pub fn set_diagnostic_target_stall(
&mut self,
core: usize,
pc: usize,
iteration: Option<u32>,
cycles: u64,
) {
self.diagnostic_schedule.fixed_target_stall = Some((core, pc, iteration, cycles));
}
pub fn set_diagnostic_schedule_policy(&mut self, policy: DiagnosticSchedulePolicy) {
self.diagnostic_schedule.policy = policy;
}
pub fn set_diagnostic_schedule_seed(&mut self, seed: u64) {
self.diagnostic_schedule.seed = seed;
}
pub fn set_diagnostic_schedule_target(&mut self, target: DiagnosticScheduleTarget) {
self.diagnostic_schedule.target = Some(target);
}
pub fn set_diagnostic_schedule_deferral_budget(&mut self, budget: u64) {
self.diagnostic_schedule.deferral_budget = budget;
}
pub fn execute<'b>(&'b mut self) -> Result<()> pub fn execute<'b>(&'b mut self) -> Result<()>
where where
'a: 'b, 'a: 'b,
{ {
self.execute_batch(&[&[]], &[], &[]).map(|_| ())
}
pub fn execute_batch<'b>(
&'b mut self,
inputs: &[&[u8]],
input_regions: &[(usize, usize)],
dump_ranges: &[(usize, usize)],
) -> Result<Vec<Vec<u8>>>
where
'a: 'b,
{
validate_inputs(inputs, input_regions)?;
self.execute_iterations(inputs, input_regions, dump_ranges)
}
fn execute_iterations<'b>(
&'b mut self,
inputs: &[&[u8]],
input_regions: &[(usize, usize)],
dump_ranges: &[(usize, usize)],
) -> Result<Vec<Vec<u8>>>
where
'a: 'b,
{
let _execution_lock = EXECUTION_LOCK.lock().unwrap();
let batch_size = u32::try_from(inputs.len().max(1)).context("batch size exceeds u32")?;
GLOBAL_ITERATION.store(0, Ordering::SeqCst);
self.cpu.begin_provenance_batch(batch_size as usize);
if let Some(input) = inputs.first() {
store_input(&mut self.cpu, input, input_regions, 0)?;
}
self.cpu
.begin_host_store_recording(batch_size as usize, dump_ranges)?;
let provenance_global_barrier = self.provenance_global_barrier;
let mut scheduler = DiagnosticScheduler::new(self.diagnostic_schedule);
self.cpu.provenance_schedule_config(scheduler.config());
let Self { let Self {
cpu, cpu,
core_instructions: cores_instructions, core_instructions: cores_instructions,
send_recv, send_recv,
..
} = self; } = self;
let active_cores: Vec<usize> = cores_instructions
.iter()
.enumerate()
.filter_map(|(index, core)| (!core.instructions.is_empty()).then_some(index))
.collect();
let mut barrier_iteration = None;
let mut cpu_progressed = 0; let mut cpu_progressed = 0;
let max_core = cpu.num_core(); let max_core = cpu.num_core();
let mut sync_events: SyncEvents = vec![[0; 32]; max_core]; let mut sync_events: SyncEvents = vec![[0; 32]; max_core];
let mut cpu_index = 0; let mut cpu_index = 0;
let mut cycle = 0;
let mut scheduler_no_progress = 0usize;
let scheduler_no_progress_limit = max_core.saturating_mul(4).max(8);
let mut now = SystemTime::now(); let mut now = SystemTime::now();
while (cpu_progressed > -2) { while (cpu_progressed > -2) {
let mut core_result = InstructionStatus::Completed; let mut core_result = InstructionStatus::Completed;
while core_result.is_completed() let mut scheduler_next = None;
&& let Some(core_instruction) = cores_instructions.get_mut(cpu_index) if provenance_global_barrier
&& barrier_iteration.is_some()
&& active_cores.iter().all(|&index| {
cores_instructions[index].current_iteration >= barrier_iteration.unwrap()
})
{ {
core_result = InstructionStatus::NotExecuted; barrier_iteration = None;
let CoreInstructions { }
instructions, while core_result.is_completed() {
program_counter, let barrier_ready = if provenance_global_barrier {
} = core_instruction; let current_iteration = cores_instructions[cpu_index].current_iteration;
core_result = instructions active_cores.iter().all(|&index| {
.get(*program_counter) let core = &cores_instructions[index];
.map_or(InstructionStatus::default(), |inst: &Instruction| { core.program_counter == core.instructions.len()
inst.execute(cpu) && core.current_iteration == current_iteration
}); })
if core_result.is_completed() { } else {
cpu_progressed = 0; false
*program_counter += 1; };
let current_pc = cores_instructions[cpu_index].program_counter;
let current_iteration = cores_instructions[cpu_index].current_iteration;
let states = if scheduler.config().policy == DiagnosticSchedulePolicy::Greedy
&& scheduler.config().fixed_stall.is_none()
&& scheduler.config().fixed_target_stall.is_none()
{
None
} else {
Some(DiagnosticScheduler::state(cores_instructions))
};
let schedule_action = states.as_deref().map_or(ScheduleAction::Execute, |states| {
scheduler.before_instruction(cpu_index, states, batch_size)
});
if let ScheduleAction::Defer { next_core, reason } = schedule_action {
cpu.set_execution_context(cycle, cpu_index, current_pc, current_iteration);
cpu.provenance_scheduler_event(
"scheduler_defer",
cpu_index,
current_pc,
current_iteration,
reason,
Some(next_core),
scheduler.config().target,
scheduler.deferrals,
);
scheduler_next = Some(next_core);
break;
} }
if (now.elapsed().unwrap() > Duration::from_secs(5)) { if let ScheduleAction::Stall { remaining, target } = schedule_action {
print_status(cores_instructions); cpu_progressed = 0;
if let Some(deadlock) = detect_deadlock(cores_instructions) { cpu.set_execution_context(cycle, cpu_index, current_pc, current_iteration);
bail!( if target {
"Deadlock cycle detected: {} [{}]", cpu.provenance_schedule_target_stall(cpu_index, current_pc, remaining);
deadlock.cycle, } else {
deadlock.states cpu.provenance_schedule_stall(cpu_index, remaining);
); }
break;
}
if let ScheduleAction::Force { reason } = schedule_action {
cpu.set_execution_context(cycle, cpu_index, current_pc, current_iteration);
cpu.provenance_scheduler_event(
"scheduler_force",
cpu_index,
current_pc,
current_iteration,
reason,
None,
scheduler.config().target,
scheduler.deferrals,
);
}
let Some(core_instruction) = cores_instructions.get_mut(cpu_index) else {
break;
};
core_result = InstructionStatus::NotExecuted;
if core_instruction.program_counter == core_instruction.instructions.len() {
if core_instruction.instructions.is_empty()
|| core_instruction.current_iteration + 1 >= batch_size
{
break;
}
let next_iteration = core_instruction.current_iteration + 1;
if provenance_global_barrier {
if barrier_iteration != Some(next_iteration) {
if !barrier_ready {
break;
}
barrier_iteration = Some(next_iteration);
}
}
core_instruction.current_iteration += 1;
core_instruction.program_counter = 0;
let iteration = core_instruction.current_iteration;
if iteration > GLOBAL_ITERATION.fetch_max(iteration, Ordering::SeqCst) {
cpu.set_execution_context(cycle, cpu_index, 0, iteration);
store_input(cpu, inputs[iteration as usize], input_regions, iteration)?;
}
}
if !matches!(
schedule_action,
ScheduleAction::Stall { .. } | ScheduleAction::Defer { .. }
) {
cpu.set_current_iteration(core_instruction.current_iteration);
let CoreInstructions {
instructions,
program_counter,
..
} = core_instruction;
cpu.set_execution_context(
cycle,
cpu_index,
*program_counter,
core_instruction.current_iteration,
);
cycle += 1;
core_result = instructions
.get(*program_counter)
.map_or(InstructionStatus::default(), |inst: &Instruction| {
inst.execute(cpu)
});
if core_result.is_completed() {
scheduler.note_completed(
cpu_index,
*program_counter,
core_instruction.current_iteration,
);
cpu_progressed = 0;
scheduler_no_progress = 0;
*program_counter += 1;
}
if (now.elapsed().unwrap() > Duration::from_secs(5)) {
print_status(cores_instructions);
if let Some(deadlock) = detect_deadlock(cores_instructions) {
bail!(
"Deadlock cycle detected: {} [{}]",
deadlock.cycle,
deadlock.states
);
}
now = SystemTime::now();
} }
now = SystemTime::now();
} }
} }
if handle_wait_sync(cores_instructions, &mut sync_events, core_result) { if handle_wait_sync(cores_instructions, &mut sync_events, core_result) {
cpu_progressed = 0; cpu_progressed = 0;
scheduler_no_progress = 0;
} }
match handle_send_recv(cpu, cores_instructions, send_recv, core_result) { if let Some(next_core) = scheduler_next {
(true, other_cpu_index) => { cpu_index = next_core;
cpu_progressed = 0; continue;
cpu_index = other_cpu_index; }
} let send_recv_result =
handle_send_recv(cpu, cores_instructions, send_recv, core_result);
if let (true, other_cpu_index) = send_recv_result {
cpu_progressed = 0;
scheduler_no_progress = 0;
cpu_index = other_cpu_index;
continue;
}
if !core_result.is_completed() {
scheduler_no_progress += 1;
}
let states = if scheduler.config().policy == DiagnosticSchedulePolicy::Greedy {
None
} else {
Some(DiagnosticScheduler::state(cores_instructions))
};
let scheduler_choice = (scheduler_no_progress <= scheduler_no_progress_limit)
.then(|| {
states
.as_deref()
.and_then(|states| scheduler.after_block(cpu_index, states, batch_size))
})
.flatten();
if let Some(choice) = scheduler_choice {
cpu.provenance_scheduler_event(
"scheduler_prefer",
cpu_index,
cores_instructions[cpu_index].program_counter,
cores_instructions[cpu_index].current_iteration,
choice.reason,
Some(choice.core),
scheduler.config().target,
scheduler.deferrals,
);
cpu_index = choice.core;
continue;
}
if scheduler_no_progress == scheduler_no_progress_limit + 1
&& scheduler.config().policy != DiagnosticSchedulePolicy::Greedy
{
cpu.provenance_scheduler_event(
"scheduler_force",
cpu_index,
cores_instructions[cpu_index].program_counter,
cores_instructions[cpu_index].current_iteration,
"NO_ALTERNATIVE_READY_EVENT",
None,
scheduler.config().target,
scheduler.deferrals,
);
}
match send_recv_result {
(true, _) => unreachable!("completed SEND/RECV was handled above"),
(false, 0) => { (false, 0) => {
cpu_index = if cpu_index + 1 >= cores_instructions.len() { cpu_index = if cpu_index + 1 >= cores_instructions.len() {
cpu_progressed -= 1; cpu_progressed -= 1;
@@ -211,7 +787,8 @@ impl<'a> Executable<'a> {
#[cfg(feature = "profile_time")] #[cfg(feature = "profile_time")]
TRACER.lock().unwrap().report(); TRACER.lock().unwrap().report();
Ok(()) cpu.finish_provenance();
Ok(cpu.finish_host_store_recording())
} }
pub fn cpu(&self) -> &CPU<'a> { pub fn cpu(&self) -> &CPU<'a> {
@@ -225,7 +802,7 @@ impl<'a> Executable<'a> {
pub fn dump(&self) { pub fn dump(&self) {
let core_instructions = &self.core_instructions; let core_instructions = &self.core_instructions;
for (i, core_instruction) in core_instructions.iter().enumerate() { for (i, core_instruction) in core_instructions.iter().enumerate() {
eprintln!("INST OF CORE {}:", i); eprintln!("Instructions for core {}:", i);
for inst in &core_instruction.instructions { for inst in &core_instruction.instructions {
inst.dump(); inst.dump();
} }
@@ -233,6 +810,35 @@ impl<'a> Executable<'a> {
} }
} }
fn validate_inputs(inputs: &[&[u8]], input_regions: &[(usize, usize)]) -> Result<()> {
let input_size = input_regions.iter().try_fold(0usize, |total, (_, size)| {
total.checked_add(*size).context("input size overflow")
})?;
if inputs.is_empty() {
bail!("at least one input is required");
}
if inputs.iter().any(|input| input.len() != input_size) {
bail!("each input must contain exactly {input_size} bytes");
}
Ok(())
}
fn store_input(
cpu: &mut CPU,
input: &[u8],
input_regions: &[(usize, usize)],
sample: u32,
) -> Result<()> {
let mut offset = 0;
for &(address, size) in input_regions {
cpu.host()
.execute_store(address, &input[offset..offset + size])?;
cpu.provenance_input_store(address, size, sample);
offset += size;
}
Ok(())
}
fn detect_deadlock(cores_instructions: &[CoreInstructions]) -> Option<DeadlockInfo> { fn detect_deadlock(cores_instructions: &[CoreInstructions]) -> Option<DeadlockInfo> {
#[derive(Debug, PartialEq, Eq)] #[derive(Debug, PartialEq, Eq)]
enum CoreState { enum CoreState {
@@ -382,3 +988,143 @@ fn handle_wait_sync(
_ => false, _ => false,
} }
} }
#[cfg(test)]
mod scheduler_tests {
use super::*;
fn target() -> DiagnosticScheduleTarget {
DiagnosticScheduleTarget {
writer_core: 0,
writer_pc: 3,
reader_core: 1,
reader_pc: 2,
address_begin: 100,
address_end: 200,
reader_iteration: Some(0),
writer_min_iteration: Some(1),
}
}
fn states() -> Vec<ScheduleCoreState> {
vec![
ScheduleCoreState {
program_counter: 3,
instruction_count: 8,
current_iteration: 1,
},
ScheduleCoreState {
program_counter: 2,
instruction_count: 8,
current_iteration: 0,
},
]
}
#[test]
fn adversarial_policy_defers_reader_until_writer_store() {
let mut scheduler = DiagnosticScheduler::new(DiagnosticScheduleConfig {
policy: DiagnosticSchedulePolicy::Adversarial,
target: Some(target()),
deferral_budget: 4,
..Default::default()
});
assert_eq!(
scheduler.before_instruction(1, &states(), 4),
ScheduleAction::Defer {
next_core: 0,
reason: "target_consumer"
}
);
scheduler.note_completed(0, 3, 1);
assert_eq!(
scheduler.before_instruction(1, &states(), 4),
ScheduleAction::Execute
);
}
#[test]
fn randomized_policy_is_deterministic_for_fixed_seed() {
let config = DiagnosticScheduleConfig {
policy: DiagnosticSchedulePolicy::Randomized,
seed: 17,
..Default::default()
};
let mut left = DiagnosticScheduler::new(config);
let mut right = DiagnosticScheduler::new(config);
let states = vec![
ScheduleCoreState {
program_counter: 0,
instruction_count: 4,
current_iteration: 0,
},
ScheduleCoreState {
program_counter: 1,
instruction_count: 4,
current_iteration: 0,
},
ScheduleCoreState {
program_counter: 2,
instruction_count: 4,
current_iteration: 0,
},
];
for current in [0, 1, 2, 0, 1] {
assert_eq!(
left.before_instruction(current, &states, 2),
right.before_instruction(current, &states, 2)
);
}
}
#[test]
fn target_stall_is_consumed_without_changing_program_order() {
let mut scheduler = DiagnosticScheduler::new(DiagnosticScheduleConfig {
fixed_target_stall: Some((1, 2, None, 2)),
..Default::default()
});
let states = states();
assert_eq!(
scheduler.before_instruction(1, &states, 4),
ScheduleAction::Stall {
remaining: 1,
target: true
}
);
assert_eq!(
scheduler.before_instruction(1, &states, 4),
ScheduleAction::Stall {
remaining: 0,
target: true
}
);
assert_eq!(
scheduler.before_instruction(1, &states, 4),
ScheduleAction::Execute
);
assert_eq!(states[1].program_counter, 2);
}
#[test]
fn adversarial_policy_releases_reader_when_writer_state_stagnates() {
let mut scheduler = DiagnosticScheduler::new(DiagnosticScheduleConfig {
policy: DiagnosticSchedulePolicy::Adversarial,
target: Some(target()),
deferral_budget: 10_000,
..Default::default()
});
let states = states();
for _ in 0..256 {
assert!(matches!(
scheduler.before_instruction(1, &states, 4),
ScheduleAction::Defer { .. }
));
}
assert_eq!(
scheduler.before_instruction(1, &states, 4),
ScheduleAction::Force {
reason: "PRODUCER_BLOCKED_BY_REAL_DEPENDENCY"
}
);
}
}
@@ -0,0 +1,639 @@
use serde_json::{Value, json};
use std::{
collections::BTreeSet,
fs::File,
io::{BufWriter, Write},
path::Path,
sync::{Arc, Mutex},
};
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, PartialOrd, Ord)]
pub enum Provenance {
#[default]
Uninitialized,
Unknown,
Samples(u64),
}
impl Provenance {
pub fn sample(sample: u32) -> Self {
if sample < 64 {
Self::Samples(1_u64 << sample)
} else {
Self::Unknown
}
}
fn merge(self, other: Self) -> Self {
match (self, other) {
(Self::Unknown, _) | (_, Self::Unknown) => Self::Unknown,
(Self::Uninitialized, value) | (value, Self::Uninitialized) => value,
(Self::Samples(left), Self::Samples(right)) => Self::Samples(left | right),
}
}
fn merge_all(values: impl IntoIterator<Item = Self>) -> Self {
let mut values = values.into_iter();
values
.next()
.map_or(Self::Uninitialized, |first| values.fold(first, Self::merge))
}
fn samples(self) -> Vec<u32> {
match self {
Self::Samples(mask) => (0..64)
.filter(|sample| mask & (1_u64 << sample) != 0)
.collect(),
Self::Unknown | Self::Uninitialized => Vec::new(),
}
}
fn state(self) -> &'static str {
match self {
Self::Samples(_) => "known",
Self::Unknown => "unknown",
Self::Uninitialized => "uninitialized",
}
}
fn is_mixed(self) -> bool {
matches!(self, Self::Samples(mask) if mask.count_ones() > 1)
}
fn json(self) -> Value {
json!({
"provenance": self.samples(),
"provenance_state": self.state(),
"mixed": self.is_mixed(),
})
}
}
#[derive(Clone, Copy, Debug)]
struct ExecutionContext {
cycle: u64,
core: usize,
pc: usize,
iteration: u32,
}
impl Default for ExecutionContext {
fn default() -> Self {
Self {
cycle: 0,
core: 0,
pc: 0,
iteration: 0,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord)]
struct Writer {
version: u64,
cycle: u64,
core: usize,
pc: usize,
iteration: u32,
provenance: Provenance,
}
#[derive(Clone, Copy, Debug, Default)]
struct GlobalCell {
provenance: Provenance,
writer: Option<Writer>,
}
#[derive(Debug)]
struct TraceSink {
output: BufWriter<File>,
}
impl TraceSink {
fn new(path: impl AsRef<Path>) -> std::io::Result<Self> {
let file = File::create(path)?;
Ok(Self {
output: BufWriter::new(file),
})
}
fn event(&mut self, value: Value) {
serde_json::to_writer(&mut self.output, &value).expect("write provenance event");
self.output
.write_all(b"\n")
.expect("write provenance newline");
}
fn flush(&mut self) {
self.output.flush().expect("flush provenance trace");
}
}
#[derive(Debug, Clone)]
pub struct ProvenanceTracker {
global: Vec<GlobalCell>,
local: Vec<Vec<Provenance>>,
next_version: u64,
context: ExecutionContext,
sink: Arc<Mutex<TraceSink>>,
}
impl ProvenanceTracker {
pub fn new(core_count: usize, path: impl AsRef<Path>) -> std::io::Result<Self> {
if let Some(parent) = path.as_ref().parent() {
std::fs::create_dir_all(parent)?;
}
let sink = Arc::new(Mutex::new(TraceSink::new(path)?));
let tracker = Self {
global: Vec::new(),
local: vec![Vec::new(); core_count],
next_version: 0,
context: ExecutionContext::default(),
sink,
};
tracker.write(json!({
"event": "provenance_trace_start",
"schema": 1,
"core_count": core_count,
}));
Ok(tracker)
}
pub fn begin_batch(&mut self, batch_size: usize) {
self.global.fill(GlobalCell::default());
for memory in &mut self.local {
memory.fill(Provenance::Uninitialized);
}
self.next_version = 0;
self.write(json!({
"event": "batch_start",
"batch_size": batch_size,
}));
}
pub fn set_context(&mut self, cycle: u64, core: usize, pc: usize, iteration: u32) {
self.context = ExecutionContext {
cycle,
core,
pc,
iteration,
};
}
pub fn cycle(&self) -> u64 {
self.context.cycle
}
pub fn flush(&mut self) {
self.sink.lock().unwrap().flush();
}
pub fn schedule_stall(&self, core: usize, remaining_cycles: u64) {
self.write(json!({
"event": "diagnostic_core_stall",
"cycle": self.context.cycle,
"core": core,
"pc": self.context.pc,
"core_iteration": self.context.iteration,
"remaining_cycles": remaining_cycles,
}));
}
pub fn schedule_target_stall(&self, core: usize, pc: usize, remaining_cycles: u64) {
self.write(json!({
"event": "diagnostic_target_stall",
"cycle": self.context.cycle,
"core": core,
"pc": pc,
"core_iteration": self.context.iteration,
"remaining_cycles": remaining_cycles,
}));
}
pub fn schedule_config(
&self,
policy: &'static str,
seed: u64,
target: Option<Value>,
deferral_budget: u64,
fixed_stall: Option<(usize, u64)>,
fixed_target_stall: Option<(usize, usize, Option<u32>, u64)>,
) {
self.write(json!({
"event": "scheduler_config",
"schedule_policy": policy,
"schedule_seed": seed,
"target_dependency": target,
"deferral_budget": deferral_budget,
"fixed_stall": fixed_stall.map(|(core, cycles)| json!({
"core": core,
"cycles": cycles,
})),
"fixed_target_stall": fixed_target_stall.map(|(core, pc, iteration, cycles)| json!({
"core": core,
"pc": pc,
"iteration": iteration,
"cycles": cycles,
})),
}));
}
pub fn scheduler_event(&self, event: Value) {
self.write(event);
}
fn write(&self, value: Value) {
self.sink.lock().unwrap().event(value);
}
fn ensure_global(&mut self, end: usize) {
if self.global.len() < end {
self.global.resize(end, GlobalCell::default());
}
}
fn ensure_local(&mut self, core: usize, end: usize) {
if let Some(memory) = self.local.get_mut(core)
&& memory.len() < end
{
memory.resize(end, Provenance::Uninitialized);
}
}
fn local_tags(&mut self, core: usize, address: usize, size: usize) -> Vec<Provenance> {
let Some(end) = address.checked_add(size) else {
return vec![Provenance::Unknown; size];
};
self.ensure_local(core, end);
self.local[core][address..end].to_vec()
}
fn store_local_tags(&mut self, core: usize, address: usize, tags: &[Provenance]) {
let Some(end) = address.checked_add(tags.len()) else {
return;
};
self.ensure_local(core, end);
self.local[core][address..end].copy_from_slice(tags);
}
fn unique_versions(cells: &[GlobalCell]) -> Vec<u64> {
cells
.iter()
.filter_map(|cell| cell.writer.map(|writer| writer.version))
.collect::<BTreeSet<_>>()
.into_iter()
.collect()
}
fn unique_writers(cells: &[GlobalCell]) -> Vec<Writer> {
cells
.iter()
.filter_map(|cell| cell.writer)
.collect::<BTreeSet<_>>()
.into_iter()
.collect()
}
fn writer_json(writer: Writer) -> Value {
json!({
"version": writer.version,
"cycle": writer.cycle,
"core": writer.core,
"pc": writer.pc,
"core_iteration": writer.iteration,
"provenance": writer.provenance.samples(),
"provenance_state": writer.provenance.state(),
})
}
pub fn input_store(&mut self, address: usize, size: usize, sample: u32) {
let Some(end) = address.checked_add(size) else {
return;
};
self.ensure_global(end);
let provenance = Provenance::sample(sample);
self.next_version += 1;
let writer = Writer {
version: self.next_version,
cycle: self.context.cycle,
core: 0,
pc: 0,
iteration: sample,
provenance,
};
let overwritten_versions = Self::unique_versions(&self.global[address..end]);
for cell in &mut self.global[address..end] {
*cell = GlobalCell {
provenance,
writer: Some(writer),
};
}
let mut event = json!({
"event": "external_input_store",
"cycle": self.context.cycle,
"core": 0,
"pc": 0,
"core_iteration": sample,
"address": address,
"size": size,
"sample": sample,
"version": writer.version,
"overwritten_versions": overwritten_versions,
});
if let Some(object) = event.as_object_mut() {
object.extend(provenance.json().as_object().unwrap().clone());
}
self.write(event);
}
pub fn global_store_from_local(
&mut self,
core: usize,
global_address: usize,
local_address: usize,
size: usize,
) {
let Some(end) = global_address.checked_add(size) else {
return;
};
let tags = self.local_tags(core, local_address, size);
self.ensure_global(end);
self.next_version += 1;
let provenance = Provenance::merge_all(tags.iter().copied());
let writer = Writer {
version: self.next_version,
cycle: self.context.cycle,
core,
pc: self.context.pc,
iteration: self.context.iteration,
provenance,
};
let overwritten_versions = Self::unique_versions(&self.global[global_address..end]);
let overwritten_writers = Self::unique_writers(&self.global[global_address..end]);
for (cell, tag) in self.global[global_address..end].iter_mut().zip(tags) {
*cell = GlobalCell {
provenance: tag,
writer: Some(writer),
};
}
let mut event = json!({
"event": "global_store",
"cycle": self.context.cycle,
"core": core,
"pc": self.context.pc,
"core_iteration": self.context.iteration,
"address": global_address,
"local_address": local_address,
"size": size,
"version": writer.version,
"overwritten_versions": overwritten_versions,
"overwritten_writers": overwritten_writers.into_iter().map(Self::writer_json).collect::<Vec<_>>(),
});
if let Some(object) = event.as_object_mut() {
object.extend(provenance.json().as_object().unwrap().clone());
}
self.write(event);
}
pub fn global_load_to_local(
&mut self,
core: usize,
global_address: usize,
local_address: usize,
size: usize,
) {
let Some(end) = global_address.checked_add(size) else {
return;
};
self.ensure_global(end);
let cells = self.global[global_address..end].to_vec();
let tags: Vec<_> = cells.iter().map(|cell| cell.provenance).collect();
let provenance = Provenance::merge_all(tags.iter().copied());
let writers = Self::unique_writers(&cells);
let versions = Self::unique_versions(&cells);
self.store_local_tags(core, local_address, &tags);
let mut event = json!({
"event": "global_load",
"cycle": self.context.cycle,
"core": core,
"pc": self.context.pc,
"core_iteration": self.context.iteration,
"address": global_address,
"local_address": local_address,
"size": size,
"versions": versions,
"last_writers": writers.into_iter().map(Self::writer_json).collect::<Vec<_>>(),
});
if let Some(object) = event.as_object_mut() {
object.extend(provenance.json().as_object().unwrap().clone());
}
self.write(event);
}
pub fn local_copy(
&mut self,
core: usize,
destination: usize,
source: usize,
size: usize,
operation: &'static str,
) {
let tags = self.local_tags(core, source, size);
let provenance = Provenance::merge_all(tags.iter().copied());
self.store_local_tags(core, destination, &tags);
self.local_event(
operation,
core,
destination,
size,
provenance,
&[provenance],
);
}
pub fn local_strided_copy(
&mut self,
core: usize,
destination: usize,
source: usize,
element_size: usize,
stride: usize,
element_count: usize,
operation: &'static str,
) {
let mut tags = Vec::with_capacity(element_size.saturating_mul(element_count));
for index in 0..element_count {
let address =
source.saturating_add(index.saturating_mul(stride).saturating_mul(element_size));
tags.extend(self.local_tags(core, address, element_size));
}
let provenance = Provenance::merge_all(tags.iter().copied());
self.store_local_tags(core, destination, &tags);
self.local_event(
operation,
core,
destination,
tags.len(),
provenance,
&[provenance],
);
}
pub fn local_transform(
&mut self,
core: usize,
destination: usize,
sources: &[(usize, usize)],
output_size: usize,
operation: &'static str,
) {
let source_tags: Vec<Vec<_>> = sources
.iter()
.map(|&(address, size)| self.local_tags(core, address, size))
.collect();
let mut output = Vec::with_capacity(output_size);
for index in 0..output_size {
let provenance = Provenance::merge_all(
source_tags
.iter()
.filter_map(|tags| tags.get(index).copied()),
);
output.push(provenance);
}
let provenance = Provenance::merge_all(output.iter().copied());
self.store_local_tags(core, destination, &output);
let operands: Vec<_> = source_tags
.iter()
.map(|tags| Provenance::merge_all(tags.iter().copied()))
.collect();
self.local_event(
operation,
core,
destination,
output_size,
provenance,
&operands,
);
}
pub fn local_broadcast_transform(
&mut self,
core: usize,
destination: usize,
source: usize,
source_size: usize,
output_size: usize,
operation: &'static str,
) {
let tags = self.local_tags(core, source, source_size);
let provenance = Provenance::merge_all(tags.iter().copied());
self.store_local_tags(core, destination, &vec![provenance; output_size]);
self.local_event(
operation,
core,
destination,
output_size,
provenance,
&[provenance],
);
}
pub fn local_mvm_transform(
&mut self,
core: usize,
destination: usize,
source: usize,
element_size: usize,
output_size: usize,
used_rows: &[bool],
operation: &'static str,
) {
let mut provenance = Provenance::Uninitialized;
for (row, used) in used_rows.iter().copied().enumerate() {
if used {
provenance = provenance.merge(Provenance::merge_all(self.local_tags(
core,
source + row * element_size,
element_size,
)));
}
}
self.store_local_tags(core, destination, &vec![provenance; output_size]);
self.local_event(
operation,
core,
destination,
output_size,
provenance,
&[provenance],
);
}
pub fn send_transfer(
&mut self,
sender: usize,
receiver: usize,
source: usize,
destination: usize,
size: usize,
) {
let tags = self.local_tags(sender, source, size);
let provenance = Provenance::merge_all(tags.iter().copied());
self.store_local_tags(receiver, destination, &tags);
let mut event = json!({
"event": "send_recv_transfer",
"cycle": self.context.cycle,
"pc": self.context.pc,
"core_iteration": self.context.iteration,
"sender_core": sender,
"receiver_core": receiver,
"source_address": source,
"destination_address": destination,
"size": size,
});
if let Some(object) = event.as_object_mut() {
object.extend(provenance.json().as_object().unwrap().clone());
}
self.write(event);
}
fn local_event(
&self,
operation: &'static str,
core: usize,
destination: usize,
size: usize,
provenance: Provenance,
operands: &[Provenance],
) {
let mut event = json!({
"event": "local_compute",
"operation": operation,
"cycle": self.context.cycle,
"core": core,
"pc": self.context.pc,
"core_iteration": self.context.iteration,
"destination_address": destination,
"size": size,
"operand_provenance": operands.iter().map(|tag| tag.json()).collect::<Vec<_>>(),
});
if let Some(object) = event.as_object_mut() {
object.extend(provenance.json().as_object().unwrap().clone());
}
self.write(event);
if provenance.is_mixed() {
self.write(json!({
"event": "cross_sample_data_mix",
"cycle": self.context.cycle,
"core": core,
"pc": self.context.pc,
"core_iteration": self.context.iteration,
"operation": operation,
"destination_address": destination,
"size": size,
"operand_provenance": operands.iter().map(|tag| tag.json()).collect::<Vec<_>>(),
"provenance": provenance.samples(),
}));
}
}
}
@@ -33,10 +33,10 @@ pub struct SendRecv {
impl SendRecv { impl SendRecv {
pub fn new(num_core: usize) -> Self { pub fn new(num_core: usize) -> Self {
let sending = [Option::None].repeat(num_core); let sending = [Option::None].repeat(num_core);
let reciving = [Option::None].repeat(num_core); let receiving = [Option::None].repeat(num_core);
Self { Self {
sending: sending.into(), sending: sending.into(),
receiving: reciving.into(), receiving: receiving.into(),
} }
} }
} }
@@ -73,18 +73,27 @@ where
let data = inst.data; let data = inst.data;
TRACER.lock().unwrap().pre_recv(cpu, data); TRACER.lock().unwrap().pre_recv(cpu, data);
} }
let [sender_core, reciver_core] = {
cpu.get_multiple_cores([sender.internal_core, receiver.internal_core]); let [sender_core, receiver_core] =
let memory = sender_core cpu.get_multiple_cores([sender.internal_core, receiver.internal_core]);
.load::<u8>(sender.address, sender.size) let memory = sender_core
.with_context(|| { .load::<u8>(sender.address, sender.size)
format!( .with_context(|| {
"Sender crash tranfering memroy from {} with size {}", format!(
sender.address, sender.size "Sender crashed while transferring memory from {} with size {}",
) sender.address, sender.size
}) )
.unwrap(); })
reciver_core.execute_store(receiver.address, memory[0]); .unwrap();
receiver_core.execute_store(receiver.address, memory[0]);
}
cpu.provenance_send_transfer(
sender.internal_core,
receiver.internal_core,
sender.address,
receiver.address,
sender.size,
);
{ {
let sender = &mut core_instructions[sender.internal_core]; let sender = &mut core_instructions[sender.internal_core];
let pc = sender.program_counter; let pc = sender.program_counter;
@@ -124,19 +133,19 @@ where
let receiver: usize = imm_core.try_into().expect("imm_core can not be negative"); let receiver: usize = imm_core.try_into().expect("imm_core can not be negative");
assert_ne!(receiver, 0, "Host can not use receive"); assert_ne!(receiver, 0, "Host can not use receive");
send_recv.sending[sender] = Some(SendRecvInfo::new(sender, receiver, address, imm_len)); send_recv.sending[sender] = Some(SendRecvInfo::new(sender, receiver, address, imm_len));
let transfered = transfer_memory( let transferred = transfer_memory(
cpu, cpu,
core_instructions, core_instructions,
send_recv.sending[sender], send_recv.sending[sender],
send_recv.receiving[receiver], send_recv.receiving[receiver],
); );
if transfered { if transferred {
send_recv.sending[sender] = None; send_recv.sending[sender] = None;
send_recv.receiving[receiver] = None; send_recv.receiving[receiver] = None;
} }
(transfered, if transfered { receiver } else { 0 }) (transferred, if transferred { receiver } else { 0 })
} }
InstructionStatus::Reciving(instruction_data) => { InstructionStatus::Receiving(instruction_data) => {
let (core_idx, imm_core) = instruction_data.get_core_immcore(); let (core_idx, imm_core) = instruction_data.get_core_immcore();
let rd = instruction_data.rd(); let rd = instruction_data.rd();
let imm_len = instruction_data let imm_len = instruction_data
@@ -153,17 +162,17 @@ where
assert_ne!(sender, 0, "Host can not use send"); assert_ne!(sender, 0, "Host can not use send");
send_recv.receiving[receiver] = send_recv.receiving[receiver] =
Some(SendRecvInfo::new(receiver, sender, address, imm_len)); Some(SendRecvInfo::new(receiver, sender, address, imm_len));
let transfered = transfer_memory( let transferred = transfer_memory(
cpu, cpu,
core_instructions, core_instructions,
send_recv.sending[sender], send_recv.sending[sender],
send_recv.receiving[receiver], send_recv.receiving[receiver],
); );
if transfered { if transferred {
send_recv.sending[sender] = None; send_recv.sending[sender] = None;
send_recv.receiving[receiver] = None; send_recv.receiving[receiver] = None;
} }
(transfered, if transfered { sender } else { 0 }) (transferred, if transferred { sender } else { 0 })
} }
_ => (false, 0), _ => (false, 0),
} }
@@ -0,0 +1,70 @@
mod common;
use pimcore::{
CoreInstructionsBuilder, Executable,
instruction_set::{InstructionsBuilder, instruction_data::InstructionDataBuilder, isa::*},
};
#[test]
fn restarts_cores_and_loads_each_input() {
let cpu = common::empty_cpu(1);
let mut cores = CoreInstructionsBuilder::new(1);
let mut instructions = InstructionsBuilder::new();
let mut data = InstructionDataBuilder::new();
data.set_core_indx(1).fix_core_indx();
instructions.make_inst(sldi, data.set_rdimm(1, 0).build());
instructions.make_inst(sldi, data.set_rdimm(2, 0).build());
instructions.make_inst(ld, data.set_rdr1(2, 1).set_imm_len(4).build());
instructions.make_inst(sldi, data.set_rdimm(3, 4).build());
instructions.make_inst(st, data.set_rdr1(3, 2).set_imm_len(4).build());
cores.set_core(1, instructions.build());
let mut executable = Executable::new(cpu, cores.build());
let first = 1.0f32.to_ne_bytes();
let second = 2.0f32.to_ne_bytes();
assert!(
executable
.execute_batch(&[&first[..3]], &[(0, 4)], &[])
.is_err()
);
executable
.execute_batch(&[&first, &second], &[(0, 4)], &[])
.unwrap();
assert_eq!(
executable.cpu_mut().host().load::<f32>(4, 4).unwrap()[0],
[2.0]
);
}
#[test]
fn records_each_iteration_output() {
let cpu = common::empty_cpu(1);
let mut cores = CoreInstructionsBuilder::new(1);
let mut instructions = InstructionsBuilder::new();
let mut data = InstructionDataBuilder::new();
data.set_core_indx(1).fix_core_indx();
instructions.make_inst(sldi, data.set_rdimm(1, 0).build());
instructions.make_inst(sldi, data.set_rdimm(2, 0).build());
instructions.make_inst(ld, data.set_rdr1(2, 1).set_imm_len(4).build());
instructions.make_inst(sldi, data.set_rdimm(3, 4).build());
instructions.make_inst(st, data.set_rdr1(3, 2).set_imm_len(4).build());
cores.set_core(1, instructions.build());
let mut executable = Executable::new(cpu, cores.build());
let first = 1.0f32.to_ne_bytes();
let second = 2.0f32.to_ne_bytes();
let outputs = executable
.execute_batch(&[&first, &second], &[(0, 4)], &[(4, 2), (6, 2)])
.unwrap();
assert_eq!(outputs.len(), 2);
assert_eq!(
f32::from_ne_bytes(outputs[0].as_slice().try_into().unwrap()),
1.0
);
assert_eq!(
f32::from_ne_bytes(outputs[1].as_slice().try_into().unwrap()),
2.0
);
}
+1 -1
View File
@@ -19,7 +19,7 @@ struct ResolvedContiguousAddress {
}; };
/// Records compile-time facts used when interpreting address arithmetic and /// Records compile-time facts used when interpreting address arithmetic and
/// loop-carried aliases inside PIM regions. /// loop-carried aliases inside Pim regions.
struct StaticValueKnowledge { struct StaticValueKnowledge {
llvm::DenseMap<mlir::Value, int64_t> indexValues; llvm::DenseMap<mlir::Value, int64_t> indexValues;
llvm::DenseMap<mlir::Value, mlir::Value> aliases; llvm::DenseMap<mlir::Value, mlir::Value> aliases;
+4 -4
View File
@@ -85,12 +85,12 @@ static mlir::LogicalResult walkPimCoreBlockImpl(mlir::Block& block,
auto step = resolveIndexValue(forOp.getStep(), knowledge); auto step = resolveIndexValue(forOp.getStep(), knowledge);
if (failed(lower) || failed(upper) || failed(step) if (failed(lower) || failed(upper) || failed(step)
|| (mode == CoreWalkMode::ExecuteCommunication && *step <= 0)) { || (mode == CoreWalkMode::ExecuteCommunication && *step <= 0)) {
forOp.emitOpError() << "requires statically evaluable scf.for bounds for PIM " << purpose; forOp.emitOpError() << "requires statically evaluable scf.for bounds for Pim " << purpose;
hasFailure = true; hasFailure = true;
continue; continue;
} }
if (*step <= 0) { if (*step <= 0) {
forOp.emitOpError("requires positive scf.for step for PIM verification"); forOp.emitOpError("requires positive scf.for step for Pim verification");
hasFailure = true; hasFailure = true;
continue; continue;
} }
@@ -126,7 +126,7 @@ static mlir::LogicalResult walkPimCoreBlockImpl(mlir::Block& block,
if (auto ifOp = mlir::dyn_cast<mlir::scf::IfOp>(op)) { if (auto ifOp = mlir::dyn_cast<mlir::scf::IfOp>(op)) {
auto condition = resolveIndexValue(ifOp.getCondition(), knowledge); auto condition = resolveIndexValue(ifOp.getCondition(), knowledge);
if (failed(condition)) { if (failed(condition)) {
ifOp.emitOpError() << "requires statically evaluable scf.if condition for PIM " << purpose; ifOp.emitOpError() << "requires statically evaluable scf.if condition for Pim " << purpose;
hasFailure = true; hasFailure = true;
continue; continue;
} }
@@ -147,7 +147,7 @@ static mlir::LogicalResult walkPimCoreBlockImpl(mlir::Block& block,
if (auto switchOp = mlir::dyn_cast<mlir::scf::IndexSwitchOp>(op)) { if (auto switchOp = mlir::dyn_cast<mlir::scf::IndexSwitchOp>(op)) {
auto selector = resolveIndexValue(switchOp.getArg(), knowledge); auto selector = resolveIndexValue(switchOp.getArg(), knowledge);
if (failed(selector)) { if (failed(selector)) {
switchOp.emitOpError() << "requires a statically evaluable scf.index_switch selector for PIM " << purpose; switchOp.emitOpError() << "requires a statically evaluable scf.index_switch selector for Pim " << purpose;
hasFailure = true; hasFailure = true;
continue; continue;
} }
+1 -1
View File
@@ -14,7 +14,7 @@ namespace onnx_mlir {
using PimCoreCommunicationPlan = llvm::DenseMap<mlir::Block*, llvm::SmallVector<mlir::Operation*, 8>>; using PimCoreCommunicationPlan = llvm::DenseMap<mlir::Block*, llvm::SmallVector<mlir::Operation*, 8>>;
/// Returns true for ops in a `pim.core` body that only participate in static /// Returns true for ops in a `pim.core` body that only participate in static
/// address or index computation and therefore do not emit PIM instructions. /// address or index computation and therefore do not emit Pim instructions.
bool isCoreStaticAddressOp(mlir::Operation* op); bool isCoreStaticAddressOp(mlir::Operation* op);
/// Walks a `pim.core` body's communication stream, statically unrolling /// Walks a `pim.core` body's communication stream, statically unrolling
+2 -2
View File
@@ -9,7 +9,7 @@ llvm::FailureOr<mlir::func::FuncOp> getPimEntryFunc(mlir::ModuleOp moduleOp) {
llvm::SmallVector<mlir::ONNXEntryPointOp> entryPoints(moduleOp.getOps<mlir::ONNXEntryPointOp>()); llvm::SmallVector<mlir::ONNXEntryPointOp> entryPoints(moduleOp.getOps<mlir::ONNXEntryPointOp>());
if (entryPoints.size() > 1) { if (entryPoints.size() > 1) {
moduleOp.emitError("PIM pipeline requires a single ONNX entry point, but found ") << entryPoints.size(); moduleOp.emitError("Pim pipeline requires a single ONNX entry point, but found ") << entryPoints.size();
return mlir::failure(); return mlir::failure();
} }
if (!entryPoints.empty()) { if (!entryPoints.empty()) {
@@ -38,7 +38,7 @@ llvm::FailureOr<mlir::func::FuncOp> getPimEntryFunc(mlir::ModuleOp moduleOp) {
if (nonExternalFuncs.size() == 1) if (nonExternalFuncs.size() == 1)
return nonExternalFuncs.front(); return nonExternalFuncs.front();
moduleOp.emitError("could not resolve a unique PIM entry function"); moduleOp.emitError("could not resolve a unique Pim entry function");
return mlir::failure(); return mlir::failure();
} }
+1 -1
View File
@@ -5,7 +5,7 @@
namespace onnx_mlir { namespace onnx_mlir {
/// Resolves the function the PIM pipeline should treat as its entry point. /// Resolves the function the Pim pipeline should treat as its entry point.
/// Prefers ONNX entry-point metadata, then `main_graph`, then the only /// Prefers ONNX entry-point metadata, then `main_graph`, then the only
/// non-external function if the module is otherwise unambiguous. /// non-external function if the module is otherwise unambiguous.
llvm::FailureOr<mlir::func::FuncOp> getPimEntryFunc(mlir::ModuleOp moduleOp); llvm::FailureOr<mlir::func::FuncOp> getPimEntryFunc(mlir::ModuleOp moduleOp);
+1 -1
View File
@@ -32,7 +32,7 @@ struct ResolvedWeightView {
bool hasWeightAlways(mlir::Operation* op); bool hasWeightAlways(mlir::Operation* op);
/// Tags an op as producing a value that should stay materialized as a reusable /// Tags an op as producing a value that should stay materialized as a reusable
/// weight across later PIM lowering/codegen stages. /// weight across later Pim lowering/codegen stages.
void markWeightAlways(mlir::Operation* op); void markWeightAlways(mlir::Operation* op);
bool isSpatialMvmVmmWeightUse(mlir::OpOperand& use); bool isSpatialMvmVmmWeightUse(mlir::OpOperand& use);
+8 -8
View File
@@ -11,7 +11,7 @@ namespace onnx_mlir::pim {
namespace { namespace {
static void emitCrashMessage(llvm::StringRef fieldName, llvm::StringRef message) { static void emitCrashMessage(llvm::StringRef fieldName, llvm::StringRef message) {
llvm::errs() << "PIM " << fieldName << " " << message << "\n"; llvm::errs() << "Pim " << fieldName << " " << message << "\n";
} }
template <typename To, typename From> template <typename To, typename From>
@@ -65,7 +65,7 @@ InFlightDiagnostic emitCheckedArithmeticError(Operation* anchor, llvm::StringRef
} }
InFlightDiagnostic emitCheckedArithmeticError(Location loc, llvm::StringRef fieldName, llvm::StringRef message) { InFlightDiagnostic emitCheckedArithmeticError(Location loc, llvm::StringRef fieldName, llvm::StringRef message) {
return emitError(loc) << "PIM " << fieldName << " " << message; return emitError(loc) << "Pim " << fieldName << " " << message;
} }
FailureOr<int32_t> checkedI32(int64_t value, Operation* anchor, llvm::StringRef fieldName) { FailureOr<int32_t> checkedI32(int64_t value, Operation* anchor, llvm::StringRef fieldName) {
@@ -174,7 +174,7 @@ FailureOr<uint64_t> getCheckedShapedTypeSizeInBytes(ShapedType type, Location lo
int32_t checkedI32OrCrash(int64_t value, llvm::StringRef fieldName) { int32_t checkedI32OrCrash(int64_t value, llvm::StringRef fieldName) {
if (value < std::numeric_limits<int32_t>::min() || value > std::numeric_limits<int32_t>::max()) { if (value < std::numeric_limits<int32_t>::min() || value > std::numeric_limits<int32_t>::max()) {
emitCrashMessage(fieldName, "is outside representable range"); emitCrashMessage(fieldName, "is outside representable range");
llvm_unreachable("PIM checked arithmetic failure"); llvm_unreachable("Pim checked arithmetic failure");
} }
return static_cast<int32_t>(value); return static_cast<int32_t>(value);
} }
@@ -182,7 +182,7 @@ int32_t checkedI32OrCrash(int64_t value, llvm::StringRef fieldName) {
int32_t checkedI32OrCrash(uint64_t value, llvm::StringRef fieldName) { int32_t checkedI32OrCrash(uint64_t value, llvm::StringRef fieldName) {
if (value > static_cast<uint64_t>(std::numeric_limits<int32_t>::max())) { if (value > static_cast<uint64_t>(std::numeric_limits<int32_t>::max())) {
emitCrashMessage(fieldName, "is outside representable range"); emitCrashMessage(fieldName, "is outside representable range");
llvm_unreachable("PIM checked arithmetic failure"); llvm_unreachable("Pim checked arithmetic failure");
} }
return static_cast<int32_t>(value); return static_cast<int32_t>(value);
} }
@@ -190,7 +190,7 @@ int32_t checkedI32OrCrash(uint64_t value, llvm::StringRef fieldName) {
uint8_t checkedU8OrCrash(uint64_t value, llvm::StringRef fieldName) { uint8_t checkedU8OrCrash(uint64_t value, llvm::StringRef fieldName) {
if (value > static_cast<uint64_t>(std::numeric_limits<uint8_t>::max())) { if (value > static_cast<uint64_t>(std::numeric_limits<uint8_t>::max())) {
emitCrashMessage(fieldName, "is outside representable range"); emitCrashMessage(fieldName, "is outside representable range");
llvm_unreachable("PIM checked arithmetic failure"); llvm_unreachable("Pim checked arithmetic failure");
} }
return static_cast<uint8_t>(value); return static_cast<uint8_t>(value);
} }
@@ -198,7 +198,7 @@ uint8_t checkedU8OrCrash(uint64_t value, llvm::StringRef fieldName) {
size_t checkedSizeOrCrash(int64_t value, llvm::StringRef fieldName) { size_t checkedSizeOrCrash(int64_t value, llvm::StringRef fieldName) {
if (value < 0) { if (value < 0) {
emitCrashMessage(fieldName, "is outside representable range"); emitCrashMessage(fieldName, "is outside representable range");
llvm_unreachable("PIM checked arithmetic failure"); llvm_unreachable("Pim checked arithmetic failure");
} }
return static_cast<size_t>(value); return static_cast<size_t>(value);
} }
@@ -206,7 +206,7 @@ size_t checkedSizeOrCrash(int64_t value, llvm::StringRef fieldName) {
size_t checkedAddOrCrash(size_t lhs, size_t rhs, llvm::StringRef fieldName) { size_t checkedAddOrCrash(size_t lhs, size_t rhs, llvm::StringRef fieldName) {
if (rhs > std::numeric_limits<size_t>::max() - lhs) { if (rhs > std::numeric_limits<size_t>::max() - lhs) {
emitCrashMessage(fieldName, "addition overflow"); emitCrashMessage(fieldName, "addition overflow");
llvm_unreachable("PIM checked arithmetic failure"); llvm_unreachable("Pim checked arithmetic failure");
} }
return lhs + rhs; return lhs + rhs;
} }
@@ -214,7 +214,7 @@ size_t checkedAddOrCrash(size_t lhs, size_t rhs, llvm::StringRef fieldName) {
size_t checkedMulOrCrash(size_t lhs, size_t rhs, llvm::StringRef fieldName) { size_t checkedMulOrCrash(size_t lhs, size_t rhs, llvm::StringRef fieldName) {
if (lhs != 0 && rhs > std::numeric_limits<size_t>::max() / lhs) { if (lhs != 0 && rhs > std::numeric_limits<size_t>::max() / lhs) {
emitCrashMessage(fieldName, "multiplication overflow"); emitCrashMessage(fieldName, "multiplication overflow");
llvm_unreachable("PIM checked arithmetic failure"); llvm_unreachable("Pim checked arithmetic failure");
} }
return lhs * rhs; return lhs * rhs;
} }
+1 -1
View File
@@ -4,7 +4,7 @@
namespace onnx_mlir { namespace onnx_mlir {
/// Returns the directory that should hold PIM artifacts/debug dumps for the /// Returns the directory that should hold Pim artifacts/debug dumps for the
/// current compiler invocation. /// current compiler invocation.
std::string getOutputDir(); std::string getOutputDir();
+5 -1
View File
@@ -140,9 +140,13 @@ OnnxMlirCompilerErrorCodes writeConfigJson(func::FuncOp funcOp,
configJson["array_group_map"] = std::move(xbarsPerArrayGroup); configJson["array_group_map"] = std::move(xbarsPerArrayGroup);
json::Array inputsAddresses; json::Array inputsAddresses;
for (BlockArgument input : funcOp.getArguments()) json::Array inputsSizes;
for (BlockArgument input : funcOp.getArguments()) {
inputsAddresses.push_back(memory.getValueAddress(input)); inputsAddresses.push_back(memory.getValueAddress(input));
inputsSizes.push_back(memory.hostMem.getMemEntry({input, std::nullopt}).size);
}
configJson["inputs_addresses"] = std::move(inputsAddresses); configJson["inputs_addresses"] = std::move(inputsAddresses);
configJson["inputs_sizes"] = std::move(inputsSizes);
json::Array outputsAddresses; json::Array outputsAddresses;
for (func::ReturnOp returnOp : funcOp.getOps<func::ReturnOp>()) for (func::ReturnOp returnOp : funcOp.getOps<func::ReturnOp>())
+3 -3
View File
@@ -171,19 +171,19 @@ inline Opcode opcodeFromString(llvm::StringRef opName) {
for (auto [index, name] : llvm::enumerate(kOpcodeNames)) for (auto [index, name] : llvm::enumerate(kOpcodeNames))
if (opName == name) if (opName == name)
return static_cast<Opcode>(index); return static_cast<Opcode>(index);
llvm_unreachable("Unsupported PIM binary opcode"); llvm_unreachable("Unsupported Pim binary opcode");
} }
inline llvm::StringRef opcodeToString(Opcode opcode) { inline llvm::StringRef opcodeToString(Opcode opcode) {
size_t index = static_cast<size_t>(opcode); size_t index = static_cast<size_t>(opcode);
assert(index < kOpcodeNames.size() && "Unsupported PIM binary opcode"); assert(index < kOpcodeNames.size() && "Unsupported Pim binary opcode");
return kOpcodeNames[index]; return kOpcodeNames[index];
} }
inline InstructionRecord makeInstructionRecord(const llvm::json::Object& instruction) { inline InstructionRecord makeInstructionRecord(const llvm::json::Object& instruction) {
InstructionRecord record; InstructionRecord record;
std::optional<llvm::StringRef> opName = instruction.getString("op"); std::optional<llvm::StringRef> opName = instruction.getString("op");
assert(opName && "Missing op field in PIM instruction"); assert(opName && "Missing op field in Pim instruction");
record.opcode = opcodeFromString(*opName); record.opcode = opcodeFromString(*opName);
const auto& format = kInstructionJsonFormats[static_cast<size_t>(record.opcode)]; const auto& format = kInstructionJsonFormats[static_cast<size_t>(record.opcode)];
if (format.rd) if (format.rd)
+51 -31
View File
@@ -125,7 +125,7 @@ static bool isZeroSplatGlobal(mlir::Value value) {
return false; return false;
} }
// PIM instruction immediates are serialized as signed int32_t fields today // Pim instruction immediates are serialized as signed int32_t fields today
// (`sldi` goes through checkedI32OrCrash), so local addresses must stay within // (`sldi` goes through checkedI32OrCrash), so local addresses must stay within
// the non-negative int32_t range. // the non-negative int32_t range.
static FailureOr<size_t> checkedAlignTo(size_t value, size_t alignment, Operation* anchor, StringRef fieldName) { static FailureOr<size_t> checkedAlignTo(size_t value, size_t alignment, Operation* anchor, StringRef fieldName) {
@@ -141,7 +141,7 @@ static void printMemoryOverflowDiagnostic(const MemoryValueKey& key,
size_t requestedSize, size_t requestedSize,
size_t currentFirstAvailableAddress, size_t currentFirstAvailableAddress,
size_t alignedEndAddress) { size_t alignedEndAddress) {
llvm::errs() << "PIM local memory allocation overflow\n"; llvm::errs() << "Pim local memory allocation overflow\n";
llvm::errs() << "Requested allocation size: " << requestedSize << " bytes\n"; llvm::errs() << "Requested allocation size: " << requestedSize << " bytes\n";
llvm::errs() << "Current firstAvailableAddress: " << currentFirstAvailableAddress << "\n"; llvm::errs() << "Current firstAvailableAddress: " << currentFirstAvailableAddress << "\n";
llvm::errs() << "Aligned end address: " << alignedEndAddress << "\n"; llvm::errs() << "Aligned end address: " << alignedEndAddress << "\n";
@@ -187,7 +187,7 @@ size_t PimMemory::allocateAddress(size_t size, const MemoryValueKey& key) {
size, size,
firstAvailableAddress, firstAvailableAddress,
succeeded(checkedAlignedEnd) ? *checkedAlignedEnd : kPimLocalMemoryAddressLimit); succeeded(checkedAlignedEnd) ? *checkedAlignedEnd : kPimLocalMemoryAddressLimit);
llvm_unreachable("PIM local memory allocation overflow"); llvm_unreachable("Pim local memory allocation overflow");
} }
firstAvailableAddress = *checkedAlignedEnd; firstAvailableAddress = *checkedAlignedEnd;
return address; return address;
@@ -276,7 +276,7 @@ void PimMemory::allocateCore(const CompiledCoreMemoryPlan& plan, std::optional<u
} }
else if (*localArenaSize != plan.arenaSize || reportRow.logicalLocalAllocationCount != plan.logicalAllocationCount else if (*localArenaSize != plan.arenaSize || reportRow.logicalLocalAllocationCount != plan.logicalAllocationCount
|| reportRow.logicalLocalBytes != plan.logicalBytes) || reportRow.logicalLocalBytes != plan.logicalBytes)
llvm_unreachable("inconsistent PIM local-memory plan across core-batch lanes"); llvm_unreachable("inconsistent Pim local-memory plan across core-batch lanes");
for (const CompiledLocalMemoryEntry& entry : plan.entries) { for (const CompiledLocalMemoryEntry& entry : plan.entries) {
MemoryValueKey key = getMemoryValueKey(entry.value, lane); MemoryValueKey key = getMemoryValueKey(entry.value, lane);
ownedMemEntriesMap[key] = entry.memory; ownedMemEntriesMap[key] = entry.memory;
@@ -352,8 +352,8 @@ size_t PimAcceleratorMemory::getValueAddress(mlir::Value value,
llvm_unreachable("Missing mem entry"); llvm_unreachable("Missing mem entry");
} }
size_t byteOffset = pim::checkedSizeOrCrash(resolvedAddress->byteOffset, "resolved PIM byte offset"); size_t byteOffset = pim::checkedSizeOrCrash(resolvedAddress->byteOffset, "resolved Pim byte offset");
return pim::checkedAddOrCrash(iter->second.address, byteOffset, "resolved PIM address"); return pim::checkedAddOrCrash(iter->second.address, byteOffset, "resolved Pim address");
} }
llvm::FailureOr<int64_t> PimAcceleratorMemory::getIndexValue(mlir::Value value, llvm::FailureOr<int64_t> PimAcceleratorMemory::getIndexValue(mlir::Value value,
@@ -544,11 +544,23 @@ void PimCodeGen::setupRdRs1(size_t rdAddress, size_t rdOffset, size_t rs1Address
genSetRegisterImmediateUnsigned(1, pim::checkedAddOrCrash(rs1Address, rs1Offset, "rs1 address")); genSetRegisterImmediateUnsigned(1, pim::checkedAddOrCrash(rs1Address, rs1Offset, "rs1 address"));
} }
void PimCodeGen::setupRdRs1Rs2( std::array<uint8_t, 3> PimCodeGen::setupRdRs1Rs2(
size_t rdAddress, size_t rdOffset, size_t rs1Address, size_t rs1Offset, size_t rs2Address, size_t rs2Offset) const { size_t rdAddress, size_t rdOffset, size_t rs1Address, size_t rs1Offset, size_t rs2Address, size_t rs2Offset) const {
genSetRegisterImmediateUnsigned(0, pim::checkedAddOrCrash(rdAddress, rdOffset, "rd address")); size_t rd = pim::checkedAddOrCrash(rdAddress, rdOffset, "rd address");
genSetRegisterImmediateUnsigned(1, pim::checkedAddOrCrash(rs1Address, rs1Offset, "rs1 address")); size_t rs1 = pim::checkedAddOrCrash(rs1Address, rs1Offset, "rs1 address");
genSetRegisterImmediateUnsigned(2, pim::checkedAddOrCrash(rs2Address, rs2Offset, "rs2 address")); size_t rs2 = pim::checkedAddOrCrash(rs2Address, rs2Offset, "rs2 address");
genSetRegisterImmediateUnsigned(0, rd);
uint8_t rs1Register = 0;
if (rd != rs1) {
genSetRegisterImmediateUnsigned(1, rs1);
rs1Register = 1;
}
if (rd == rs2)
return {0, rs1Register, 0};
if (rs1 == rs2)
return {0, rs1Register, rs1Register};
genSetRegisterImmediateUnsigned(2, rs2);
return {0, rs1Register, 2};
} }
void PimCodeGen::emitMemCopyOp(pim_binary::Opcode opcode, void PimCodeGen::emitMemCopyOp(pim_binary::Opcode opcode,
@@ -664,13 +676,13 @@ void PimCodeGen::codeGenVMVOp(pim::PimVMVOp vmvOp, const StaticValueKnowledge& k
auto sourceType = cast<ShapedType>(vmvOp.getSource().getType()); auto sourceType = cast<ShapedType>(vmvOp.getSource().getType());
int32_t bitwidth = getVectorElementBitwidthOrCrash(sourceType); int32_t bitwidth = getVectorElementBitwidthOrCrash(sourceType);
ensureVectorBitwidth(bitwidth, bitwidth); ensureVectorBitwidth(bitwidth, bitwidth);
setupRdRs1Rs2(addressOf(vmvOp.getTarget(), knowledge), *targetOffset, auto registers = setupRdRs1Rs2(addressOf(vmvOp.getTarget(), knowledge), *targetOffset,
addressOf(vmvOp.getSource(), knowledge), *sourceOffset, 0, *sourceStride); addressOf(vmvOp.getSource(), knowledge), *sourceOffset, 0, *sourceStride);
pim_binary::InstructionRecord instruction; pim_binary::InstructionRecord instruction;
instruction.opcode = pim_binary::Opcode::vmv; instruction.opcode = pim_binary::Opcode::vmv;
instruction.rd = 0; instruction.rd = registers[0];
instruction.r1 = 1; instruction.r1 = registers[1];
instruction.r2OrImm = 2; instruction.r2OrImm = registers[2];
instruction.generic3 = vmvOp.getLength(); instruction.generic3 = vmvOp.getLength();
emitInstruction(instruction); emitInstruction(instruction);
} }
@@ -694,19 +706,26 @@ void PimCodeGen::codeGenSendOp(pim::PimSendOp sendOp, const StaticValueKnowledge
void PimCodeGen::codeGenWaitOp( void PimCodeGen::codeGenWaitOp(
pim::PimWaitOp waitOp, const StaticValueKnowledge& knowledge) const { pim::PimWaitOp waitOp, const StaticValueKnowledge& knowledge) const {
if (pimDisableSynchronization)
return;
auto eventRegister = indexOf(waitOp.getEventRegister(), knowledge); auto eventRegister = indexOf(waitOp.getEventRegister(), knowledge);
assert(succeeded(eventRegister) auto waitValue = indexOf(waitOp.getWaitValue(), knowledge);
&& "pim.wait event register must be statically resolvable during codegen"); assert(succeeded(eventRegister) && succeeded(waitValue)
&& "pim.wait operands must be statically resolvable during codegen");
if (*waitValue == 0)
return;
pim_binary::InstructionRecord instruction; pim_binary::InstructionRecord instruction;
instruction.opcode = pim_binary::Opcode::wait; instruction.opcode = pim_binary::Opcode::wait;
instruction.generic1 = pim::checkedI32OrCrash( instruction.generic1 = pim::checkedI32OrCrash(
*eventRegister, "wait event register"); *eventRegister, "wait event register");
instruction.generic2 = waitOp.getWaitValue(); instruction.generic2 = pim::checkedI32OrCrash(*waitValue, "wait value");
emitInstruction(instruction); emitInstruction(instruction);
} }
void PimCodeGen::codeGenSyncOp( void PimCodeGen::codeGenSyncOp(
pim::PimSyncOp syncOp, const StaticValueKnowledge& knowledge) const { pim::PimSyncOp syncOp, const StaticValueKnowledge& knowledge) const {
if (pimDisableSynchronization)
return;
auto targetCoreId = indexOf(syncOp.getTargetCoreId(), knowledge); auto targetCoreId = indexOf(syncOp.getTargetCoreId(), knowledge);
auto eventRegister = indexOf(syncOp.getEventRegister(), knowledge); auto eventRegister = indexOf(syncOp.getEventRegister(), knowledge);
assert(succeeded(targetCoreId) && succeeded(eventRegister) assert(succeeded(targetCoreId) && succeeded(eventRegister)
@@ -777,12 +796,13 @@ void PimCodeGen::emitBinaryVectorOp(pim_binary::Opcode opcode,
auto inputType = cast<ShapedType>(lhs.getType()); auto inputType = cast<ShapedType>(lhs.getType());
ensureVectorBitwidth(getVectorElementBitwidthOrCrash(inputType), ensureVectorBitwidth(getVectorElementBitwidthOrCrash(inputType),
getVectorElementBitwidthOrCrash(cast<ShapedType>(output.getType()))); getVectorElementBitwidthOrCrash(cast<ShapedType>(output.getType())));
setupRdRs1Rs2(addressOf(output, knowledge), 0, addressOf(lhs, knowledge), 0, addressOf(rhs, knowledge), 0); auto registers = setupRdRs1Rs2(
addressOf(output, knowledge), 0, addressOf(lhs, knowledge), 0, addressOf(rhs, knowledge), 0);
pim_binary::InstructionRecord instruction; pim_binary::InstructionRecord instruction;
instruction.opcode = opcode; instruction.opcode = opcode;
instruction.rd = 0; instruction.rd = registers[0];
instruction.r1 = 1; instruction.r1 = registers[1];
instruction.r2OrImm = 2; instruction.r2OrImm = registers[2];
instruction.generic3 = getVectorElementCountOrCrash(inputType); instruction.generic3 = getVectorElementCountOrCrash(inputType);
emitInstruction(instruction); emitInstruction(instruction);
} }
@@ -942,7 +962,7 @@ static LogicalResult executeCompiledCorePlan(
auto step = node.step.evaluate(knowledge); auto step = node.step.evaluate(knowledge);
auto forOp = cast<mlir::scf::ForOp>(node.op); auto forOp = cast<mlir::scf::ForOp>(node.op);
if (failed(lowerBound) || failed(upperBound) || failed(step) || *step <= 0) { if (failed(lowerBound) || failed(upperBound) || failed(step) || *step <= 0) {
forOp.emitOpError("requires statically evaluable scf.for bounds for PIM codegen"); forOp.emitOpError("requires statically evaluable scf.for bounds for Pim codegen");
return failure(); return failure();
} }
@@ -968,7 +988,7 @@ static LogicalResult executeCompiledCorePlan(
auto condition = node.condition.evaluate(knowledge); auto condition = node.condition.evaluate(knowledge);
auto ifOp = cast<mlir::scf::IfOp>(node.op); auto ifOp = cast<mlir::scf::IfOp>(node.op);
if (failed(condition)) { if (failed(condition)) {
ifOp.emitOpError("requires statically evaluable scf.if condition for PIM codegen"); ifOp.emitOpError("requires statically evaluable scf.if condition for Pim codegen");
return failure(); return failure();
} }
@@ -982,7 +1002,7 @@ static LogicalResult executeCompiledCorePlan(
auto selector = node.condition.evaluate(knowledge); auto selector = node.condition.evaluate(knowledge);
auto switchOp = cast<mlir::scf::IndexSwitchOp>(node.op); auto switchOp = cast<mlir::scf::IndexSwitchOp>(node.op);
if (failed(selector)) { if (failed(selector)) {
switchOp.emitOpError("requires a statically evaluable scf.index_switch selector for PIM codegen"); switchOp.emitOpError("requires a statically evaluable scf.index_switch selector for Pim codegen");
return failure(); return failure();
} }
const llvm::SmallVectorImpl<CompiledCoreNode>* selectedBody = node.defaultBody.get(); const llvm::SmallVectorImpl<CompiledCoreNode>* selectedBody = node.defaultBody.get();
@@ -1168,12 +1188,12 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
} }
auto getCompiledProgram = [&](Operation* op) { auto getCompiledProgram = [&](Operation* op) {
auto it = compiledPrograms.find(op); auto it = compiledPrograms.find(op);
assert(it != compiledPrograms.end() && "missing compiled PIM core program"); assert(it != compiledPrograms.end() && "missing compiled Pim core program");
return it->second.get(); return it->second.get();
}; };
auto getMemoryPlan = [&](Operation* op) { auto getMemoryPlan = [&](Operation* op) {
auto it = memoryPlans.find(op); auto it = memoryPlans.find(op);
assert(it != memoryPlans.end() && "missing PIM core memory plan"); assert(it != memoryPlans.end() && "missing Pim core memory plan");
return it->second.get(); return it->second.get();
}; };
@@ -1251,7 +1271,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
if (failed(weightView)) { if (failed(weightView)) {
std::string message; std::string message;
llvm::raw_string_ostream os(message); llvm::raw_string_ostream os(message);
os << "requires a statically resolvable dense global weight view during PIM codegen; weight=" os << "requires a statically resolvable dense global weight view during Pim codegen; weight="
<< vmmOp.getWeight() << " type=" << vmmOp.getWeight().getType(); << vmmOp.getWeight() << " type=" << vmmOp.getWeight().getType();
result.recordDiagnostic(vmmOp, os.str()); result.recordDiagnostic(vmmOp, os.str());
return failure(); return failure();
@@ -1259,7 +1279,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
if (weightView->shape.size() != 2) { if (weightView->shape.size() != 2) {
std::string message; std::string message;
llvm::raw_string_ostream os(message); llvm::raw_string_ostream os(message);
os << "requires a rank-2 matrix weight view during PIM codegen; resolved shape=["; os << "requires a rank-2 matrix weight view during Pim codegen; resolved shape=[";
llvm::interleaveComma(weightView->shape, os); llvm::interleaveComma(weightView->shape, os);
os << "] weight=" << vmmOp.getWeight() << " type=" << vmmOp.getWeight().getType(); os << "] weight=" << vmmOp.getWeight() << " type=" << vmmOp.getWeight().getType();
result.recordDiagnostic(vmmOp, os.str()); result.recordDiagnostic(vmmOp, os.str());
@@ -1371,7 +1391,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
} }
if (diagnostics.hasFailure()) if (diagnostics.hasFailure())
diagnostics.emitSuppressedSummary(summaryAnchor ? summaryAnchor : moduleOp.getOperation(), diagnostics.emitSuppressedSummary(summaryAnchor ? summaryAnchor : moduleOp.getOperation(),
"PIM codegen diagnostic(s)"); "Pim codegen diagnostic(s)");
for (size_t jobIndex = 0; jobIndex < jobs.size(); ++jobIndex) for (size_t jobIndex = 0; jobIndex < jobs.size(); ++jobIndex)
if (jobResults[jobIndex].status != CompilerSuccess) if (jobResults[jobIndex].status != CompilerSuccess)
@@ -1437,7 +1457,7 @@ OnnxMlirCompilerErrorCodes onnx_mlir::compileToPimCode(ModuleOp& moduleOp, std::
if (!batchPerCoreRow) if (!batchPerCoreRow)
batchPerCoreRow = result.reportRow; batchPerCoreRow = result.reportRow;
else if (!(*batchPerCoreRow == result.reportRow)) else if (!(*batchPerCoreRow == result.reportRow))
llvm_unreachable("one PIM core batch produced inconsistent per-core memory reports"); llvm_unreachable("one Pim core batch produced inconsistent per-core memory reports");
} }
uint64_t batchReportId = jobs[group.front()].batchReportId.value_or(0); uint64_t batchReportId = jobs[group.front()].batchReportId.value_or(0);
+1 -1
View File
@@ -176,7 +176,7 @@ class PimCodeGen {
void genSetRegisterImmediateUnsigned(size_t registerNumber, size_t immediate) const; void genSetRegisterImmediateUnsigned(size_t registerNumber, size_t immediate) const;
void setupRd(size_t rdAddress, size_t rdOffset) const; void setupRd(size_t rdAddress, size_t rdOffset) const;
void setupRdRs1(size_t rdAddress, size_t rdOffset, size_t rs1Address, size_t rs1Offset) const; void setupRdRs1(size_t rdAddress, size_t rdOffset, size_t rs1Address, size_t rs1Offset) const;
void setupRdRs1Rs2( std::array<uint8_t, 3> setupRdRs1Rs2(
size_t rdAddress, size_t rdOffset, size_t rs1Address, size_t rs1Offset, size_t rs2Address, size_t rs2Offset) const; size_t rdAddress, size_t rdOffset, size_t rs1Address, size_t rs1Offset, size_t rs2Address, size_t rs2Offset) const;
void emitMemCopyOp(pim_binary::Opcode opcode, void emitMemCopyOp(pim_binary::Opcode opcode,
+52 -33
View File
@@ -9,25 +9,25 @@
namespace onnx_mlir { namespace onnx_mlir {
llvm::cl::opt<PimEmissionTargetType> pimEmissionTarget( llvm::cl::opt<PimEmissionTargetType> pimEmissionTarget(
llvm::cl::desc("[Optional] Choose PIM-related target to emit (once selected it will cancel the other targets):"), llvm::cl::desc("[Optional] Choose Pim-related target to emit (once selected it will cancel the other targets):"),
llvm::cl::values(clEnumVal(EmitSpatial, "Lower model to spatial IR")), llvm::cl::values(clEnumVal(EmitSpatial, "Lower model to Spatial IR")),
llvm::cl::values(clEnumVal(EmitPim, "Lower model to PIM IR")), llvm::cl::values(clEnumVal(EmitPim, "Lower model to Pim IR")),
llvm::cl::values(clEnumVal(EmitPimBufferized, "Lower model to PIM IR and bufferize it")), llvm::cl::values(clEnumVal(EmitPimBufferized, "Lower model to Pim IR and bufferize it")),
llvm::cl::values(clEnumVal(EmitPimCodegen, "Lower model to PIM IR and generate code for PIM")), llvm::cl::values(clEnumVal(EmitPimCodegen, "Lower model to Pim IR and generate code for Pim")),
llvm::cl::init(EmitPimCodegen), llvm::cl::init(EmitPimCodegen),
llvm::cl::cat(OnnxMlirOptions)); llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<PimMemoryReportLevel> pimMemoryReport( llvm::cl::opt<PimMemoryReportLevel> pimMemoryReport(
"pim-memory-report", "pim-memory-report",
llvm::cl::desc("Emit a human-readable PIM memory planning report"), llvm::cl::desc("Emit a human-readable Pim memory planning report"),
llvm::cl::values(clEnumValN(PimMemoryReportNone, "none", "Do not emit any PIM memory planning report")), llvm::cl::values(clEnumValN(PimMemoryReportNone, "none", "Do not emit any Pim memory planning report")),
llvm::cl::values(clEnumValN(PimMemoryReportSummary, "summary", "Emit a concise PIM memory summary")), llvm::cl::values(clEnumValN(PimMemoryReportSummary, "summary", "Emit a concise Pim memory summary")),
llvm::cl::init(PimMemoryReportSummary), llvm::cl::init(PimMemoryReportSummary),
llvm::cl::cat(OnnxMlirOptions)); llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<PimConvLoweringType> pimConvLowering( llvm::cl::opt<PimConvLoweringType> pimConvLowering(
"pim-conv-lowering", "pim-conv-lowering",
llvm::cl::desc("Convolution lowering strategy for PIM"), llvm::cl::desc("Convolution lowering strategy for Pim"),
llvm::cl::values(clEnumValN(PimConvLoweringAuto, "auto", "Select the Conv lowering strategy automatically")), llvm::cl::values(clEnumValN(PimConvLoweringAuto, "auto", "Select the Conv lowering strategy automatically")),
llvm::cl::values(clEnumValN(PimConvLoweringLegacy, "legacy", "Use the legacy explicit-im2col Conv lowering")), llvm::cl::values(clEnumValN(PimConvLoweringLegacy, "legacy", "Use the legacy explicit-im2col Conv lowering")),
llvm::cl::values(clEnumValN(PimConvLoweringDepthwise, "depthwise", "Force the depthwise-specialized Conv lowering")), llvm::cl::values(clEnumValN(PimConvLoweringDepthwise, "depthwise", "Force the depthwise-specialized Conv lowering")),
@@ -55,20 +55,20 @@ llvm::cl::opt<PimSpatialDataflowExportType> pimExportSpatialDataflow(
llvm::cl::desc("Emit Gephi-importable CSV dataflow reports for Spatial pipeline snapshots"), llvm::cl::desc("Emit Gephi-importable CSV dataflow reports for Spatial pipeline snapshots"),
llvm::cl::values(clEnumValN(SpatialDataflowExportNone, "none", "Do not emit Spatial dataflow CSV reports")), llvm::cl::values(clEnumValN(SpatialDataflowExportNone, "none", "Do not emit Spatial dataflow CSV reports")),
llvm::cl::values( llvm::cl::values(
clEnumValN(SpatialDataflowExportSpatial1, "spatial1", "Emit spatial1 graph dataflow CSV reports")), clEnumValN(SpatialDataflowExportSpatial1, "spatial1", "Emit Spatial1 graph dataflow CSV reports")),
llvm::cl::values( llvm::cl::values(
clEnumValN(SpatialDataflowExportSpatial2, "spatial2", "Emit spatial2 trivially merged graph dataflow CSV reports")), clEnumValN(SpatialDataflowExportSpatial2, "spatial2", "Emit Spatial2 trivially merged graph dataflow CSV reports")),
llvm::cl::values( llvm::cl::values(
clEnumValN(SpatialDataflowExportSpatial3, "spatial3", "Emit spatial3 scheduled dataflow CSV reports")), clEnumValN(SpatialDataflowExportSpatial3, "spatial3", "Emit Spatial3 scheduled dataflow CSV reports")),
llvm::cl::values( llvm::cl::values(
clEnumValN(SpatialDataflowExportSpatial4, "spatial4", "Emit spatial4 realized dataflow CSV reports")), clEnumValN(SpatialDataflowExportSpatial4, "spatial4", "Emit Spatial4 realized dataflow CSV reports")),
llvm::cl::values(clEnumValN(SpatialDataflowExportAll, "all", "Emit all Spatial dataflow CSV reports")), llvm::cl::values(clEnumValN(SpatialDataflowExportAll, "all", "Emit all Spatial dataflow CSV reports")),
llvm::cl::init(SpatialDataflowExportNone), llvm::cl::init(SpatialDataflowExportNone),
llvm::cl::cat(OnnxMlirOptions)); llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<bool> llvm::cl::opt<bool>
pimOnlyCodegen("pim-only-codegen", pimOnlyCodegen("pim-only-codegen",
llvm::cl::desc("Only generate code for PIM (assume input is already in bufferized PIM IR)"), llvm::cl::desc("Only generate code for Pim (assume input is already in bufferized Pim IR)"),
llvm::cl::init(false), llvm::cl::init(false),
llvm::cl::cat(OnnxMlirOptions)); llvm::cl::cat(OnnxMlirOptions));
@@ -96,21 +96,37 @@ llvm::cl::opt<bool> pimEmitJson("pim-emit-json",
llvm::cl::opt<bool> pimDetectCommunicationDeadlock( llvm::cl::opt<bool> pimDetectCommunicationDeadlock(
"pim-detect-communication-deadlock", "pim-detect-communication-deadlock",
llvm::cl::desc("Expensively simulate the statically expanded PIM send/receive order at verification time and fail if a blocking communication deadlock is found"), llvm::cl::desc("Expensively simulate the statically expanded Pim send/receive order at verification time and fail if a blocking communication deadlock is found"),
llvm::cl::init(false), llvm::cl::init(false),
llvm::cl::cat(OnnxMlirOptions)); llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<bool> pimVerifyBufferizationCopyFreedom( llvm::cl::opt<bool> pimVerifyBufferizationCopyFreedom(
"pim-verify-bufferization-copy-freedom", "pim-verify-bufferization-copy-freedom",
llvm::cl::desc("Run the expensive official PIM tensor-copy freedom proof before bufferization"), llvm::cl::desc("Run the expensive official Pim tensor-copy freedom proof before bufferization"),
llvm::cl::init(false),
llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<bool> pimDisableSynchronization(
"pim-disable-synchronization",
llvm::cl::desc("Omit Pim wait/sync instructions from generated code for performance ablation"),
llvm::cl::init(false),
llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<bool> pimDisableSpatialPlanning(
"pim-disable-spatial-planning",
llvm::cl::desc("Select the trivial Spatial layout plan for performance ablation"),
llvm::cl::init(false), llvm::cl::init(false),
llvm::cl::cat(OnnxMlirOptions)); llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<size_t> llvm::cl::opt<size_t>
crossbarSize("crossbar-size", llvm::cl::desc("Width and height of a single crossbar"), llvm::cl::init(128)); crossbarSize("crossbar-size",
llvm::cl::desc("Width and height of a single crossbar (required for Pim compilation)"),
llvm::cl::init(0));
llvm::cl::opt<size_t> llvm::cl::opt<size_t>
crossbarCountInCore("crossbar-count", llvm::cl::desc("Number of crossbars in each core"), llvm::cl::init(64)); crossbarCountInCore("crossbar-count",
llvm::cl::desc("Number of crossbars in each core (required for Pim compilation)"),
llvm::cl::init(0));
llvm::cl::opt<size_t> pipelineStages( llvm::cl::opt<size_t> pipelineStages(
"pipeline", "pipeline",
@@ -119,32 +135,35 @@ llvm::cl::opt<size_t> pipelineStages(
llvm::cl::cat(OnnxMlirOptions)); llvm::cl::cat(OnnxMlirOptions));
llvm::cl::opt<long> coresCount("core-count", llvm::cl::opt<long> coresCount("core-count",
llvm::cl::desc("Number of cores in the chip. Required for PIM compilation."), llvm::cl::desc("Number of cores in the chip. Required for Pim compilation."),
llvm::cl::init(-1)); llvm::cl::init(-1));
llvm::cl::opt<std::string> pimTargetConfig( llvm::cl::opt<std::string> pimTargetConfig(
"pim-target-config", "pim-target-config",
llvm::cl::desc("PIM target configuration used to construct the Spatial scheduling cost model"), llvm::cl::desc("Pim target configuration used to construct the Spatial scheduling cost model"),
llvm::cl::init(""), llvm::cl::init(""),
llvm::cl::cat(OnnxMlirOptions)); llvm::cl::cat(OnnxMlirOptions));
bool hasExplicitPimCoreCount() { return coresCount.getNumOccurrences() != 0; } void verifyPimCompilerOptions() {
if (coresCount.getNumOccurrences() == 0)
void verifyExplicitPimCoreCount() { llvm::report_fatal_error("Pim compilation requires an explicit --core-count=<positive integer>");
if (!hasExplicitPimCoreCount())
llvm::report_fatal_error("PIM compilation requires an explicit --core-count=<positive integer>");
if (coresCount.getValue() <= 0) if (coresCount.getValue() <= 0)
llvm::report_fatal_error("PIM compilation requires --core-count to be a positive integer"); llvm::report_fatal_error("Pim compilation requires --core-count to be a positive integer");
} if (crossbarSize.getNumOccurrences() == 0)
llvm::report_fatal_error("Pim compilation requires an explicit --crossbar-size=<positive integer>");
void verifyPimPipelineStages() { if (crossbarSize.getValue() == 0)
llvm::report_fatal_error("Pim compilation requires --crossbar-size to be a positive integer");
if (crossbarCountInCore.getNumOccurrences() == 0)
llvm::report_fatal_error("Pim compilation requires an explicit --crossbar-count=<positive integer>");
if (crossbarCountInCore.getValue() == 0)
llvm::report_fatal_error("Pim compilation requires --crossbar-count to be a positive integer");
if (pipelineStages.getValue() == 0) if (pipelineStages.getValue() == 0)
llvm::report_fatal_error("PIM compilation requires --pipeline to be positive"); llvm::report_fatal_error("Pim compilation requires --pipeline to be positive");
if (static_cast<size_t>(coresCount.getValue()) % pipelineStages.getValue() != 0) if (static_cast<size_t>(coresCount.getValue()) < pipelineStages.getValue())
llvm::report_fatal_error("PIM compilation requires --core-count to be divisible by --pipeline"); llvm::report_fatal_error("Pim compilation requires --pipeline not to exceed --core-count");
if (crossbarCountInCore.getValue() if (crossbarCountInCore.getValue()
> std::numeric_limits<size_t>::max() / pipelineStages.getValue()) > std::numeric_limits<size_t>::max() / pipelineStages.getValue())
llvm::report_fatal_error("PIM compilation --crossbar-count * --pipeline overflows"); llvm::report_fatal_error("Pim compilation --crossbar-count * --pipeline overflows");
} }
} // namespace onnx_mlir } // namespace onnx_mlir
+3 -3
View File
@@ -59,6 +59,8 @@ extern llvm::cl::opt<bool> pimEmitJson;
extern llvm::cl::opt<bool> pimReportConvLowering; extern llvm::cl::opt<bool> pimReportConvLowering;
extern llvm::cl::opt<bool> pimDetectCommunicationDeadlock; extern llvm::cl::opt<bool> pimDetectCommunicationDeadlock;
extern llvm::cl::opt<bool> pimVerifyBufferizationCopyFreedom; extern llvm::cl::opt<bool> pimVerifyBufferizationCopyFreedom;
extern llvm::cl::opt<bool> pimDisableSynchronization;
extern llvm::cl::opt<bool> pimDisableSpatialPlanning;
extern llvm::cl::opt<size_t> crossbarSize; extern llvm::cl::opt<size_t> crossbarSize;
extern llvm::cl::opt<size_t> crossbarCountInCore; extern llvm::cl::opt<size_t> crossbarCountInCore;
@@ -68,8 +70,6 @@ extern llvm::cl::opt<std::string> pimTargetConfig;
extern llvm::cl::opt<uint64_t> pimConvIm2colMaxElements; extern llvm::cl::opt<uint64_t> pimConvIm2colMaxElements;
extern llvm::cl::opt<uint64_t> pimConvStreamChunkPositions; extern llvm::cl::opt<uint64_t> pimConvStreamChunkPositions;
bool hasExplicitPimCoreCount(); void verifyPimCompilerOptions();
void verifyExplicitPimCoreCount();
void verifyPimPipelineStages();
} // namespace onnx_mlir } // namespace onnx_mlir
+23 -21
View File
@@ -12,6 +12,7 @@
#include <limits> #include <limits>
#include <tuple> #include <tuple>
#include "src/Accelerators/PIM/Common/PimCommon.hpp"
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp" #include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
#include "src/Accelerators/PIM/Compiler/PimCompilerUtils.hpp" #include "src/Accelerators/PIM/Compiler/PimCompilerUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/ONNXToSpatialOptions.hpp" #include "src/Accelerators/PIM/Conversion/ONNXToSpatial/ONNXToSpatialOptions.hpp"
@@ -78,6 +79,7 @@ spatial::SchedulingTarget getDefaultPimSchedulingTarget() {
target.residentWeightCapacity = crossbarCountInCore.getValue(); target.residentWeightCapacity = crossbarCountInCore.getValue();
target.matrixRows = crossbarSize.getValue(); target.matrixRows = crossbarSize.getValue();
target.matrixColumns = crossbarSize.getValue(); target.matrixColumns = crossbarSize.getValue();
target.synchronizationRegisterCount = kPimEventRegisterCount;
setDefaultPimInterProcessorLatencies(target); setDefaultPimInterProcessorLatencies(target);
return target; return target;
@@ -95,7 +97,7 @@ spatial::ConvLoweringStrategy getSpatialConvLoweringStrategy(PimConvLoweringType
case PimConvLoweringInputKTiled: return spatial::ConvLoweringStrategy::InputKTiled; case PimConvLoweringInputKTiled: return spatial::ConvLoweringStrategy::InputKTiled;
case PimConvLoweringTiled2D: return spatial::ConvLoweringStrategy::Tiled2D; case PimConvLoweringTiled2D: return spatial::ConvLoweringStrategy::Tiled2D;
} }
llvm_unreachable("unknown PIM Conv lowering strategy"); llvm_unreachable("unknown Pim Conv lowering strategy");
} }
spatial::SpatialDataflowExportStage getPimSpatialDataflowExportStage( spatial::SpatialDataflowExportStage getPimSpatialDataflowExportStage(
@@ -108,7 +110,7 @@ spatial::SpatialDataflowExportStage getPimSpatialDataflowExportStage(
case SpatialDataflowExportSpatial4: return spatial::SpatialDataflowExportStage::Spatial4; case SpatialDataflowExportSpatial4: return spatial::SpatialDataflowExportStage::Spatial4;
case SpatialDataflowExportAll: return spatial::SpatialDataflowExportStage::All; case SpatialDataflowExportAll: return spatial::SpatialDataflowExportStage::All;
} }
llvm_unreachable("unknown PIM Spatial dataflow export stage"); llvm_unreachable("unknown Pim Spatial dataflow export stage");
} }
spatial::SpatialTargetResources getPimSpatialTargetResources(const spatial::SchedulingTarget& target) { spatial::SpatialTargetResources getPimSpatialTargetResources(const spatial::SchedulingTarget& target) {
@@ -118,7 +120,7 @@ spatial::SpatialTargetResources getPimSpatialTargetResources(const spatial::Sche
resources.processorCount = target.processorCount; resources.processorCount = target.processorCount;
resources.vectorWidth = target.vectorWidth; resources.vectorWidth = target.vectorWidth;
if (failed(resources.verify())) if (failed(resources.verify()))
llvm::report_fatal_error("PIM target resources are incomplete"); llvm::report_fatal_error("Pim target resources are incomplete");
return resources; return resources;
} }
@@ -136,7 +138,7 @@ const llvm::json::Object& requireObject(const llvm::json::Object& object,
llvm::StringRef path) { llvm::StringRef path) {
const llvm::json::Object* nested = object.getObject(key); const llvm::json::Object* nested = object.getObject(key);
if (!nested) if (!nested)
llvm::report_fatal_error("PIM target config is missing object '" + path + "." + key + "'"); llvm::report_fatal_error("Pim target config is missing object '" + path + "." + key + "'");
return *nested; return *nested;
} }
@@ -149,7 +151,7 @@ Cost getConfigCost(const llvm::json::Object& object,
return fallback; return fallback;
if (!std::isfinite(*number) || *number < 0.0 || (!allowZero && *number == 0.0) if (!std::isfinite(*number) || *number < 0.0 || (!allowZero && *number == 0.0)
|| *number > static_cast<double>(std::numeric_limits<Cost>::max())) || *number > static_cast<double>(std::numeric_limits<Cost>::max()))
llvm::report_fatal_error("PIM target config field '" + key + "' must be a valid positive number"); llvm::report_fatal_error("Pim target config field '" + key + "' must be a valid positive number");
return static_cast<Cost>(std::ceil(*number)); return static_cast<Cost>(std::ceil(*number));
} }
@@ -157,11 +159,11 @@ std::pair<size_t, size_t> getConfigPair(const llvm::json::Object& object,
llvm::StringRef key) { llvm::StringRef key) {
const llvm::json::Array* values = object.getArray(key); const llvm::json::Array* values = object.getArray(key);
if (!values || values->size() != 2) if (!values || values->size() != 2)
llvm::report_fatal_error("PIM target config field '" + key + "' must contain two integers"); llvm::report_fatal_error("Pim target config field '" + key + "' must contain two integers");
std::optional<int64_t> first = (*values)[0].getAsInteger(); std::optional<int64_t> first = (*values)[0].getAsInteger();
std::optional<int64_t> second = (*values)[1].getAsInteger(); std::optional<int64_t> second = (*values)[1].getAsInteger();
if (!first || !second || *first <= 0 || *second <= 0) if (!first || !second || *first <= 0 || *second <= 0)
llvm::report_fatal_error("PIM target config field '" + key + "' must contain two positive integers"); llvm::report_fatal_error("Pim target config field '" + key + "' must contain two positive integers");
return {static_cast<size_t>(*first), static_cast<size_t>(*second)}; return {static_cast<size_t>(*first), static_cast<size_t>(*second)};
} }
@@ -172,7 +174,7 @@ void loadPimInterProcessorLatencies(
network.getString("net_config_file_path"); network.getString("net_config_file_path");
if (!filename) if (!filename)
llvm::report_fatal_error( llvm::report_fatal_error(
"PIM target config is missing network latency file path"); "Pim target config is missing network latency file path");
llvm::SmallString<256> networkPath(*filename); llvm::SmallString<256> networkPath(*filename);
if (!llvm::sys::path::is_absolute(networkPath)) { if (!llvm::sys::path::is_absolute(networkPath)) {
@@ -185,19 +187,19 @@ void loadPimInterProcessorLatencies(
auto buffer = llvm::MemoryBuffer::getFile(networkPath); auto buffer = llvm::MemoryBuffer::getFile(networkPath);
if (!buffer) if (!buffer)
llvm::report_fatal_error( llvm::report_fatal_error(
llvm::Twine("failed to read PIM network config '") llvm::Twine("failed to read Pim network config '")
+ networkPath + "': " + buffer.getError().message()); + networkPath + "': " + buffer.getError().message());
auto parsed = llvm::json::parse(buffer.get()->getBuffer()); auto parsed = llvm::json::parse(buffer.get()->getBuffer());
if (!parsed) if (!parsed)
llvm::report_fatal_error( llvm::report_fatal_error(
llvm::Twine("failed to parse PIM network config '") llvm::Twine("failed to parse Pim network config '")
+ networkPath + "': " + llvm::toString(parsed.takeError())); + networkPath + "': " + llvm::toString(parsed.takeError()));
const llvm::json::Object* root = parsed->getAsObject(); const llvm::json::Object* root = parsed->getAsObject();
const llvm::json::Object* latencies = const llvm::json::Object* latencies =
root ? root->getObject("latency") : nullptr; root ? root->getObject("latency") : nullptr;
if (!latencies) if (!latencies)
llvm::report_fatal_error( llvm::report_fatal_error(
"PIM network config is missing its latency matrix"); "Pim network config is missing its latency matrix");
target.interProcessorLatencyNs.assign( target.interProcessorLatencyNs.assign(
target.processorCount * target.processorCount, 0); target.processorCount * target.processorCount, 0);
@@ -208,7 +210,7 @@ void loadPimInterProcessorLatencies(
const llvm::json::Object* row = latencies->getObject(sourceKey); const llvm::json::Object* row = latencies->getObject(sourceKey);
if (!row) if (!row)
llvm::report_fatal_error( llvm::report_fatal_error(
llvm::Twine("PIM network config is missing latency row ") llvm::Twine("Pim network config is missing latency row ")
+ sourceKey); + sourceKey);
for (size_t destination = 0; for (size_t destination = 0;
destination < target.processorCount; ++destination) { destination < target.processorCount; ++destination) {
@@ -218,7 +220,7 @@ void loadPimInterProcessorLatencies(
std::optional<double> latency = row->getNumber(destinationKey); std::optional<double> latency = row->getNumber(destinationKey);
if (!latency || !std::isfinite(*latency) || *latency <= 0.0) if (!latency || !std::isfinite(*latency) || *latency <= 0.0)
llvm::report_fatal_error( llvm::report_fatal_error(
llvm::Twine("PIM network config is missing latency ") llvm::Twine("Pim network config is missing latency ")
+ sourceKey + " -> " + destinationKey); + sourceKey + " -> " + destinationKey);
Cost roundedLatency = static_cast<Cost>(std::ceil(*latency)); Cost roundedLatency = static_cast<Cost>(std::ceil(*latency));
target.interProcessorLatencyNs[ target.interProcessorLatencyNs[
@@ -242,17 +244,17 @@ spatial::SchedulingTarget getPimSchedulingTarget() {
auto buffer = llvm::MemoryBuffer::getFile(pimTargetConfig); auto buffer = llvm::MemoryBuffer::getFile(pimTargetConfig);
if (!buffer) if (!buffer)
llvm::report_fatal_error( llvm::report_fatal_error(
llvm::Twine("failed to read PIM target config '") llvm::Twine("failed to read Pim target config '")
+ pimTargetConfig.getValue() + "': " + buffer.getError().message()); + pimTargetConfig.getValue() + "': " + buffer.getError().message());
auto parsed = llvm::json::parse(buffer.get()->getBuffer()); auto parsed = llvm::json::parse(buffer.get()->getBuffer());
if (!parsed) if (!parsed)
llvm::report_fatal_error( llvm::report_fatal_error(
llvm::Twine("failed to parse PIM target config '") llvm::Twine("failed to parse Pim target config '")
+ pimTargetConfig.getValue() + "': " + pimTargetConfig.getValue() + "': "
+ llvm::toString(parsed.takeError())); + llvm::toString(parsed.takeError()));
const llvm::json::Object* root = parsed->getAsObject(); const llvm::json::Object* root = parsed->getAsObject();
if (!root) if (!root)
llvm::report_fatal_error("PIM target config must contain a JSON object"); llvm::report_fatal_error("Pim target config must contain a JSON object");
const llvm::json::Object& chip = requireObject(*root, "chip_config", "root"); const llvm::json::Object& chip = requireObject(*root, "chip_config", "root");
const llvm::json::Object& core = requireObject(chip, "core_config", "chip_config"); const llvm::json::Object& core = requireObject(chip, "core_config", "chip_config");
@@ -265,7 +267,7 @@ spatial::SchedulingTarget getPimSchedulingTarget() {
std::optional<int64_t> coreCount = chip.getInteger("core_cnt"); std::optional<int64_t> coreCount = chip.getInteger("core_cnt");
if (!coreCount || *coreCount <= 0) if (!coreCount || *coreCount <= 0)
llvm::report_fatal_error("PIM target config field 'core_cnt' must be a positive integer"); llvm::report_fatal_error("Pim target config field 'core_cnt' must be a positive integer");
target.processorCount = static_cast<size_t>(*coreCount); target.processorCount = static_cast<size_t>(*coreCount);
target.residentWeightCapacity = target.residentWeightCapacity =
getConfigCost(matrix, "xbar_array_count", target.residentWeightCapacity); getConfigCost(matrix, "xbar_array_count", target.residentWeightCapacity);
@@ -276,7 +278,7 @@ spatial::SchedulingTarget getPimSchedulingTarget() {
|| target.residentWeightCapacity != crossbarCountInCore.getValue() || target.residentWeightCapacity != crossbarCountInCore.getValue()
|| target.matrixRows != crossbarSize.getValue() || target.matrixRows != crossbarSize.getValue()
|| target.matrixColumns != crossbarSize.getValue()) || target.matrixColumns != crossbarSize.getValue())
llvm::report_fatal_error("PIM target config resources do not match --core-count, " llvm::report_fatal_error("Pim target config resources do not match --core-count, "
"--crossbar-count, and --crossbar-size"); "--crossbar-count, and --crossbar-size");
loadPimInterProcessorLatencies(target, network); loadPimInterProcessorLatencies(target, network);
@@ -329,8 +331,7 @@ void addPassesPim(OwningOpRef<ModuleOp>& module,
PassManager& pm, PassManager& pm,
EmissionTargetType& emissionTarget, EmissionTargetType& emissionTarget,
std::string outputNameNoExt) { std::string outputNameNoExt) {
verifyExplicitPimCoreCount(); verifyPimCompilerOptions();
verifyPimPipelineStages();
spatial::SchedulingTarget schedulingTarget = getPimSchedulingTarget(); spatial::SchedulingTarget schedulingTarget = getPimSchedulingTarget();
spatial::SpatialTargetResources targetResources = getPimSpatialTargetResources(schedulingTarget); spatial::SpatialTargetResources targetResources = getPimSpatialTargetResources(schedulingTarget);
@@ -350,7 +351,8 @@ void addPassesPim(OwningOpRef<ModuleOp>& module,
spatial::SpatialDataflowExportStage exportStage = spatial::SpatialDataflowExportStage exportStage =
getPimSpatialDataflowExportStage(pimExportSpatialDataflow.getValue()); getPimSpatialDataflowExportStage(pimExportSpatialDataflow.getValue());
pm.addPass(createONNXToSpatialPass(targetResources, planningOptions)); pm.addPass(createONNXToSpatialPass(targetResources, planningOptions));
pm.addPass(createSpatialLayoutPlanningPass(targetResources)); pm.addPass(createSpatialLayoutPlanningPass(
targetResources, pimDisableSpatialPlanning.getValue()));
pm.addPass(createLowerSpatialPlansPass(targetResources, planningOptions, exportStage)); pm.addPass(createLowerSpatialPlansPass(targetResources, planningOptions, exportStage));
pm.addPass(createTrivialGraphComputeMergePass( pm.addPass(createTrivialGraphComputeMergePass(
schedulingTarget.residentWeightCapacity, exportStage)); schedulingTarget.residentWeightCapacity, exportStage));
+3 -3
View File
@@ -46,7 +46,7 @@ static LogicalResult compileCoreEmissionPlan(Block& block, SmallVectorImpl<Compi
auto upper = compileIndexExpr(forOp.getUpperBound()); auto upper = compileIndexExpr(forOp.getUpperBound());
auto step = compileIndexExpr(forOp.getStep()); auto step = compileIndexExpr(forOp.getStep());
if (failed(lower) || failed(upper) || failed(step)) { if (failed(lower) || failed(upper) || failed(step)) {
forOp.emitOpError("requires statically evaluable scf.for bounds for PIM codegen"); forOp.emitOpError("requires statically evaluable scf.for bounds for Pim codegen");
return failure(); return failure();
} }
CompiledCoreNode node; CompiledCoreNode node;
@@ -63,7 +63,7 @@ static LogicalResult compileCoreEmissionPlan(Block& block, SmallVectorImpl<Compi
if (auto ifOp = dyn_cast<scf::IfOp>(op)) { if (auto ifOp = dyn_cast<scf::IfOp>(op)) {
auto condition = compileIndexExpr(ifOp.getCondition()); auto condition = compileIndexExpr(ifOp.getCondition());
if (failed(condition)) { if (failed(condition)) {
ifOp.emitOpError("requires statically evaluable scf.if condition for PIM codegen"); ifOp.emitOpError("requires statically evaluable scf.if condition for Pim codegen");
return failure(); return failure();
} }
CompiledCoreNode node; CompiledCoreNode node;
@@ -82,7 +82,7 @@ static LogicalResult compileCoreEmissionPlan(Block& block, SmallVectorImpl<Compi
if (auto switchOp = dyn_cast<scf::IndexSwitchOp>(op)) { if (auto switchOp = dyn_cast<scf::IndexSwitchOp>(op)) {
auto selector = compileIndexExpr(switchOp.getArg()); auto selector = compileIndexExpr(switchOp.getArg());
if (failed(selector)) { if (failed(selector)) {
switchOp.emitOpError("requires a statically evaluable scf.index_switch selector for PIM codegen"); switchOp.emitOpError("requires a statically evaluable scf.index_switch selector for Pim codegen");
return failure(); return failure();
} }
CompiledCoreNode node; CompiledCoreNode node;
@@ -249,7 +249,7 @@ auto createEmptySpatGraphComputeBatch(RewriterT& rewriter,
if (laneCount <= 0 || laneCount > std::numeric_limits<int32_t>::max()) if (laneCount <= 0 || laneCount > std::numeric_limits<int32_t>::max())
return mlir::FailureOr<spatial::SpatGraphComputeBatch>(mlir::failure()); return mlir::FailureOr<spatial::SpatGraphComputeBatch>(mlir::failure());
auto laneCountAttr = pim::getCheckedI32Attr(rewriter, loc, laneCount, "spatial compute_batch lane count"); auto laneCountAttr = pim::getCheckedI32Attr(rewriter, loc, laneCount, "Spatial compute_batch lane count");
if (mlir::failed(laneCountAttr)) if (mlir::failed(laneCountAttr))
return mlir::FailureOr<spatial::SpatGraphComputeBatch>(mlir::failure()); return mlir::FailureOr<spatial::SpatGraphComputeBatch>(mlir::failure());
@@ -1,39 +0,0 @@
#include "ContractionMaterialization.hpp"
#include "src/Accelerators/PIM/Common/IR/ConstantUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/CompileTime.hpp"
#include "MatrixProductLowering.hpp"
namespace onnx_mlir {
mlir::Value materializePaddedContractionInput(
mlir::Value input,
mlir::RankedTensorType paddedType,
mlir::PatternRewriter& rewriter,
mlir::Location loc) {
return createPaddedInputCompute(input, paddedType, rewriter, loc);
}
mlir::FailureOr<mlir::Value> materializeTransposedContractionConstant(
mlir::Value input,
mlir::RankedTensorType resultType,
llvm::ArrayRef<int64_t> permutation,
mlir::PatternRewriter& rewriter,
mlir::Location loc) {
auto denseAttr = getHostConstDenseElementsAttr(input);
auto inputType = denseAttr ? mlir::dyn_cast<mlir::RankedTensorType>(denseAttr.getType()) : nullptr;
if (!inputType || !inputType.hasStaticShape() || !resultType || !resultType.hasStaticShape()
|| inputType.getRank() != resultType.getRank())
return mlir::failure();
auto transposedAttr = transposeDenseElementsAttr(denseAttr, permutation);
if (mlir::failed(transposedAttr) || transposedAttr->getType() != resultType)
return mlir::failure();
return getOrCreateConstant(rewriter,
rewriter.getInsertionBlock()->getParentOp(),
*transposedAttr,
resultType);
}
} // namespace onnx_mlir
@@ -1,23 +0,0 @@
#pragma once
#include "llvm/ADT/ArrayRef.h"
#include "mlir/IR/BuiltinTypes.h"
#include "mlir/IR/PatternMatch.h"
namespace onnx_mlir {
mlir::Value materializePaddedContractionInput(
mlir::Value input,
mlir::RankedTensorType paddedType,
mlir::PatternRewriter& rewriter,
mlir::Location loc);
mlir::FailureOr<mlir::Value> materializeTransposedContractionConstant(
mlir::Value input,
mlir::RankedTensorType resultType,
llvm::ArrayRef<int64_t> permutation,
mlir::PatternRewriter& rewriter,
mlir::Location loc);
} // namespace onnx_mlir
@@ -25,7 +25,7 @@ llvm::SmallVector<mlir::Value> sliceVector(const mlir::Value& vectorToSlice,
mlir::Location loc); mlir::Location loc);
/// Partitions one logical vector into per-core crossbar-sized slices using the /// Partitions one logical vector into per-core crossbar-sized slices using the
/// current PIM target geometry. /// current Pim target geometry.
llvm::DenseMap<CoreId, llvm::SmallVector<mlir::Value>> sliceVectorPerCrossbarPerCore( llvm::DenseMap<CoreId, llvm::SmallVector<mlir::Value>> sliceVectorPerCrossbarPerCore(
const mlir::Value& vectorToSlice, const mlir::Value& vectorToSlice,
mlir::PatternRewriter& rewriter, mlir::PatternRewriter& rewriter,
@@ -1,902 +0,0 @@
#include "mlir/Dialect/Affine/IR/AffineOps.h"
#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/Linalg/IR/Linalg.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/Dialect/Tensor/IR/Tensor.h"
#include "mlir/Pass/Pass.h"
#include "mlir/Transforms/GreedyPatternRewriteDriver.h"
#include "mlir/Transforms/DialectConversion.h"
#include "Conversion/ONNXToSpatial/ONNXToSpatialVerifier.hpp"
#include "mlir/Transforms/Passes.h"
#include "src/Accelerators/PIM/Common/PimCommon.hpp"
#include "src/Accelerators/PIM/Common/Support/DebugDump.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/BiasAddUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/Common.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/MatrixProductLowering.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Patterns.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp"
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
#include "src/Accelerators/PIM/Dialect/Spatial/Transforms/MergeComputeNodes/SpatialDataflowCsvExporter.hpp"
#include "src/Accelerators/PIM/Pass/PIMPasses.h"
using namespace mlir;
namespace onnx_mlir {
namespace {
static FailureOr<RowStripPhysicalValue> getRowStripValue(Value value) {
return getRowStripPhysicalValue(value);
}
static FailureOr<Value> publishRowStripValue(Operation* planOp,
Value storage,
PatternRewriter& rewriter) {
auto logicalType = dyn_cast<RankedTensorType>(planOp->getResult(0).getType());
if (!logicalType)
return planOp->emitOpError("requires ranked logical output type"), failure();
FailureOr<RowStripPhysicalValue> value = describeRowStripPhysicalValue(storage, logicalType);
if (failed(value))
return planOp->emitOpError("lowering produced invalid row-strip physical storage"), failure();
FailureOr<Value> blueprint = createRowStripStorageBlueprint(
storage, logicalType, rewriter, planOp->getLoc());
if (failed(blueprint))
return planOp->emitOpError("failed to create row-strip storage Blueprint"), failure();
rewriter.replaceOp(planOp, *blueprint);
return *blueprint;
}
static bool isRowStripSelected(Operation* op) {
auto selected = spatial::getSelectedPhysicalLayout(op);
return selected && *selected == spatial::PhysicalLayout::NHWCRowStrip;
}
static bool isDenseSelected(Operation* op) {
auto selected = spatial::getSelectedPhysicalLayout(op);
return selected && *selected == spatial::PhysicalLayout::DenseNCHW;
}
static spatial::PhysicalLayout getKnownPhysicalLayout(Value value) {
if (auto materialize = value.getDefiningOp<spatial::SpatMaterializeLayoutOp>())
return materialize.getTargetPhysicalLayout();
if (auto blueprint = value.getDefiningOp<spatial::SpatBlueprintOp>())
return blueprint.getPhysicalLayout();
if (Operation* producer = value.getDefiningOp()) {
if (auto selected = spatial::getSelectedPhysicalLayout(producer))
return *selected;
}
return spatial::PhysicalLayout::DenseNCHW;
}
static LogicalResult verifySelectedLayouts(
func::FuncOp funcOp, const spatial::SpatialTargetInfo& target) {
LogicalResult result = success();
funcOp.walk([&](Operation* op) {
auto capability = dyn_cast<spatial::SpatialLayoutCapabilityInterface>(op);
if (!capability)
return;
auto selected = spatial::getSelectedPhysicalLayout(op);
if (!selected) {
op->emitOpError("requires a selected physical layout from SpatialLayoutPlanning");
result = failure();
return;
}
if (*selected != spatial::PhysicalLayout::DenseNCHW
&& *selected != spatial::PhysicalLayout::NHWCRowStrip) {
op->emitOpError("has an unsupported selected physical layout");
result = failure();
return;
}
SmallVector<spatial::PhysicalLayout> operandLayouts;
operandLayouts.reserve(op->getNumOperands());
for (Value operand : op->getOperands())
operandLayouts.push_back(getKnownPhysicalLayout(operand));
auto alternatives = capability.getLayoutAlternatives(target, operandLayouts);
if (llvm::none_of(alternatives, [&](const spatial::LayoutAlternative& alternative) {
return alternative.resultLayout == *selected
&& alternative.operandLayouts == operandLayouts;
})) {
op->emitOpError("selected physical layout is not lowerable for its explicit operand layouts");
result = failure();
}
});
return result;
}
static FailureOr<Value>
lowerRowStripRelu(const RowStripPhysicalValue& input, spatial::SpatReluPlanOp planOp, PatternRewriter& rewriter) {
return applyRowStripRelu(input, rewriter, planOp.getLoc());
}
static FailureOr<Value>
lowerRowStripSilu(const RowStripPhysicalValue& input, spatial::SpatSiluPlanOp planOp, PatternRewriter& rewriter) {
return applyRowStripSilu(input, rewriter, planOp.getLoc());
}
static FailureOr<Value> lowerRowStripBiasAdd(const RowStripPhysicalValue& input,
spatial::SpatBiasAddPlanOp planOp,
PatternRewriter& rewriter) {
return applyRowStripBiasAdd(input, planOp.getBias(), rewriter, planOp.getLoc());
}
static FailureOr<Value> lowerRowStripAdd(const RowStripPhysicalValue& lhs,
const RowStripPhysicalValue& rhs,
spatial::SpatAddPlanOp planOp,
PatternRewriter& rewriter) {
return applyRowStripAdd(lhs, rhs, rewriter, planOp.getLoc());
}
static FailureOr<Value> lowerRowStripConcat(ArrayRef<RowStripPhysicalValue> inputs,
spatial::SpatConcatPlanOp planOp,
PatternRewriter& rewriter) {
auto outputType = dyn_cast<RankedTensorType>(planOp.getOutput().getType());
if (!outputType)
return failure();
return applyRowStripConcat(inputs, outputType, rewriter, planOp.getLoc());
}
static FailureOr<Value>
materializeRowStripToDense(const RowStripPhysicalValue& rowStripValue, Location loc, PatternRewriter& rewriter) {
if (rowStripValue.logicalType.getRank() != 4 || !rowStripValue.logicalType.hasStaticShape())
return failure();
return createRowStripAssemblyBlueprint(rowStripValue, rewriter, loc);
}
static FailureOr<Value> materializeDenseToRowStrip(
Value input, RankedTensorType logicalType, Location loc, PatternRewriter& rewriter) {
if (!logicalType || !logicalType.hasStaticShape() || logicalType.getRank() != 4
|| logicalType.getDimSize(0) != 1)
return failure();
auto nhwcType = RankedTensorType::get(
{1, logicalType.getDimSize(2), logicalType.getDimSize(3), logicalType.getDimSize(1)},
logicalType.getElementType(), logicalType.getEncoding());
auto rowsType = RankedTensorType::get(
{logicalType.getDimSize(2) * logicalType.getDimSize(3), logicalType.getDimSize(1)},
logicalType.getElementType(), logicalType.getEncoding());
auto rowsCompute = createSpatCompute<1>(
rewriter, loc, rowsType, {}, input, [&](Value denseInput) {
Value nhwc = createLinalgTranspose(
denseInput, nhwcType, {0, 2, 3, 1}, rewriter, loc);
Value rows = tensor::CollapseShapeOp::create(
rewriter, loc, rowsType, nhwc,
SmallVector<ReassociationIndices> {{0, 1, 2}, {3}});
spatial::SpatYieldOp::create(rewriter, loc, rows);
});
Value rows = rowsCompute->getResult(0);
FailureOr<Value> storage = createRowStripStorageFromRows(rows, logicalType, rewriter, loc);
if (failed(storage))
return failure();
return createRowStripStorageBlueprint(*storage, logicalType, rewriter, loc);
}
static FailureOr<Value> lowerDenseBatchBiasAdd(Value input, Value bias, RankedTensorType resultType,
PatternRewriter& rewriter, Location loc) {
auto producer = input.getDefiningOp<spatial::SpatGraphComputeBatch>();
auto inputType = dyn_cast<RankedTensorType>(input.getType());
auto biasType = dyn_cast<RankedTensorType>(bias.getType());
if (!producer || !inputType || !biasType || !inputType.hasStaticShape() || !biasType.hasStaticShape()
|| !resultType.hasStaticShape() || inputType.getDimSize(0) != producer.getLaneCount()
|| biasType.getDimSize(0) != producer.getLaneCount() || resultType.getDimSize(0) != producer.getLaneCount())
return failure();
auto inputFragmentType = spatial::getGraphBatchFragmentType(inputType, producer.getLaneCount());
auto outputFragmentType = spatial::getGraphBatchFragmentType(resultType, producer.getLaneCount());
if (failed(inputFragmentType) || failed(outputFragmentType) || inputFragmentType->getRank() != biasType.getRank()
|| inputFragmentType->getDimSize(0) != 1 || inputFragmentType->getShape().drop_front() != biasType.getShape().drop_front()
|| inputFragmentType->getRank() != outputFragmentType->getRank() + 1)
return failure();
for (auto [inputDim, outputDim] : llvm::zip(inputFragmentType->getShape().drop_front(), outputFragmentType->getShape()))
if (outputDim > inputDim)
return failure();
auto batch = createSpatComputeBatch(rewriter, loc, TypeRange {resultType}, producer.getLaneCount(), {}, ValueRange {input, bias},
[&](detail::SpatComputeBatchBodyArgs args) -> LogicalResult {
FailureOr<Value> fragment = extractGraphBatchPhysicalFragment(rewriter, loc, args.inputs[0], args.lane, *inputFragmentType);
if (failed(fragment))
return failure();
MixedSliceGeometry biasSlice;
for (int64_t dim : inputFragmentType->getShape()) {
biasSlice.offsets.push_back(biasSlice.offsets.empty() ? OpFoldResult(args.lane) : rewriter.getIndexAttr(0));
biasSlice.sizes.push_back(rewriter.getIndexAttr(dim));
biasSlice.strides.push_back(rewriter.getIndexAttr(1));
}
Value biasFragment = extractMixedSliceOrIdentity(rewriter, loc, args.inputs[1], *inputFragmentType, biasSlice);
if (!biasFragment)
return failure();
Value added = spatial::SpatVAddOp::create(rewriter, loc, *inputFragmentType, *fragment, biasFragment);
MixedSliceGeometry outputSlice;
outputSlice.offsets.assign(inputFragmentType->getRank(), rewriter.getIndexAttr(0));
outputSlice.sizes.push_back(rewriter.getIndexAttr(1));
outputSlice.strides.assign(inputFragmentType->getRank(), rewriter.getIndexAttr(1));
for (int64_t dim : outputFragmentType->getShape())
outputSlice.sizes.push_back(rewriter.getIndexAttr(dim));
Value output = extractMixedSliceOrIdentity(rewriter, loc, added, *outputFragmentType, outputSlice);
if (!output)
return failure();
publishGraphBatchPhysicalFragment(rewriter, loc, output, args.outputs.front(), args.lane);
return success();
});
if (failed(batch))
return failure();
return batch->getResult(0);
}
struct LowerDenseReluPlan final : OpRewritePattern<spatial::SpatReluPlanOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatReluPlanOp planOp,
PatternRewriter& rewriter) const override {
auto selected = spatial::getSelectedPhysicalLayout(planOp.getOperation());
if (!selected || *selected != spatial::PhysicalLayout::DenseNCHW)
return failure();
auto computeOp = createSpatCompute<1>(
rewriter, planOp.getLoc(), planOp.getOutput().getType(), {}, planOp.getInput(), [&](Value x) {
auto relu = spatial::SpatReluOp::create(rewriter, planOp.getLoc(), planOp.getOutput().getType(), x);
spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), relu.getResult());
});
rewriter.replaceOp(planOp, computeOp.getResults());
return success();
}
};
struct LowerDenseSiluPlan final : OpRewritePattern<spatial::SpatSiluPlanOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatSiluPlanOp planOp,
PatternRewriter& rewriter) const override {
auto selected = spatial::getSelectedPhysicalLayout(planOp.getOperation());
if (!selected || *selected != spatial::PhysicalLayout::DenseNCHW)
return failure();
auto computeOp = createSpatCompute<1>(
rewriter, planOp.getLoc(), planOp.getOutput().getType(), {}, planOp.getInput(), [&](Value x) {
Value sigmoid = spatial::SpatSigmoidOp::create(
rewriter, planOp.getLoc(), planOp.getOutput().getType(), x).getResult();
Value silu = spatial::SpatVMulOp::create(
rewriter, planOp.getLoc(), planOp.getOutput().getType(), x, sigmoid).getResult();
spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), silu);
});
rewriter.replaceOp(planOp, computeOp.getResults());
return success();
}
};
struct LowerDenseResizePlan final : OpRewritePattern<spatial::SpatResizeNearestPlanOp> {
explicit LowerDenseResizePlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target)
: OpRewritePattern<spatial::SpatResizeNearestPlanOp>(ctx), target(target) {}
LogicalResult matchAndRewrite(spatial::SpatResizeNearestPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isDenseSelected(planOp.getOperation()))
return failure();
FailureOr<Value> lowered = lowerSelectedResizeNearestPlan(planOp, std::nullopt, target, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected dense nearest Resize plan");
rewriter.replaceOp(planOp, *lowered);
return success();
}
const spatial::SpatialTargetInfo& target;
};
struct LowerDenseBiasAddPlan final : OpRewritePattern<spatial::SpatBiasAddPlanOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatBiasAddPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isDenseSelected(planOp.getOperation()))
return failure();
auto resultType = dyn_cast<RankedTensorType>(planOp.getOutput().getType());
if (!resultType)
return planOp.emitOpError("requires ranked output type");
FailureOr<Value> denseBias = materializeDenseBiasAddTensor(
planOp.getBias(), resultType, rewriter, planOp.getLoc());
if (failed(denseBias))
return planOp.emitOpError("failed to materialize dense Conv-style bias");
if (planOp.getInput().getDefiningOp<spatial::SpatGraphComputeBatch>()) {
FailureOr<Value> lowered = lowerDenseBatchBiasAdd(
planOp.getInput(), *denseBias, resultType, rewriter, planOp.getLoc());
if (succeeded(lowered)) {
rewriter.replaceOp(planOp, *lowered);
return success();
}
}
auto computeOp = createSpatCompute<2>(
rewriter,
planOp.getLoc(),
planOp.getOutput().getType(),
{},
ValueRange {planOp.getInput(), *denseBias},
[&](Value x, Value y) {
auto added = spatial::SpatVAddOp::create(
rewriter, planOp.getLoc(), planOp.getOutput().getType(), x, y);
spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), added.getResult());
});
rewriter.replaceOp(planOp, computeOp.getResults());
return success();
}
};
struct LowerDenseAddPlan final : OpRewritePattern<spatial::SpatAddPlanOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatAddPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isDenseSelected(planOp.getOperation()))
return failure();
auto compute = createSpatCompute<2>(
rewriter,
planOp.getLoc(),
planOp.getOutput().getType(),
{},
ValueRange {planOp.getLhs(), planOp.getRhs()},
[&](Value lhsValue, Value rhsValue) {
Value added = spatial::SpatVAddOp::create(
rewriter, planOp.getLoc(), planOp.getOutput().getType(), lhsValue, rhsValue);
spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), added);
});
rewriter.replaceOp(planOp, compute.getResults());
return success();
}
};
struct LowerDenseConcatPlan final : OpRewritePattern<spatial::SpatConcatPlanOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatConcatPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isDenseSelected(planOp.getOperation()))
return failure();
auto compute = createSpatCompute(
rewriter,
planOp.getLoc(),
TypeRange {planOp.getOutput().getType()},
{},
planOp.getInputs(),
[&](ValueRange values) {
Value concatenated = spatial::SpatConcatOp::create(
rewriter,
planOp.getLoc(),
planOp.getOutput().getType(),
rewriter.getI64IntegerAttr(planOp.getAxis()),
values);
spatial::SpatYieldOp::create(rewriter, planOp.getLoc(), concatenated);
});
rewriter.replaceOp(planOp, compute.getResults());
return success();
}
};
static LogicalResult lowerAddPlan(spatial::SpatAddPlanOp planOp,
PatternRewriter& rewriter) {
FailureOr<RowStripPhysicalValue> lhs = getRowStripValue(planOp.getLhs());
FailureOr<RowStripPhysicalValue> rhs = getRowStripValue(planOp.getRhs());
if (isRowStripSelected(planOp.getOperation()) && failed(lhs)) {
if (getKnownPhysicalLayout(planOp.getLhs()) == spatial::PhysicalLayout::NHWCRowStrip)
return failure();
return planOp.emitOpError("selected row-strip Add plan requires row-strip inputs");
}
if (isRowStripSelected(planOp.getOperation()) && failed(rhs)) {
if (getKnownPhysicalLayout(planOp.getRhs()) == spatial::PhysicalLayout::NHWCRowStrip)
return failure();
return planOp.emitOpError("selected row-strip Add plan requires row-strip inputs");
}
if (isRowStripSelected(planOp.getOperation())) {
rewriter.setInsertionPoint(planOp);
FailureOr<Value> lowered = lowerRowStripAdd(*lhs, *rhs, planOp, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected row-strip Spatial add plan");
if (failed(publishRowStripValue(planOp, *lowered, rewriter)))
return failure();
return success();
}
return planOp.emitOpError("dense Add plan was not lowered by the selected-plan patterns");
}
static LogicalResult lowerConcatPlan(spatial::SpatConcatPlanOp planOp,
PatternRewriter& rewriter) {
SmallVector<RowStripPhysicalValue> inputs;
for (Value input : planOp.getInputs()) {
FailureOr<RowStripPhysicalValue> physical = getRowStripValue(input);
if (failed(physical)) {
inputs.clear();
break;
}
inputs.push_back(*physical);
}
if (isRowStripSelected(planOp.getOperation()) && inputs.size() != planOp.getInputs().size()) {
if (llvm::any_of(planOp.getInputs(), [](Value input) {
return getKnownPhysicalLayout(input) == spatial::PhysicalLayout::NHWCRowStrip;
}))
return failure();
return planOp.emitOpError("selected row-strip Concat plan requires row-strip inputs");
}
if (isRowStripSelected(planOp.getOperation())) {
rewriter.setInsertionPoint(planOp);
FailureOr<Value> lowered = lowerRowStripConcat(inputs, planOp, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected row-strip Spatial concat plan");
if (failed(publishRowStripValue(planOp, *lowered, rewriter)))
return failure();
return success();
}
return planOp.emitOpError("dense Concat plan was not lowered by the selected-plan patterns");
}
struct LowerSelectedConvPlan final : OpRewritePattern<spatial::SpatConv2DPlanOp> {
explicit LowerSelectedConvPlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target)
: OpRewritePattern<spatial::SpatConv2DPlanOp>(ctx), target(target) {}
LogicalResult matchAndRewrite(spatial::SpatConv2DPlanOp planOp,
PatternRewriter& rewriter) const override {
if (isDenseSelected(planOp.getOperation())) {
FailureOr<Value> lowered = lowerSelectedConv2DPlan(
planOp, std::nullopt, /*emitRowStripLayout=*/false, target, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected dense Spatial Conv plan");
rewriter.replaceOp(planOp, *lowered);
return success();
}
if (!isRowStripSelected(planOp.getOperation()))
return failure();
FailureOr<RowStripPhysicalValue> rowStripInput = getRowStripValue(planOp.getInput());
if (failed(rowStripInput)
&& getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip)
return failure();
std::optional<Value> physicalInput;
if (succeeded(rowStripInput))
physicalInput = rowStripInput->storage;
FailureOr<Value> lowered = lowerSelectedConv2DPlan(
planOp, physicalInput, /*emitRowStripLayout=*/true, target, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected row-strip Spatial Conv plan");
if (failed(publishRowStripValue(planOp, *lowered, rewriter)))
return failure();
return success();
}
const spatial::SpatialTargetInfo& target;
};
struct LowerRowStripReluPlan final : OpRewritePattern<spatial::SpatReluPlanOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatReluPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isRowStripSelected(planOp.getOperation()))
return failure();
FailureOr<RowStripPhysicalValue> input = getRowStripValue(planOp.getInput());
if (failed(input)) {
if (getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip)
return failure();
return planOp.emitOpError("selected row-strip ReLU plan requires a row-strip input");
}
FailureOr<Value> lowered = lowerRowStripRelu(*input, planOp, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected row-strip Spatial ReLU plan");
if (failed(publishRowStripValue(planOp, *lowered, rewriter)))
return failure();
return success();
}
};
struct LowerRowStripSiluPlan final : OpRewritePattern<spatial::SpatSiluPlanOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatSiluPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isRowStripSelected(planOp.getOperation()))
return failure();
FailureOr<RowStripPhysicalValue> input = getRowStripValue(planOp.getInput());
if (failed(input)) {
if (getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip)
return failure();
return planOp.emitOpError("selected row-strip SiLU plan requires a row-strip input");
}
FailureOr<Value> lowered = lowerRowStripSilu(*input, planOp, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected row-strip Spatial SiLU plan");
if (failed(publishRowStripValue(planOp, *lowered, rewriter)))
return failure();
return success();
}
};
struct LowerRowStripResizePlan final : OpRewritePattern<spatial::SpatResizeNearestPlanOp> {
explicit LowerRowStripResizePlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target)
: OpRewritePattern<spatial::SpatResizeNearestPlanOp>(ctx), target(target) {}
LogicalResult matchAndRewrite(spatial::SpatResizeNearestPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isRowStripSelected(planOp.getOperation()))
return failure();
FailureOr<RowStripPhysicalValue> input = getRowStripValue(planOp.getInput());
if (failed(input)) {
if (getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip)
return failure();
return planOp.emitOpError("selected row-strip Resize plan requires a row-strip input");
}
FailureOr<Value> lowered = lowerSelectedResizeNearestPlan(planOp, input->storage, target, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected row-strip Resize plan");
if (failed(publishRowStripValue(planOp, *lowered, rewriter)))
return failure();
return success();
}
const spatial::SpatialTargetInfo& target;
};
struct LowerDenseMaxPoolPlan final : OpRewritePattern<spatial::SpatMaxPool2DPlanOp> {
explicit LowerDenseMaxPoolPlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target)
: OpRewritePattern<spatial::SpatMaxPool2DPlanOp>(ctx), target(target) {}
LogicalResult matchAndRewrite(spatial::SpatMaxPool2DPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isDenseSelected(planOp.getOperation()))
return failure();
FailureOr<Value> lowered = lowerDenseMaxPool2DPlan(planOp, target, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected dense Spatial MaxPool plan");
rewriter.replaceOp(planOp, *lowered);
return success();
}
const spatial::SpatialTargetInfo& target;
};
struct LowerRowStripMaxPoolPlan final : OpRewritePattern<spatial::SpatMaxPool2DPlanOp> {
explicit LowerRowStripMaxPoolPlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target)
: OpRewritePattern<spatial::SpatMaxPool2DPlanOp>(ctx), target(target) {}
LogicalResult matchAndRewrite(spatial::SpatMaxPool2DPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isRowStripSelected(planOp.getOperation()))
return failure();
FailureOr<RowStripPhysicalValue> input = getRowStripValue(planOp.getInput());
if (failed(input)
&& getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip)
return failure();
std::optional<Value> physicalInput;
if (succeeded(input))
physicalInput = input->storage;
FailureOr<Value> lowered = lowerSelectedMaxPool2DPlan(planOp, physicalInput, target, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected row-strip Spatial MaxPool plan");
if (failed(publishRowStripValue(planOp, *lowered, rewriter)))
return failure();
return success();
}
const spatial::SpatialTargetInfo& target;
};
struct LowerRowStripGlobalAveragePoolPlan
final : OpRewritePattern<spatial::SpatGlobalAveragePoolPlanOp> {
explicit LowerRowStripGlobalAveragePoolPlan(MLIRContext* ctx, const spatial::SpatialTargetInfo& target)
: OpRewritePattern<spatial::SpatGlobalAveragePoolPlanOp>(ctx), target(target) {}
LogicalResult matchAndRewrite(spatial::SpatGlobalAveragePoolPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isRowStripSelected(planOp.getOperation()))
return failure();
FailureOr<RowStripPhysicalValue> input = getRowStripValue(planOp.getInput());
if (failed(input)
&& getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip)
return failure();
std::optional<Value> physicalInput;
if (succeeded(input))
physicalInput = input->storage;
FailureOr<Value> lowered = lowerSelectedGlobalAveragePoolPlan(planOp, physicalInput, target, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected row-strip Spatial global AveragePool plan");
if (failed(publishRowStripValue(planOp, *lowered, rewriter)))
return failure();
return success();
}
const spatial::SpatialTargetInfo& target;
};
struct LowerDenseGlobalAveragePoolPlan
final : OpRewritePattern<spatial::SpatGlobalAveragePoolPlanOp> {
explicit LowerDenseGlobalAveragePoolPlan(MLIRContext* ctx,
const spatial::SpatialTargetInfo& target)
: OpRewritePattern<spatial::SpatGlobalAveragePoolPlanOp>(ctx), target(target) {}
LogicalResult matchAndRewrite(spatial::SpatGlobalAveragePoolPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isDenseSelected(planOp.getOperation()))
return failure();
FailureOr<Value> lowered = lowerDenseGlobalAveragePoolPlan(planOp, target, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected dense Spatial global AveragePool plan");
rewriter.replaceOp(planOp, *lowered);
return success();
}
const spatial::SpatialTargetInfo& target;
};
struct LowerRowStripBiasAddPlan final : OpRewritePattern<spatial::SpatBiasAddPlanOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatBiasAddPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isRowStripSelected(planOp.getOperation()))
return failure();
FailureOr<RowStripPhysicalValue> input = getRowStripValue(planOp.getInput());
if (failed(input)) {
if (getKnownPhysicalLayout(planOp.getInput()) == spatial::PhysicalLayout::NHWCRowStrip)
return failure();
return planOp.emitOpError("selected row-strip bias_add plan requires a row-strip input");
}
FailureOr<Value> lowered = lowerRowStripBiasAdd(*input, planOp, rewriter);
if (failed(lowered))
return planOp.emitOpError("failed to lower selected row-strip Spatial bias_add plan");
if (failed(publishRowStripValue(planOp, *lowered, rewriter)))
return failure();
return success();
}
};
struct LowerRowStripAddPlan final : OpRewritePattern<spatial::SpatAddPlanOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatAddPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isRowStripSelected(planOp.getOperation()))
return failure();
return lowerAddPlan(planOp, rewriter);
}
};
struct LowerRowStripConcatPlan final : OpRewritePattern<spatial::SpatConcatPlanOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatConcatPlanOp planOp,
PatternRewriter& rewriter) const override {
if (!isRowStripSelected(planOp.getOperation()))
return failure();
return lowerConcatPlan(planOp, rewriter);
}
};
struct LowerMaterializeLayout final
: OpRewritePattern<spatial::SpatMaterializeLayoutOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatMaterializeLayoutOp materializeOp,
PatternRewriter& rewriter) const override {
auto source = materializeOp.getSourcePhysicalLayout();
auto target = materializeOp.getTargetPhysicalLayout();
if (source == spatial::PhysicalLayout::DenseNCHW
&& target == spatial::PhysicalLayout::DenseNCHW) {
rewriter.replaceOp(materializeOp, materializeOp.getInput());
return success();
}
if (source == spatial::PhysicalLayout::DenseNCHW
&& target == spatial::PhysicalLayout::NHWCRowStrip) {
auto logicalType = dyn_cast<RankedTensorType>(materializeOp.getInput().getType());
if (!logicalType)
return materializeOp.emitOpError("requires a ranked dense input"), failure();
FailureOr<Value> rowStrip = materializeDenseToRowStrip(
materializeOp.getInput(), logicalType, materializeOp.getLoc(), rewriter);
if (failed(rowStrip))
return materializeOp.emitOpError(
"failed to materialize dense NCHW storage to row-strip layout"), failure();
rewriter.replaceOp(materializeOp, *rowStrip);
return success();
}
if (source != spatial::PhysicalLayout::NHWCRowStrip
|| target != spatial::PhysicalLayout::DenseNCHW)
return materializeOp.emitOpError(
"unsupported Spatial layout materialization direction"), failure();
auto inputType = dyn_cast<RankedTensorType>(materializeOp.getInput().getType());
if (!inputType)
return materializeOp.emitOpError("requires a ranked row-strip input"), failure();
FailureOr<RowStripPhysicalValue> rowStripValue =
getRowStripValue(materializeOp.getInput());
if (failed(rowStripValue))
return materializeOp.emitOpError(
"requires an explicitly defining row-strip physical value"), failure();
FailureOr<Value> dense = materializeRowStripToDense(
*rowStripValue, materializeOp.getLoc(), rewriter);
if (failed(dense))
return materializeOp.emitOpError(
"failed to materialize row-strip storage to dense NCHW"), failure();
rewriter.replaceOp(materializeOp, *dense);
return success();
}
};
struct LowerRowStripFlatten final
: OpRewritePattern<spatial::SpatGraphCompute> {
explicit LowerRowStripFlatten(MLIRContext* context,
const spatial::SpatialTargetInfo& target)
: OpRewritePattern<spatial::SpatGraphCompute>(context), target(target) {}
LogicalResult matchAndRewrite(spatial::SpatGraphCompute flattenOp,
PatternRewriter& rewriter) const override {
if (flattenOp.getInputs().size() != 1)
return failure();
FailureOr<RowStripPhysicalValue> input =
getRowStripValue(flattenOp.getInputs().front());
if (failed(input) || failed(canLowerFlattenFromRowStrip(flattenOp, target)))
return failure();
if (failed(lowerFlattenFromRowStrip(*input, flattenOp, target, rewriter)))
return flattenOp.emitOpError(
"failed to preserve row-strip layout through Flatten"), failure();
return success();
}
const spatial::SpatialTargetInfo& target;
};
struct LowerSpatialPlansPass final : PassWrapper<LowerSpatialPlansPass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(LowerSpatialPlansPass)
StringRef getArgument() const override { return "lower-spatial-plans"; }
StringRef getDescription() const override { return "Lower selected Spatial planning ops to low-level Spatial IR."; }
LowerSpatialPlansPass() = default;
explicit LowerSpatialPlansPass(const spatial::SpatialTargetInfo& target)
: target(target), hasTarget(true) {}
void runOnOperation() override {
ModuleOp moduleOp = getOperation();
if (!hasTarget) {
moduleOp.emitError("Spatial plan lowering requires an injected SpatialTargetInfo");
signalPassFailure();
return;
}
MLIRContext* ctx = moduleOp.getContext();
auto entryFunc = getPimEntryFunc(moduleOp);
if (failed(entryFunc)) {
moduleOp.emitError("failed to locate the PIM entry function during LowerSpatialPlans");
signalPassFailure();
return;
}
func::FuncOp funcOp = *entryFunc;
PatternRewriter rewriter(ctx);
auto verifyLogicalPhase = [&](StringRef stage) -> bool {
if (succeeded(verifyLogicalSpatialGraphInvariants(*entryFunc)))
return true;
moduleOp.emitError() << "logical Spatial graph verification failed " << stage;
signalPassFailure();
return false;
};
if (!verifyLogicalPhase("at the start of LowerSpatialPlans"))
return;
if (failed(verifySelectedLayouts(funcOp, target))) {
moduleOp.emitError("selected Spatial layout verification failed");
signalPassFailure();
return;
}
RewritePatternSet selectedPlanPatterns(ctx);
selectedPlanPatterns.add<LowerDenseReluPlan,
LowerRowStripReluPlan,
LowerDenseSiluPlan,
LowerRowStripSiluPlan,
LowerDenseBiasAddPlan,
LowerRowStripBiasAddPlan,
LowerDenseAddPlan,
LowerRowStripAddPlan,
LowerDenseConcatPlan,
LowerRowStripConcatPlan>(ctx);
selectedPlanPatterns.add<LowerSelectedConvPlan,
LowerDenseResizePlan,
LowerRowStripResizePlan,
LowerDenseMaxPoolPlan,
LowerRowStripMaxPoolPlan,
LowerDenseGlobalAveragePoolPlan,
LowerRowStripGlobalAveragePoolPlan>(ctx, target);
if (failed(applyPatternsGreedily(funcOp, std::move(selectedPlanPatterns)))) {
moduleOp.emitError("failed to lower selected Spatial plans");
signalPassFailure();
return;
}
RewritePatternSet layoutPatterns(ctx);
layoutPatterns.add<LowerMaterializeLayout>(ctx);
layoutPatterns.add<LowerRowStripFlatten>(ctx, target);
ConversionTarget layoutTarget(*ctx);
layoutTarget.addLegalDialect<spatial::SpatialDialect,
tensor::TensorDialect,
linalg::LinalgDialect,
affine::AffineDialect,
arith::ArithDialect,
scf::SCFDialect,
func::FuncDialect>();
layoutTarget.addIllegalDialect<ONNXDialect>();
layoutTarget.addIllegalOp<spatial::SpatMaterializeLayoutOp>();
layoutTarget.addDynamicallyLegalOp<spatial::SpatGraphCompute>(
[&](spatial::SpatGraphCompute computeOp) {
if (computeOp.getInputs().size() != 1)
return true;
FailureOr<RowStripPhysicalValue> input =
getRowStripValue(computeOp.getInputs().front());
return failed(input) || failed(canLowerFlattenFromRowStrip(computeOp, target));
});
FrozenRewritePatternSet frozenLayoutPatterns(std::move(layoutPatterns));
if (failed(applyFullConversion(funcOp, layoutTarget,
frozenLayoutPatterns))) {
moduleOp.emitError("failed to lower explicit Spatial layout materialization");
signalPassFailure();
return;
}
if (!verifyLogicalPhase("after selected-plan conversion"))
return;
SmallVector<spatial::SpatBlueprintOp> deadPhysicalViews;
funcOp.walk([&](spatial::SpatBlueprintOp blueprint) {
if (spatial::isPhysicalView(blueprint.getMode()) && blueprint.use_empty())
deadPhysicalViews.push_back(blueprint);
});
for (spatial::SpatBlueprintOp blueprint : deadPhysicalViews)
rewriter.eraseOp(blueprint);
bool hasIllegalOps = false;
moduleOp.walk([&](Operation* op) {
if (isa<ONNXEntryPointOp>(op))
return;
if (auto blueprint = dyn_cast<spatial::SpatBlueprintOp>(op)) {
if (spatial::isFragmentAssembly(blueprint.getMode()))
return;
op->emitOpError("planning blueprint must not remain after LowerSpatialPlans");
hasIllegalOps = true;
}
else if (isa<spatial::SpatConv2DPlanOp,
spatial::SpatBiasAddPlanOp,
spatial::SpatAddPlanOp,
spatial::SpatReluPlanOp,
spatial::SpatSiluPlanOp,
spatial::SpatResizeNearestPlanOp,
spatial::SpatMaxPool2DPlanOp,
spatial::SpatGlobalAveragePoolPlanOp,
spatial::SpatMaterializeLayoutOp>(op)
|| op->getDialect()->getNamespace() == "onnx") {
op->emitOpError("operation must not remain after LowerSpatialPlans");
hasIllegalOps = true;
}
});
PassManager canonicalizationPM(ctx);
canonicalizationPM.addPass(createCanonicalizerPass());
if (failed(canonicalizationPM.run(moduleOp)))
moduleOp.emitWarning("failed to run LowerSpatialPlansPass canonicalization; continuing");
if (hasIllegalOps) {
signalPassFailure();
} else {
dumpModule(moduleOp, "spatial1_graph");
spatial::SpatialDataflowExportStage exportMode = spatial::getSpatialDataflowExportStage();
if (spatial::shouldExportSpatialDataflowStage(exportMode, spatial::SpatialDataflowExportStage::Spatial1)
&& failed(spatial::exportSpatialDataflowCsvGraph(funcOp, "spatial1_graph"))) {
signalPassFailure();
return;
}
}
if (!verifyLogicalPhase("at the end of LowerSpatialPlans"))
return;
}
spatial::SpatialTargetInfo target;
bool hasTarget = false;
};
} // namespace
std::unique_ptr<Pass> createLowerSpatialPlansPass() { return std::make_unique<LowerSpatialPlansPass>(); }
std::unique_ptr<Pass> createLowerSpatialPlansPass(const spatial::SpatialTargetInfo& target) {
return std::make_unique<LowerSpatialPlansPass>(target);
}
} // namespace onnx_mlir
@@ -46,7 +46,7 @@ struct LowerSpatialPlansPass final
} }
auto entryFunc = getPimEntryFunc(moduleOp); auto entryFunc = getPimEntryFunc(moduleOp);
if (failed(entryFunc)) { if (failed(entryFunc)) {
moduleOp.emitError("failed to locate the PIM entry function during LowerSpatialPlans"); moduleOp.emitError("failed to locate the Pim entry function during LowerSpatialPlans");
signalPassFailure(); signalPassFailure();
return; return;
} }
@@ -158,7 +158,7 @@ void ONNXToSpatialPass::runOnOperation() {
auto entryFunc = getPimEntryFunc(moduleOp); auto entryFunc = getPimEntryFunc(moduleOp);
if (failed(entryFunc)) { if (failed(entryFunc)) {
moduleOp.emitError("failed to locate the PIM entry function during ONNX-to-Spatial lowering"); moduleOp.emitError("failed to locate the Pim entry function during ONNX-to-Spatial lowering");
signalPassFailure(); signalPassFailure();
return; return;
} }
@@ -245,7 +245,7 @@ void ONNXToSpatialPass::runOnOperation() {
RewritePatternSet postPatterns(ctx); RewritePatternSet postPatterns(ctx);
populatePostPatterns(postPatterns, ctx); populatePostPatterns(postPatterns, ctx);
if (failed(applyPartialConversion(*entryFunc, postTarget, std::move(postPatterns)))) { if (failed(applyPartialConversion(*entryFunc, postTarget, std::move(postPatterns)))) {
moduleOp.emitError("failed to normalize weight-like Spatial compute operands before Spatial-to-PIM lowering"); moduleOp.emitError("failed to normalize weight-like Spatial compute operands before Spatial-to-Pim lowering");
signalPassFailure(); signalPassFailure();
return; return;
} }
@@ -42,8 +42,9 @@ static SmallVector<spatial::PhysicalLayout> getOperandLayouts(
class SpatialLayoutAnalysis { class SpatialLayoutAnalysis {
public: public:
SpatialLayoutAnalysis(func::FuncOp funcOp, SpatialLayoutAnalysis(func::FuncOp funcOp,
const spatial::SpatialTargetResources& target) const spatial::SpatialTargetResources& target,
: funcOp(funcOp), target(target) {} bool selectTrivialPlan)
: funcOp(funcOp), target(target), selectTrivialPlan(selectTrivialPlan) {}
FailureOr<SpatialLayoutSelection> run() { FailureOr<SpatialLayoutSelection> run() {
SpatialLayoutSelection selection; SpatialLayoutSelection selection;
@@ -56,6 +57,9 @@ public:
selection.selectedAlternative[&op] = 0; selection.selectedAlternative[&op] = 0;
} }
if (selectTrivialPlan)
return selection;
const size_t maxRounds = 2 * planOps.size() + 1; const size_t maxRounds = 2 * planOps.size() + 1;
for (size_t round = 0; round < maxRounds; ++round) { for (size_t round = 0; round < maxRounds; ++round) {
bool changed = false; bool changed = false;
@@ -168,6 +172,7 @@ private:
func::FuncOp funcOp; func::FuncOp funcOp;
const spatial::SpatialTargetResources& target; const spatial::SpatialTargetResources& target;
bool selectTrivialPlan;
}; };
static LogicalResult materializeMismatchedUses( static LogicalResult materializeMismatchedUses(
@@ -251,8 +256,9 @@ struct SpatialLayoutPlanningPass final
} }
SpatialLayoutPlanningPass() = default; SpatialLayoutPlanningPass() = default;
explicit SpatialLayoutPlanningPass(const spatial::SpatialTargetResources& target) SpatialLayoutPlanningPass(const spatial::SpatialTargetResources& target,
: target(target), hasTarget(true) {} bool selectTrivialPlan)
: target(target), selectTrivialPlan(selectTrivialPlan), hasTarget(true) {}
void runOnOperation() override { void runOnOperation() override {
ModuleOp moduleOp = getOperation(); ModuleOp moduleOp = getOperation();
@@ -263,13 +269,13 @@ struct SpatialLayoutPlanningPass final
} }
auto entryFunc = getPimEntryFunc(moduleOp); auto entryFunc = getPimEntryFunc(moduleOp);
if (failed(entryFunc)) { if (failed(entryFunc)) {
moduleOp.emitError("failed to locate the PIM entry function during Spatial layout planning"); moduleOp.emitError("failed to locate the Pim entry function during Spatial layout planning");
signalPassFailure(); signalPassFailure();
return; return;
} }
func::FuncOp funcOp = *entryFunc; func::FuncOp funcOp = *entryFunc;
SpatialLayoutAnalysis analysis(funcOp, target); SpatialLayoutAnalysis analysis(funcOp, target, selectTrivialPlan);
FailureOr<SpatialLayoutSelection> selection = analysis.run(); FailureOr<SpatialLayoutSelection> selection = analysis.run();
if (failed(selection)) { if (failed(selection)) {
signalPassFailure(); signalPassFailure();
@@ -301,6 +307,7 @@ struct SpatialLayoutPlanningPass final
} }
spatial::SpatialTargetResources target; spatial::SpatialTargetResources target;
bool selectTrivialPlan = false;
bool hasTarget = false; bool hasTarget = false;
}; };
@@ -311,8 +318,8 @@ std::unique_ptr<Pass> createSpatialLayoutPlanningPass() {
} }
std::unique_ptr<Pass> createSpatialLayoutPlanningPass( std::unique_ptr<Pass> createSpatialLayoutPlanningPass(
const spatial::SpatialTargetResources& target) { const spatial::SpatialTargetResources& target, bool selectTrivialPlan) {
return std::make_unique<SpatialLayoutPlanningPass>(target); return std::make_unique<SpatialLayoutPlanningPass>(target, selectTrivialPlan);
} }
} // namespace onnx_mlir } // namespace onnx_mlir
@@ -199,7 +199,7 @@ static bool writeConvLoweringReport(const ConvLoweringReportEntry& entry,
return false; return false;
} }
reportFile << "# PIM Conv Lowering Report (bounded to 512 rows)\n\n"; reportFile << "# Pim conv lowering report (bounded to 512 rows)\n\n";
reportFile << "## Plan selection\n"; reportFile << "## Plan selection\n";
writeConvReportTableHeader(reportFile, "Selector"); writeConvReportTableHeader(reportFile, "Selector");
bool realizationSectionStarted = false; bool realizationSectionStarted = false;
@@ -370,6 +370,14 @@ struct ReduceMeanToSpatialCompute : OpConversionPattern<ReduceMeanOp> {
Location loc = reduceMeanOp.getLoc(); Location loc = reduceMeanOp.getLoc();
RankedTensorType leafType = getAllOnesType(inputType, resultType.getElementType()); RankedTensorType leafType = getAllOnesType(inputType, resultType.getElementType());
RankedTensorType keepdimsType = getKeepdimsType(inputType, resultType.getElementType(), reducedAxes); RankedTensorType keepdimsType = getKeepdimsType(inputType, resultType.getElementType(), reducedAxes);
if (semantics->keepdims != 0 && inputType.getRank() == 4
&& inputType.getDimSize(0) == 1 && semantics->axes == ArrayRef<int64_t>({2, 3})
&& resultType == keepdimsType) {
auto plan = spatial::SpatGlobalAveragePoolPlanOp::create(
rewriter, loc, resultType, adaptor.getData(), spatial::getNCHWLayout(rewriter.getContext()));
rewriter.replaceOp(reduceMeanOp, plan.getResult());
return success();
}
int64_t laneCount = 1; int64_t laneCount = 1;
for (auto [dim, isReduced] : llvm::zip_equal(keepdimsType.getShape(), reducedAxes)) { for (auto [dim, isReduced] : llvm::zip_equal(keepdimsType.getShape(), reducedAxes)) {
if (isReduced) if (isReduced)
@@ -1,64 +0,0 @@
#pragma once
#include <optional>
#include "mlir/IR/PatternMatch.h"
#include "mlir/Support/LogicalResult.h"
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
namespace onnx_mlir {
struct RowStripPhysicalValue;
mlir::FailureOr<mlir::Value>
lowerSelectedConv2DPlan(spatial::SpatConv2DPlanOp planOp,
std::optional<mlir::Value> rowStripInput,
bool emitRowStripLayout,
const spatial::SpatialTargetInfo& target,
mlir::PatternRewriter& rewriter);
mlir::LogicalResult canLowerConvPlanToRowStrip(spatial::SpatConv2DPlanOp planOp,
const spatial::SpatialTargetInfo& target);
mlir::LogicalResult canConsumeAndProduceRowStrip(spatial::SpatConv2DPlanOp planOp,
const spatial::SpatialTargetInfo& target);
mlir::LogicalResult canLowerResizeNearestPlanToRowStrip(
spatial::SpatResizeNearestPlanOp planOp, const spatial::SpatialTargetInfo& target);
mlir::FailureOr<mlir::Value> lowerSelectedResizeNearestPlan(
spatial::SpatResizeNearestPlanOp planOp,
std::optional<mlir::Value> rowStripInput,
const spatial::SpatialTargetInfo& target,
mlir::PatternRewriter& rewriter);
mlir::LogicalResult canLowerMaxPoolPlanToRowStrip(spatial::SpatMaxPool2DPlanOp planOp,
const spatial::SpatialTargetInfo& target);
mlir::FailureOr<mlir::Value>
lowerDenseMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
const spatial::SpatialTargetInfo& target,
mlir::PatternRewriter& rewriter);
mlir::FailureOr<mlir::Value>
lowerSelectedMaxPool2DPlan(spatial::SpatMaxPool2DPlanOp planOp,
std::optional<mlir::Value> rowStripInput,
const spatial::SpatialTargetInfo& target,
mlir::PatternRewriter& rewriter);
mlir::LogicalResult
canLowerGlobalAveragePoolPlanToRowStrip(spatial::SpatGlobalAveragePoolPlanOp planOp,
const spatial::SpatialTargetInfo& target);
mlir::FailureOr<mlir::Value>
lowerDenseGlobalAveragePoolPlan(spatial::SpatGlobalAveragePoolPlanOp planOp,
const spatial::SpatialTargetInfo& target,
mlir::PatternRewriter& rewriter);
mlir::FailureOr<mlir::Value>
lowerSelectedGlobalAveragePoolPlan(spatial::SpatGlobalAveragePoolPlanOp planOp,
std::optional<mlir::Value> rowStripInput,
const spatial::SpatialTargetInfo& target,
mlir::PatternRewriter& rewriter);
} // namespace onnx_mlir
@@ -1,133 +0,0 @@
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/BiasAddUtils.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/PlanLowering.hpp"
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
using namespace mlir;
namespace onnx_mlir::spatial {
static LayoutAlternative denseAlternative(Operation *op) {
LayoutAlternative alternative;
alternative.operandLayouts.assign(op->getNumOperands(), PhysicalLayout::DenseNCHW);
alternative.resultLayout = PhysicalLayout::DenseNCHW;
return alternative;
}
static LayoutAlternative rowStripAlternative(Operation *op,
ArrayRef<PhysicalLayout> operandLayouts) {
LayoutAlternative alternative;
alternative.operandLayouts.assign(operandLayouts.begin(), operandLayouts.end());
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
alternative.intrinsicCost = -2;
return alternative;
}
static bool hasRowStripInput(ArrayRef<PhysicalLayout> operandLayouts, unsigned index) {
return index < operandLayouts.size()
&& operandLayouts[index] == PhysicalLayout::NHWCRowStrip;
}
SmallVector<LayoutAlternative> SpatConv2DPlanOp::getLayoutAlternatives(
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (hasRowStripInput(operandLayouts, 0)) {
if (succeeded(canConsumeAndProduceRowStrip(*this, target)))
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
}
else if (succeeded(canLowerConvPlanToRowStrip(*this, target))) {
LayoutAlternative alternative = denseAlternative(getOperation());
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
alternative.intrinsicCost = -2;
alternatives.push_back(std::move(alternative));
}
return alternatives;
}
SmallVector<LayoutAlternative> SpatReluPlanOp::getLayoutAlternatives(
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (hasRowStripInput(operandLayouts, 0))
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
return alternatives;
}
SmallVector<LayoutAlternative> SpatSiluPlanOp::getLayoutAlternatives(
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (hasRowStripInput(operandLayouts, 0)) {
LayoutAlternative alternative = rowStripAlternative(getOperation(), operandLayouts);
alternative.intrinsicCost = -3;
alternatives.push_back(std::move(alternative));
}
return alternatives;
}
SmallVector<LayoutAlternative> SpatResizeNearestPlanOp::getLayoutAlternatives(
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (hasRowStripInput(operandLayouts, 0)
&& succeeded(canLowerResizeNearestPlanToRowStrip(*this, target)))
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
return alternatives;
}
SmallVector<LayoutAlternative> SpatMaxPool2DPlanOp::getLayoutAlternatives(
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (succeeded(canLowerMaxPoolPlanToRowStrip(*this, target))) {
LayoutAlternative alternative = denseAlternative(getOperation());
if (hasRowStripInput(operandLayouts, 0))
alternative = rowStripAlternative(getOperation(), operandLayouts);
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
alternative.intrinsicCost = -2;
alternatives.push_back(std::move(alternative));
}
return alternatives;
}
SmallVector<LayoutAlternative> SpatGlobalAveragePoolPlanOp::getLayoutAlternatives(
const SpatialTargetInfo& target, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (succeeded(canLowerGlobalAveragePoolPlanToRowStrip(*this, target))) {
LayoutAlternative alternative = denseAlternative(getOperation());
if (hasRowStripInput(operandLayouts, 0))
alternative = rowStripAlternative(getOperation(), operandLayouts);
alternative.resultLayout = PhysicalLayout::NHWCRowStrip;
alternative.intrinsicCost = -2;
alternatives.push_back(std::move(alternative));
}
return alternatives;
}
SmallVector<LayoutAlternative> SpatBiasAddPlanOp::getLayoutAlternatives(
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
auto resultType = dyn_cast<RankedTensorType>(getOutput().getType());
if (resultType && hasRowStripInput(operandLayouts, 0)
&& isSupportedBiasAddValue(getBias(), resultType))
alternatives.push_back(rowStripAlternative(getOperation(),
{PhysicalLayout::NHWCRowStrip,
PhysicalLayout::DenseNCHW}));
return alternatives;
}
SmallVector<LayoutAlternative> SpatAddPlanOp::getLayoutAlternatives(
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (operandLayouts.size() >= 2 && hasRowStripInput(operandLayouts, 0)
&& hasRowStripInput(operandLayouts, 1))
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
return alternatives;
}
SmallVector<LayoutAlternative> SpatConcatPlanOp::getLayoutAlternatives(
const SpatialTargetInfo&, ArrayRef<PhysicalLayout> operandLayouts) {
SmallVector<LayoutAlternative> alternatives {denseAlternative(getOperation())};
if (!operandLayouts.empty() && llvm::all_of(operandLayouts, [](PhysicalLayout layout) {
return layout == PhysicalLayout::NHWCRowStrip;
}))
alternatives.push_back(rowStripAlternative(getOperation(), operandLayouts));
return alternatives;
}
} // namespace onnx_mlir::spatial
@@ -1,265 +0,0 @@
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/IR/PatternMatch.h"
#include "mlir/Pass/Pass.h"
#include "llvm/ADT/DenseMap.h"
#include "Conversion/ONNXToSpatial/ONNXToSpatialVerifier.hpp"
#include "src/Accelerators/PIM/Common/PimCommon.hpp"
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/RowStripLayoutUtils.hpp"
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
#include "src/Accelerators/PIM/Pass/PIMPasses.h"
#include <algorithm>
using namespace mlir;
namespace onnx_mlir {
namespace {
using LayoutMap = llvm::DenseMap<Value, spatial::PhysicalLayout>;
static spatial::PhysicalLayout getSelectedLayout(const LayoutMap& layouts, Value value) {
if (auto it = layouts.find(value); it != layouts.end())
return it->second;
if (auto materialize = value.getDefiningOp<spatial::SpatMaterializeLayoutOp>())
return materialize.getTargetPhysicalLayout();
if (auto blueprint = value.getDefiningOp<spatial::SpatBlueprintOp>())
return blueprint.getPhysicalLayout();
return spatial::PhysicalLayout::DenseNCHW;
}
static SmallVector<spatial::PhysicalLayout> getOperandLayouts(
Operation* op, const LayoutMap& layouts) {
SmallVector<spatial::PhysicalLayout> operandLayouts;
operandLayouts.reserve(op->getNumOperands());
for (Value operand : op->getOperands())
operandLayouts.push_back(getSelectedLayout(layouts, operand));
return operandLayouts;
}
static FailureOr<SmallVector<spatial::LayoutAlternative>> getAlternatives(
Operation* op, const LayoutMap& layouts, const spatial::SpatialTargetInfo& target) {
auto capability = dyn_cast<spatial::SpatialLayoutCapabilityInterface>(op);
if (!capability)
return failure();
SmallVector<spatial::LayoutAlternative> alternatives =
capability.getLayoutAlternatives(target, getOperandLayouts(op, layouts));
if (alternatives.empty())
return op->emitOpError("does not advertise a legal Spatial layout alternative"), failure();
for (const spatial::LayoutAlternative& alternative : alternatives)
if (alternative.operandLayouts.size() != op->getNumOperands())
return op->emitOpError("advertises a layout alternative with the wrong operand count"), failure();
return alternatives;
}
static unsigned findCurrentAlternative(
Operation* op, ArrayRef<spatial::LayoutAlternative> alternatives,
spatial::PhysicalLayout selectedResult) {
for (auto [index, alternative] : llvm::enumerate(alternatives))
if (alternative.resultLayout == selectedResult)
return index;
return 0;
}
static int64_t alternativeCost(Operation* op,
const spatial::LayoutAlternative& alternative,
const LayoutMap& layouts,
const LayoutMap& selectedResults,
const spatial::SpatialTargetInfo& target) {
int64_t cost = alternative.intrinsicCost;
SmallVector<spatial::PhysicalLayout> operandLayouts = getOperandLayouts(op, layouts);
for (auto [actual, required] : llvm::zip(operandLayouts, alternative.operandLayouts))
cost += actual != required;
Value result = op->getResult(0);
for (OpOperand& use : result.getUses()) {
auto user = dyn_cast<spatial::SpatialLayoutCapabilityInterface>(use.getOwner());
if (!user) {
if (alternative.resultLayout != spatial::PhysicalLayout::DenseNCHW) {
auto flatten = dyn_cast<spatial::SpatGraphCompute>(use.getOwner());
if (!flatten || failed(canLowerFlattenFromRowStrip(flatten, target)))
++cost;
}
continue;
}
auto userAlternatives = getAlternatives(use.getOwner(), selectedResults, target);
if (failed(userAlternatives))
continue;
spatial::PhysicalLayout userResult =
selectedResults.lookup(use.getOwner()->getResult(0));
unsigned userIndex = findCurrentAlternative(use.getOwner(), *userAlternatives, userResult);
if (use.getOperandNumber() < (*userAlternatives)[userIndex].operandLayouts.size()
&& (*userAlternatives)[userIndex].operandLayouts[use.getOperandNumber()]
!= alternative.resultLayout)
++cost;
}
return cost;
}
static LogicalResult materializeMismatchedUses(
IRRewriter& rewriter, Value value, const LayoutMap& layouts,
const spatial::SpatialTargetInfo& target) {
spatial::PhysicalLayout sourceLayout = getSelectedLayout(layouts, value);
SmallVector<std::pair<OpOperand*, spatial::PhysicalLayout>> mismatches;
for (OpOperand& use : value.getUses()) {
Operation* userOp = use.getOwner();
spatial::PhysicalLayout required = spatial::PhysicalLayout::DenseNCHW;
if (auto capability = dyn_cast<spatial::SpatialLayoutCapabilityInterface>(userOp)) {
auto alternatives = getAlternatives(userOp, layouts, target);
if (failed(alternatives))
return failure();
spatial::PhysicalLayout selected =
getSelectedLayout(layouts, userOp->getResult(0));
unsigned selectedIndex = findCurrentAlternative(userOp, *alternatives, selected);
required = (*alternatives)[selectedIndex].operandLayouts[use.getOperandNumber()];
}
else if (auto flatten = dyn_cast<spatial::SpatGraphCompute>(userOp);
flatten && sourceLayout == spatial::PhysicalLayout::NHWCRowStrip
&& succeeded(canLowerFlattenFromRowStrip(flatten, target))) {
continue;
}
if (required != sourceLayout)
mismatches.push_back({&use, required});
}
for (auto [use, required] : mismatches) {
Operation* userOp = use->getOwner();
rewriter.setInsertionPoint(userOp);
auto materialized = spatial::SpatMaterializeLayoutOp::create(
rewriter, userOp->getLoc(), use->get().getType(), use->get(),
spatial::LogicalLayoutAttr::get(
rewriter.getContext(), spatial::LogicalLayout::NCHW),
spatial::PhysicalLayoutAttr::get(rewriter.getContext(), sourceLayout),
spatial::PhysicalLayoutAttr::get(rewriter.getContext(),
required));
use->set(materialized.getResult());
}
return success();
}
static LogicalResult verifySelectedLayouts(
ArrayRef<Operation*> planOps, const LayoutMap& layouts,
const spatial::SpatialTargetInfo& target) {
for (Operation* op : planOps) {
auto selected = spatial::getSelectedPhysicalLayout(op);
if (!selected)
return op->emitOpError("requires a selected physical layout"), failure();
auto alternatives = getAlternatives(op, layouts, target);
if (failed(alternatives))
return failure();
if (llvm::none_of(*alternatives, [&](const spatial::LayoutAlternative& alternative) {
return alternative.resultLayout == *selected;
}))
return op->emitOpError("selected physical layout is not advertised by its layout contract"), failure();
}
return success();
}
struct SpatialLayoutPlanningPass final
: PassWrapper<SpatialLayoutPlanningPass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(SpatialLayoutPlanningPass)
StringRef getArgument() const override { return "spatial-layout-planning"; }
StringRef getDescription() const override {
return "Select Spatial layout alternatives and insert explicit reconciliation barriers.";
}
SpatialLayoutPlanningPass() = default;
explicit SpatialLayoutPlanningPass(const spatial::SpatialTargetInfo& target)
: target(target), hasTarget(true) {}
void runOnOperation() override {
ModuleOp moduleOp = getOperation();
if (!hasTarget) {
moduleOp.emitError("Spatial layout planning requires an injected SpatialTargetInfo");
signalPassFailure();
return;
}
auto entryFunc = getPimEntryFunc(moduleOp);
if (failed(entryFunc)) {
moduleOp.emitError("failed to locate the PIM entry function during Spatial layout planning");
signalPassFailure();
return;
}
func::FuncOp funcOp = *entryFunc;
SmallVector<Operation*> planOps;
for (Operation& op : funcOp.getBody().front())
if (isa<spatial::SpatialLayoutCapabilityInterface>(&op))
planOps.push_back(&op);
LayoutMap layouts;
for (Operation* op : planOps)
layouts[op->getResult(0)] = spatial::PhysicalLayout::DenseNCHW;
const size_t maxRounds = 2 * planOps.size() + 1;
bool converged = false;
for (size_t round = 0; round < maxRounds && !converged; ++round) {
converged = true;
SmallVector<Operation*> order(planOps);
if (round % 2)
std::reverse(order.begin(), order.end());
for (Operation* op : order) {
auto alternatives = getAlternatives(op, layouts, target);
if (failed(alternatives)) {
signalPassFailure();
return;
}
spatial::PhysicalLayout current = layouts.lookup(op->getResult(0));
unsigned currentIndex = findCurrentAlternative(op, *alternatives, current);
int64_t bestCost = alternativeCost(
op, (*alternatives)[currentIndex], layouts, layouts, target);
unsigned bestIndex = currentIndex;
for (auto [index, alternative] : llvm::enumerate(*alternatives)) {
int64_t cost = alternativeCost(op, alternative, layouts, layouts, target);
if (cost < bestCost) {
bestCost = cost;
bestIndex = index;
}
}
spatial::PhysicalLayout selected = (*alternatives)[bestIndex].resultLayout;
if (selected != current) {
layouts[op->getResult(0)] = selected;
converged = false;
}
}
}
if (!converged) {
moduleOp.emitError("Spatial layout selection did not converge within its bounded iteration budget");
signalPassFailure();
return;
}
IRRewriter rewriter(&getContext());
for (Operation* op : planOps) {
op->setAttr(spatial::kSelectedLayoutAttrName,
spatial::PhysicalLayoutAttr::get(
rewriter.getContext(), layouts.lookup(op->getResult(0))));
if (failed(materializeMismatchedUses(rewriter, op->getResult(0), layouts, target))) {
signalPassFailure();
return;
}
}
if (failed(verifySelectedLayouts(planOps, layouts, target))
|| failed(verifyLogicalSpatialGraphInvariants(*entryFunc))) {
moduleOp.emitError("Spatial layout planning verification failed");
signalPassFailure();
}
}
spatial::SpatialTargetInfo target;
bool hasTarget = false;
};
} // namespace
std::unique_ptr<Pass> createSpatialLayoutPlanningPass() {
return std::make_unique<SpatialLayoutPlanningPass>();
}
std::unique_ptr<Pass> createSpatialLayoutPlanningPass(
const spatial::SpatialTargetInfo& target) {
return std::make_unique<SpatialLayoutPlanningPass>(target);
}
} // namespace onnx_mlir
@@ -307,7 +307,7 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul
"resultful compute_batch lowering currently requires a spat.in_parallel terminator"); "resultful compute_batch lowering currently requires a spat.in_parallel terminator");
} }
auto coreIds = getRequiredScheduledBatchCoreIds(computeBatchOp, "spatial compute_batch core id"); auto coreIds = getRequiredScheduledBatchCoreIds(computeBatchOp, "Spatial compute_batch core id");
if (failed(coreIds)) if (failed(coreIds))
return failure(); return failure();
SmallVector<Value> batchWeights(computeBatchOp.getWeights().begin(), computeBatchOp.getWeights().end()); SmallVector<Value> batchWeights(computeBatchOp.getWeights().begin(), computeBatchOp.getWeights().end());
@@ -317,7 +317,7 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeBatchOp(spatial::SpatSchedul
rewriter.setInsertionPointAfter(computeBatchOp); rewriter.setInsertionPointAfter(computeBatchOp);
auto laneCountAttr = pim::getCheckedI32Attr( auto laneCountAttr = pim::getCheckedI32Attr(
rewriter, computeBatchOp, static_cast<uint64_t>(computeBatchOp.getLaneCount()), "pim core_batch lane count"); rewriter, computeBatchOp, static_cast<uint64_t>(computeBatchOp.getLaneCount()), "Pim core_batch lane count");
if (failed(laneCountAttr)) if (failed(laneCountAttr))
return failure(); return failure();
auto coreBatchOp = auto coreBatchOp =
@@ -368,11 +368,14 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeOp(spatial::SpatScheduledCom
return failure(); return failure();
PimWaitOp::create( PimWaitOp::create(
rewriter, receiveOp->getLoc(), hostWaitLoad.getEventRegister(), rewriter, receiveOp->getLoc(), hostWaitLoad.getEventRegister(),
rewriter.getI32IntegerAttr(1)); hostWaitLoad.getWaitValue());
received = PimMemCopyHostToDevOp::create( received = PimMemCopyHostToDevOp::create(
rewriter, receiveOp->getLoc(), outputBuffer.getType(), zero, rewriter, receiveOp->getLoc(), outputBuffer.getType(), zero,
hostWaitLoad.getHostOffset(), outputBuffer, *hostBuffer, *sizeAttr) hostWaitLoad.getHostOffset(), outputBuffer, *hostBuffer, *sizeAttr)
.getOutput(); .getOutput();
PimSyncOp::create(
rewriter, receiveOp->getLoc(), hostWaitLoad.getSourceCoreId(),
hostWaitLoad.getAcknowledgementEventRegister());
} else { } else {
received = PimReceiveOp::create( received = PimReceiveOp::create(
rewriter, receiveOp->getLoc(), outputBuffer.getType(), outputBuffer, rewriter, receiveOp->getLoc(), outputBuffer.getType(), outputBuffer,
@@ -407,7 +410,7 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeOp(spatial::SpatScheduledCom
continue; continue;
} }
return computeOp.emitOpError("has an unsupported remaining result use during Spatial-to-PIM lowering"); return computeOp.emitOpError("has an unsupported remaining result use during Spatial-to-Pim lowering");
} }
rewriter.setInsertionPoint(yieldOp); rewriter.setInsertionPoint(yieldOp);
@@ -417,7 +420,7 @@ LogicalResult raptor::SpatialToPimPass::lowerComputeOp(spatial::SpatScheduledCom
if (!computeOp.getWeights().empty()) if (!computeOp.getWeights().empty())
computeWeights.append(computeOp.getWeights().begin(), computeOp.getWeights().end()); computeWeights.append(computeOp.getWeights().begin(), computeOp.getWeights().end());
rewriter.setInsertionPointAfter(computeOp); rewriter.setInsertionPointAfter(computeOp);
auto checkedCoreId = getRequiredScheduledCoreId(computeOp, "spatial compute core id"); auto checkedCoreId = getRequiredScheduledCoreId(computeOp, "Spatial compute core id");
if (failed(checkedCoreId)) if (failed(checkedCoreId))
return failure(); return failure();
auto coreIdAttr = pim::getCheckedI32Attr(rewriter, computeOp, static_cast<int64_t>(*checkedCoreId), "pim core id"); auto coreIdAttr = pim::getCheckedI32Attr(rewriter, computeOp, static_cast<int64_t>(*checkedCoreId), "pim core id");
@@ -147,15 +147,42 @@ struct HostWaitLoadLowering : OpRewritePattern<spatial::SpatHostWaitLoadOp> {
return failure(); return failure();
auto wait = pim::PimWaitOp::create( auto wait = pim::PimWaitOp::create(
rewriter, op.getLoc(), op.getEventRegister(), rewriter, op.getLoc(), op.getEventRegister(),
rewriter.getI32IntegerAttr(1)); op.getWaitValue());
copyRaptorDebugAttrs(op.getOperation(), wait.getOperation()); copyRaptorDebugAttrs(op.getOperation(), wait.getOperation());
return pim::PimMemCopyHostToDevOp::create( Value output = pim::PimMemCopyHostToDevOp::create(
rewriter, op.getLoc(), outputBuffer.getType(), zero, rewriter, op.getLoc(), outputBuffer.getType(), zero,
op.getHostOffset(), outputBuffer, *hostBuffer, sizeAttr).getOutput(); op.getHostOffset(), outputBuffer, *hostBuffer, sizeAttr).getOutput();
auto sync = pim::PimSyncOp::create(
rewriter, op.getLoc(), op.getSourceCoreId(),
op.getAcknowledgementEventRegister());
copyRaptorDebugAttrs(op.getOperation(), sync.getOperation());
return output;
}); });
} }
}; };
struct SyncLowering : OpRewritePattern<spatial::SpatSyncOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatSyncOp op,
PatternRewriter& rewriter) const override {
rewriter.replaceOpWithNewOp<pim::PimSyncOp>(
op, op.getTargetCoreId(), op.getEventRegister());
return success();
}
};
struct WaitLowering : OpRewritePattern<spatial::SpatWaitOp> {
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(spatial::SpatWaitOp op,
PatternRewriter& rewriter) const override {
rewriter.replaceOpWithNewOp<pim::PimWaitOp>(
op, op.getEventRegister(), op.getWaitValue());
return success();
}
};
struct ExtractRowsLowering : OpRewritePattern<spatial::SpatExtractRowsOp> { struct ExtractRowsLowering : OpRewritePattern<spatial::SpatExtractRowsOp> {
using OpRewritePattern::OpRewritePattern; using OpRewritePattern::OpRewritePattern;
@@ -200,7 +227,8 @@ struct ConcatLowering : OpRewritePattern<spatial::SpatConcatOp> {
void populateChannelLoweringPatterns(RewritePatternSet& patterns) { void populateChannelLoweringPatterns(RewritePatternSet& patterns) {
patterns.add<ChannelSendLowering, ChannelReceiveLowering, patterns.add<ChannelSendLowering, ChannelReceiveLowering,
HostStoreSyncLowering, HostWaitLoadLowering, HostStoreSyncLowering, HostWaitLoadLowering,
ExtractRowsLowering, ConcatLowering>(patterns.getContext()); SyncLowering, WaitLowering, ExtractRowsLowering,
ConcatLowering>(patterns.getContext());
} }
} // namespace onnx_mlir } // namespace onnx_mlir
@@ -734,7 +734,7 @@ raptor::SpatialToPimPass::ReturnPathLoweringResult raptor::SpatialToPimPass::low
auto storedType = dyn_cast<RankedTensorType>(storedValue.getType()); auto storedType = dyn_cast<RankedTensorType>(storedValue.getType());
if (!storedType) { if (!storedType) {
producerOp->emitOpError( producerOp->emitOpError(
"has an unsupported non-ranked concat-return helper yield during Spatial-to-PIM lowering"); "has an unsupported non-ranked concat-return helper yield during Spatial-to-Pim lowering");
return ReturnPathLoweringResult::Failure; return ReturnPathLoweringResult::Failure;
} }
rewriter.setInsertionPointAfterValue(storedValue); rewriter.setInsertionPointAfterValue(storedValue);
@@ -748,7 +748,7 @@ raptor::SpatialToPimPass::ReturnPathLoweringResult raptor::SpatialToPimPass::low
SmallVector<int64_t> destinationIndices; SmallVector<int64_t> destinationIndices;
if (failed(mapIndicesThroughHelperChain( if (failed(mapIndicesThroughHelperChain(
sourceIndices, concatReturnUse->concatShape, concatReturnUse->helperChain, destinationIndices))) { sourceIndices, concatReturnUse->concatShape, concatReturnUse->helperChain, destinationIndices))) {
producerOp->emitOpError("has an unsupported concat-return helper chain during Spatial-to-PIM lowering"); producerOp->emitOpError("has an unsupported concat-return helper chain during Spatial-to-Pim lowering");
return ReturnPathLoweringResult::Failure; return ReturnPathLoweringResult::Failure;
} }
@@ -88,7 +88,7 @@ void onnx_mlir::raptor::SpatialToPimPass::runOnOperation() {
operationsToRemove.clear(); operationsToRemove.clear();
ModuleOp moduleOp = getOperation(); ModuleOp moduleOp = getOperation();
if (!hasTarget || failed(targetResources.verify())) { if (!hasTarget || failed(targetResources.verify())) {
moduleOp.emitError("Spatial-to-PIM lowering requires valid injected target resources"); moduleOp.emitError("Spatial-to-Pim lowering requires valid injected target resources");
signalPassFailure(); signalPassFailure();
return; return;
} }
@@ -96,7 +96,7 @@ void onnx_mlir::raptor::SpatialToPimPass::runOnOperation() {
auto entryFunc = getPimEntryFunc(moduleOp); auto entryFunc = getPimEntryFunc(moduleOp);
if (failed(entryFunc)) { if (failed(entryFunc)) {
moduleOp.emitError("failed to locate the PIM entry function during Spatial-to-PIM lowering"); moduleOp.emitError("failed to locate the Pim entry function during Spatial-to-Pim lowering");
signalPassFailure(); signalPassFailure();
return; return;
} }
@@ -128,12 +128,14 @@ void onnx_mlir::raptor::SpatialToPimPass::runOnOperation() {
spatial::SpatChannelSendOp, spatial::SpatChannelSendOp,
spatial::SpatHostStoreSyncOp, spatial::SpatHostStoreSyncOp,
spatial::SpatHostWaitLoadOp, spatial::SpatHostWaitLoadOp,
spatial::SpatSyncOp,
spatial::SpatWaitOp,
spatial::SpatExtractRowsOp>(); spatial::SpatExtractRowsOp>();
RewritePatternSet initialPatterns(ctx); RewritePatternSet initialPatterns(ctx);
populateInitialPatterns(initialPatterns); populateInitialPatterns(initialPatterns);
if (failed(applyPartialConversion(moduleOp, target, std::move(initialPatterns)))) { if (failed(applyPartialConversion(moduleOp, target, std::move(initialPatterns)))) {
moduleOp.emitError("failed to lower required Spatial ops to the initial PIM form"); moduleOp.emitError("failed to lower required Spatial ops to the initial Pim form");
signalPassFailure(); signalPassFailure();
return; return;
} }
@@ -151,7 +153,7 @@ void onnx_mlir::raptor::SpatialToPimPass::runOnOperation() {
auto returnOp = cast<func::ReturnOp>(funcOp.front().getTerminator()); auto returnOp = cast<func::ReturnOp>(funcOp.front().getTerminator());
addReturnOutputBuffers(returnOp, rewriter); addReturnOutputBuffers(returnOp, rewriter);
if (failed(allocateAndInitializeCoreLocalVariables(funcOp, rewriter))) { if (failed(allocateAndInitializeCoreLocalVariables(funcOp, rewriter))) {
funcOp.emitOpError("failed to allocate or initialize core-local tensors during Spatial-to-PIM lowering"); funcOp.emitOpError("failed to allocate or initialize core-local tensors during Spatial-to-Pim lowering");
signalPassFailure(); signalPassFailure();
return; return;
} }
@@ -223,6 +225,8 @@ void onnx_mlir::raptor::SpatialToPimPass::runOnOperation() {
spatial::SpatChannelSendOp, spatial::SpatChannelSendOp,
spatial::SpatHostStoreSyncOp, spatial::SpatHostStoreSyncOp,
spatial::SpatHostWaitLoadOp, spatial::SpatHostWaitLoadOp,
spatial::SpatSyncOp,
spatial::SpatWaitOp,
spatial::SpatExtractRowsOp>(); spatial::SpatExtractRowsOp>();
SmallVector<pim::PimCoreOp> coreOps; SmallVector<pim::PimCoreOp> coreOps;
@@ -274,12 +278,14 @@ void onnx_mlir::raptor::SpatialToPimPass::runOnOperation() {
spatial::SpatChannelSendOp, spatial::SpatChannelSendOp,
spatial::SpatHostStoreSyncOp, spatial::SpatHostStoreSyncOp,
spatial::SpatHostWaitLoadOp, spatial::SpatHostWaitLoadOp,
spatial::SpatSyncOp,
spatial::SpatWaitOp,
spatial::SpatExtractRowsOp>(); spatial::SpatExtractRowsOp>();
RewritePatternSet communicationPatterns(ctx); RewritePatternSet communicationPatterns(ctx);
populateChannelLoweringPatterns(communicationPatterns); populateChannelLoweringPatterns(communicationPatterns);
if (failed(applyFullConversion(funcOp, communicationTarget, std::move(communicationPatterns)))) { if (failed(applyFullConversion(funcOp, communicationTarget, std::move(communicationPatterns)))) {
funcOp.emitOpError("failed to lower Spatial communication ops to PIM communication ops"); funcOp.emitOpError("failed to lower Spatial communication ops to Pim communication ops");
signalPassFailure(); signalPassFailure();
return; return;
} }
@@ -26,7 +26,7 @@ namespace raptor {
struct SpatialToPimPass : mlir::PassWrapper<SpatialToPimPass, mlir::OperationPass<mlir::ModuleOp>> { struct SpatialToPimPass : mlir::PassWrapper<SpatialToPimPass, mlir::OperationPass<mlir::ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(SpatialToPimPass) MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(SpatialToPimPass)
llvm::StringRef getArgument() const override { return "convert-spatial-to-pim"; } llvm::StringRef getArgument() const override { return "convert-spatial-to-pim"; }
llvm::StringRef getDescription() const override { return "Lower Spatial ops to PIM-ready format"; } llvm::StringRef getDescription() const override { return "Lower Spatial ops to Pim-ready format"; }
SpatialToPimPass() = default; SpatialToPimPass() = default;
explicit SpatialToPimPass(const spatial::SpatialTargetResources& target) explicit SpatialToPimPass(const spatial::SpatialTargetResources& target)
@@ -302,12 +302,19 @@ static FailureOr<int64_t> getShapedByteSize(MemRefType type) {
return static_cast<int64_t>(*byteSize); return static_cast<int64_t>(*byteSize);
} }
static FailureOr<SmallVector<int64_t>> struct LogicalCopyShape {
SmallVector<int64_t> dimensions;
Type elementType;
};
static bool isPackedByteBuffer(MemRefType type) {
return type.getRank() == 1 && type.getElementType().isInteger(8);
}
static FailureOr<LogicalCopyShape>
inferLogicalCopyShape(MemRefType targetType, MemRefType sourceType, int64_t size) { inferLogicalCopyShape(MemRefType targetType, MemRefType sourceType, int64_t size) {
if (!targetType.hasStaticShape() || !sourceType.hasStaticShape()) if (!targetType.hasStaticShape() || !sourceType.hasStaticShape())
return failure(); return failure();
if (targetType.getElementType() != sourceType.getElementType() || targetType.getRank() != sourceType.getRank())
return failure();
auto targetBytes = getShapedByteSize(targetType); auto targetBytes = getShapedByteSize(targetType);
auto sourceBytes = getShapedByteSize(sourceType); auto sourceBytes = getShapedByteSize(sourceType);
@@ -316,18 +323,37 @@ inferLogicalCopyShape(MemRefType targetType, MemRefType sourceType, int64_t size
bool targetMatches = *targetBytes == size; bool targetMatches = *targetBytes == size;
bool sourceMatches = *sourceBytes == size; bool sourceMatches = *sourceBytes == size;
if (targetMatches && sourceMatches && targetType.getShape() != sourceType.getShape()) bool matchingTypes = targetType.getElementType() == sourceType.getElementType()
&& targetType.getRank() == sourceType.getRank();
if (matchingTypes) {
if (targetMatches && sourceMatches
&& targetType.getShape() != sourceType.getShape())
return failure();
MemRefType logicalType = targetMatches ? targetType : sourceType;
if (targetMatches || sourceMatches)
return LogicalCopyShape {
SmallVector<int64_t>(logicalType.getShape()),
logicalType.getElementType()};
return failure(); return failure();
if (targetMatches) }
return SmallVector<int64_t>(targetType.getShape().begin(), targetType.getShape().end()); if (targetMatches && isPackedByteBuffer(sourceType))
if (sourceMatches) return LogicalCopyShape {
return SmallVector<int64_t>(sourceType.getShape().begin(), sourceType.getShape().end()); SmallVector<int64_t>(targetType.getShape()),
targetType.getElementType()};
if (sourceMatches && isPackedByteBuffer(targetType))
return LogicalCopyShape {
SmallVector<int64_t>(sourceType.getShape()),
sourceType.getElementType()};
return failure(); return failure();
} }
static FailureOr<int64_t> getContiguousSuffixRank(Value value, ArrayRef<int64_t> copyShape) { static FailureOr<int64_t> getContiguousSuffixRank(
Value value, ArrayRef<int64_t> copyShape, Type elementType = {}) {
auto type = dyn_cast<MemRefType>(value.getType()); auto type = dyn_cast<MemRefType>(value.getType());
if (type && elementType && isPackedByteBuffer(type))
return copyShape.size();
if (!type || !type.hasStaticShape() || !hasByteSizedElementType(type.getElementType()) if (!type || !type.hasStaticShape() || !hasByteSizedElementType(type.getElementType())
|| (elementType && type.getElementType() != elementType)
|| type.getRank() != static_cast<int64_t>(copyShape.size())) || type.getRank() != static_cast<int64_t>(copyShape.size()))
return failure(); return failure();
if (llvm::any_of(copyShape, [](int64_t dim) { return dim <= 0; })) if (llvm::any_of(copyShape, [](int64_t dim) { return dim <= 0; }))
@@ -351,6 +377,30 @@ static FailureOr<int64_t> getContiguousSuffixRank(Value value, ArrayRef<int64_t>
return contiguousSuffixRank; return contiguousSuffixRank;
} }
static FailureOr<SmallVector<int64_t>> getOuterByteStrides(
Value value, const LogicalCopyShape &copyShape, size_t outerRank) {
auto type = cast<MemRefType>(value.getType());
SmallVector<int64_t> strides;
if (isPackedByteBuffer(type))
strides = computeRowMajorStrides(copyShape.dimensions);
else {
auto proven = getProvenMemRefStrides(value);
if (failed(proven))
return failure();
strides = std::move(*proven);
}
int64_t elementByteWidth = static_cast<int64_t>(
getElementTypeSizeInBytes(copyShape.elementType));
SmallVector<int64_t> result;
for (int64_t stride : ArrayRef<int64_t>(strides).take_front(outerRank)) {
auto byteStride = checkedPositiveMul(stride, elementByteWidth);
if (failed(byteStride))
return failure();
result.push_back(*byteStride);
}
return result;
}
static FailureOr<CopyEndpointPlan> analyzeCopyEndpoint(Value value, Value initialByteOffset, MemRefType logicalType) { static FailureOr<CopyEndpointPlan> analyzeCopyEndpoint(Value value, Value initialByteOffset, MemRefType logicalType) {
if (!logicalType.hasStaticShape() || !hasByteSizedElementType(logicalType.getElementType())) if (!logicalType.hasStaticShape() || !hasByteSizedElementType(logicalType.getElementType()))
return failure(); return failure();
@@ -448,8 +498,10 @@ analyzeCopyRewrite(Value target, Value source, Value targetOffset, Value sourceO
if (failed(logicalCopyShape)) if (failed(logicalCopyShape))
return failure(); return failure();
auto targetSuffixRank = getContiguousSuffixRank(target, *logicalCopyShape); auto targetSuffixRank = getContiguousSuffixRank(
auto sourceSuffixRank = getContiguousSuffixRank(source, *logicalCopyShape); target, logicalCopyShape->dimensions, logicalCopyShape->elementType);
auto sourceSuffixRank = getContiguousSuffixRank(
source, logicalCopyShape->dimensions, logicalCopyShape->elementType);
if (failed(targetSuffixRank) || failed(sourceSuffixRank)) if (failed(targetSuffixRank) || failed(sourceSuffixRank))
return failure(); return failure();
@@ -458,23 +510,24 @@ analyzeCopyRewrite(Value target, Value source, Value targetOffset, Value sourceO
plan.source = *sourcePlan; plan.source = *sourcePlan;
int64_t contiguousSuffixRank = std::min(*targetSuffixRank, *sourceSuffixRank); int64_t contiguousSuffixRank = std::min(*targetSuffixRank, *sourceSuffixRank);
if (contiguousSuffixRank == static_cast<int64_t>(logicalCopyShape->size())) { if (contiguousSuffixRank
== static_cast<int64_t>(logicalCopyShape->dimensions.size())) {
plan.kind = CopyRewritePlan::Kind::Direct; plan.kind = CopyRewritePlan::Kind::Direct;
plan.directBytes = size; plan.directBytes = size;
return plan; return plan;
} }
auto targetStrides = getProvenMemRefStrides(target); int64_t elementByteWidth = static_cast<int64_t>(
auto sourceStrides = getProvenMemRefStrides(source); getElementTypeSizeInBytes(logicalCopyShape->elementType));
if (failed(targetStrides) || failed(sourceStrides))
return failure();
int64_t elementByteWidth = static_cast<int64_t>(getElementTypeSizeInBytes(targetType.getElementType()));
plan.kind = CopyRewritePlan::Kind::Loop; plan.kind = CopyRewritePlan::Kind::Loop;
plan.loop.targetBaseOffset = plan.target.offset; plan.loop.targetBaseOffset = plan.target.offset;
plan.loop.sourceBaseOffset = plan.source.offset; plan.loop.sourceBaseOffset = plan.source.offset;
plan.loop.outerShape.assign(logicalCopyShape->begin(), logicalCopyShape->end() - contiguousSuffixRank); plan.loop.outerShape.assign(
SmallVector<int64_t> chunkShape(logicalCopyShape->end() - contiguousSuffixRank, logicalCopyShape->end()); logicalCopyShape->dimensions.begin(),
logicalCopyShape->dimensions.end() - contiguousSuffixRank);
SmallVector<int64_t> chunkShape(
logicalCopyShape->dimensions.end() - contiguousSuffixRank,
logicalCopyShape->dimensions.end());
auto outerElements = checkedPositiveProduct(plan.loop.outerShape); auto outerElements = checkedPositiveProduct(plan.loop.outerShape);
auto chunkElements = checkedPositiveProduct(chunkShape); auto chunkElements = checkedPositiveProduct(chunkShape);
auto chunkBytes = failed(chunkElements) auto chunkBytes = failed(chunkElements)
@@ -484,18 +537,14 @@ analyzeCopyRewrite(Value target, Value source, Value targetOffset, Value sourceO
return failure(); return failure();
plan.loop.outerElements = *outerElements; plan.loop.outerElements = *outerElements;
plan.loop.chunkBytes = *chunkBytes; plan.loop.chunkBytes = *chunkBytes;
for (int64_t stride : ArrayRef<int64_t>(*targetStrides).take_front(plan.loop.outerShape.size())) { auto targetStrides = getOuterByteStrides(
auto byteStride = checkedPositiveMul(stride, elementByteWidth); target, *logicalCopyShape, plan.loop.outerShape.size());
if (failed(byteStride)) auto sourceStrides = getOuterByteStrides(
return failure(); source, *logicalCopyShape, plan.loop.outerShape.size());
plan.loop.targetOuterByteStrides.push_back(*byteStride); if (failed(targetStrides) || failed(sourceStrides))
} return failure();
for (int64_t stride : ArrayRef<int64_t>(*sourceStrides).take_front(plan.loop.outerShape.size())) { plan.loop.targetOuterByteStrides = std::move(*targetStrides);
auto byteStride = checkedPositiveMul(stride, elementByteWidth); plan.loop.sourceOuterByteStrides = std::move(*sourceStrides);
if (failed(byteStride))
return failure();
plan.loop.sourceOuterByteStrides.push_back(*byteStride);
}
if (plan.loop.chunkBytes <= 0) if (plan.loop.chunkBytes <= 0)
return failure(); return failure();
return plan; return plan;
@@ -402,7 +402,7 @@ static LogicalResult verifyPimCoresNeedNoTensorCopies(
bufferization::BufferizationState state; bufferization::BufferizationState state;
if (failed(bufferization::insertTensorCopies(*clone, options, state))) { if (failed(bufferization::insertTensorCopies(*clone, options, state))) {
moduleOp.emitError("official one-shot analysis failed while verifying PIM core copy freedom"); moduleOp.emitError("official one-shot analysis failed while verifying Pim core copy freedom");
return failure(); return failure();
} }
@@ -415,10 +415,10 @@ static LogicalResult verifyPimCoresNeedNoTensorCopies(
Operation* requiredBy = alloc->getUsers().empty() Operation* requiredBy = alloc->getUsers().empty()
? alloc.getOperation() : *alloc->getUsers().begin(); ? alloc.getOperation() : *alloc->getUsers().begin();
diagnostics.report(requiredBy, [](Operation* op) { diagnostics.report(requiredBy, [](Operation* op) {
op->emitOpError("official one-shot bufferization requires a tensor copy inside a PIM core"); op->emitOpError("official one-shot bufferization requires a tensor copy inside a Pim core");
}); });
}); });
diagnostics.emitSuppressedSummary(moduleOp, "required PIM core tensor copies"); diagnostics.emitSuppressedSummary(moduleOp, "required Pim core tensor copies");
return success(!diagnostics.hasFailure()); return success(!diagnostics.hasFailure());
} }
@@ -440,7 +440,7 @@ static LogicalResult runOneShotPimBufferization(
bufferization::BufferizationState state; bufferization::BufferizationState state;
if (failed(bufferization::insertTensorCopies(moduleOp, hostOptions, state)) if (failed(bufferization::insertTensorCopies(moduleOp, hostOptions, state))
|| failed(bufferization::bufferizeModuleOp(moduleOp, options, state))) { || failed(bufferization::bufferizeModuleOp(moduleOp, options, state))) {
moduleOp.emitError("Failed to bufferize PIM and Spatial ops"); moduleOp.emitError("Failed to bufferize Pim and Spatial ops");
return failure(); return failure();
} }
return success(); return success();
@@ -478,7 +478,7 @@ static LogicalResult verifyContiguousRuntimeOperands(ModuleOp moduleOp) {
if (succeeded(resolveContiguousAddress(operand, knowledge)) || succeeded(compileContiguousAddressExpr(operand))) if (succeeded(resolveContiguousAddress(operand, knowledge)) || succeeded(compileContiguousAddressExpr(operand)))
return; return;
op.emitOpError() << "operand #" << operandIndex op.emitOpError() << "operand #" << operandIndex
<< " is not backed by contiguous addressable storage after PIM bufferization"; << " is not backed by contiguous addressable storage after Pim bufferization";
hasFailure = true; hasFailure = true;
}; };
@@ -552,7 +552,7 @@ static LogicalResult verifyContiguousRuntimeOperands(ModuleOp moduleOp) {
}); });
if (hasFailure) { if (hasFailure) {
moduleOp.emitError("PIM bufferization must fully normalize executable runtime operand contiguity before codegen"); moduleOp.emitError("Pim bufferization must fully normalize executable runtime operand contiguity before codegen");
return failure(); return failure();
} }
return success(); return success();
@@ -589,7 +589,7 @@ static LogicalResult verifyPimCopyAddressSpaces(ModuleOp moduleOp) {
}); });
if (failureCount != 0) if (failureCount != 0)
moduleOp.emitError() << "found " << failureCount moduleOp.emitError() << "found " << failureCount
<< " PIM copy address-space violation(s); the first is reported above"; << " Pim copy address-space violation(s); the first is reported above";
return success(failureCount == 0); return success(failureCount == 0);
} }
@@ -602,18 +602,30 @@ static LogicalResult normalizePimMemory(ModuleOp moduleOp, func::FuncOp funcOp)
PatternRewriter rewriter(ctx); PatternRewriter rewriter(ctx);
SmallVector<MemRefCopyWorkItem> copyWorklist; SmallVector<MemRefCopyWorkItem> copyWorklist;
SmallVector<PimMemCopyDevToHostOp> hostToHostCopies;
llvm::SmallPtrSet<Operation*, 16> seenCopyOps; llvm::SmallPtrSet<Operation*, 16> seenCopyOps;
llvm::SmallPtrSet<Operation*, 4> seenHostToHostCopies;
auto addCopyOp = [&](memref::CopyOp copyOp, const StaticValueKnowledge& knowledge) { auto addCopyOp = [&](memref::CopyOp copyOp, const StaticValueKnowledge& knowledge) {
if (seenCopyOps.insert(copyOp.getOperation()).second) if (seenCopyOps.insert(copyOp.getOperation()).second)
copyWorklist.push_back({copyOp, knowledge}); copyWorklist.push_back({copyOp, knowledge});
}; };
auto collectCopy = [&](Operation &op,
const StaticValueKnowledge &knowledge) {
if (auto copyOp = dyn_cast<memref::CopyOp>(&op))
addCopyOp(copyOp, knowledge);
if (auto copyOp = dyn_cast<PimMemCopyDevToHostOp>(&op);
copyOp
&& isHostBackedPimAddress(copyOp.getDeviceSource(), knowledge)
&& isHostBackedPimAddress(copyOp.getHostTarget(), knowledge)
&& seenHostToHostCopies.insert(copyOp).second)
hostToHostCopies.push_back(copyOp);
};
moduleOp.walk([&](pim::PimCoreOp coreOp) { moduleOp.walk([&](pim::PimCoreOp coreOp) {
StaticValueKnowledge knowledge = seedCoreKnowledge(coreOp); StaticValueKnowledge knowledge = seedCoreKnowledge(coreOp);
(void) walkPimCoreBlockStructurally( (void) walkPimCoreBlockStructurally(
coreOp.getBody().front(), knowledge, [&](Operation& op, const StaticValueKnowledge& opKnowledge) { coreOp.getBody().front(), knowledge, [&](Operation& op, const StaticValueKnowledge& opKnowledge) {
if (auto copyOp = dyn_cast<memref::CopyOp>(&op)) collectCopy(op, opKnowledge);
addCopyOp(copyOp, opKnowledge);
return success(); return success();
}); });
}); });
@@ -622,8 +634,7 @@ static LogicalResult normalizePimMemory(ModuleOp moduleOp, func::FuncOp funcOp)
StaticValueKnowledge knowledge = seedCoreBatchKnowledge(coreBatchOp, lane); StaticValueKnowledge knowledge = seedCoreBatchKnowledge(coreBatchOp, lane);
(void) walkPimCoreBlockStructurally( (void) walkPimCoreBlockStructurally(
coreBatchOp.getBody().front(), knowledge, [&](Operation& op, const StaticValueKnowledge& opKnowledge) { coreBatchOp.getBody().front(), knowledge, [&](Operation& op, const StaticValueKnowledge& opKnowledge) {
if (auto copyOp = dyn_cast<memref::CopyOp>(&op)) collectCopy(op, opKnowledge);
addCopyOp(copyOp, opKnowledge);
return success(); return success();
}); });
} }
@@ -631,6 +642,22 @@ static LogicalResult normalizePimMemory(ModuleOp moduleOp, func::FuncOp funcOp)
bool hasFailed = false; bool hasFailed = false;
Value zeroOffset = getOrCreateIndexConstant(rewriter, funcOp, 0); Value zeroOffset = getOrCreateIndexConstant(rewriter, funcOp, 0);
for (PimMemCopyDevToHostOp copyOp : hostToHostCopies) {
rewriter.setInsertionPoint(copyOp);
auto scratchType = MemRefType::get(
{copyOp.getSize()}, rewriter.getI8Type());
Value scratch = memref::AllocOp::create(
rewriter, copyOp.getLoc(), scratchType);
auto load = PimMemCopyHostToDevOp::create(
rewriter, copyOp.getLoc(), scratchType, zeroOffset,
copyOp.getDeviceSourceOffset(), scratch, copyOp.getDeviceSource(),
copyOp.getSizeAttr());
auto store = PimMemCopyDevToHostOp::create(
rewriter, copyOp.getLoc(), copyOp.getHostTarget().getType(),
copyOp.getHostTargetOffset(), zeroOffset, copyOp.getHostTarget(),
load.getOutput(), copyOp.getSizeAttr());
rewriter.replaceOp(copyOp, store.getOutput());
}
for (const MemRefCopyWorkItem& workItem : copyWorklist) { for (const MemRefCopyWorkItem& workItem : copyWorklist) {
memref::CopyOp copyOp = workItem.copyOp; memref::CopyOp copyOp = workItem.copyOp;
rewriter.setInsertionPoint(copyOp); rewriter.setInsertionPoint(copyOp);
@@ -646,7 +673,7 @@ static LogicalResult normalizePimMemory(ModuleOp moduleOp, func::FuncOp funcOp)
GreedyRewriteConfig contiguityConfig; GreedyRewriteConfig contiguityConfig;
contiguityConfig.enableFolding(false); contiguityConfig.enableFolding(false);
if (failed(applyPatternsGreedily(moduleOp, std::move(contiguityPatterns), contiguityConfig))) { if (failed(applyPatternsGreedily(moduleOp, std::move(contiguityPatterns), contiguityConfig))) {
moduleOp.emitError("failed to normalize PIM copy contiguity during bufferization"); moduleOp.emitError("failed to normalize Pim copy contiguity during bufferization");
return failure(); return failure();
} }
annotateWeightsMemrefs(moduleOp, funcOp); annotateWeightsMemrefs(moduleOp, funcOp);
@@ -657,7 +684,7 @@ static LogicalResult normalizePimMemory(ModuleOp moduleOp, func::FuncOp funcOp)
static FailureOr<func::FuncOp> requirePimEntryFunc(ModuleOp moduleOp, StringRef phase) { static FailureOr<func::FuncOp> requirePimEntryFunc(ModuleOp moduleOp, StringRef phase) {
auto entryFunc = getPimEntryFunc(moduleOp); auto entryFunc = getPimEntryFunc(moduleOp);
if (failed(entryFunc)) { if (failed(entryFunc)) {
moduleOp.emitError("failed to locate the PIM entry function during ") << phase; moduleOp.emitError("failed to locate the Pim entry function during ") << phase;
return failure(); return failure();
} }
return *entryFunc; return *entryFunc;
@@ -674,12 +701,12 @@ struct PimBufferizationPreparationPass
StringRef getArgument() const override { return "pim-bufferization-preparation"; } StringRef getArgument() const override { return "pim-bufferization-preparation"; }
StringRef getDescription() const override { StringRef getDescription() const override {
return "Prepare writable tensor destinations for PIM one-shot bufferization."; return "Prepare writable tensor destinations for Pim one-shot bufferization.";
} }
void runOnOperation() final { void runOnOperation() final {
ModuleOp moduleOp = getOperation(); ModuleOp moduleOp = getOperation();
auto funcOp = requirePimEntryFunc(moduleOp, "PIM bufferization preparation"); auto funcOp = requirePimEntryFunc(moduleOp, "Pim bufferization preparation");
if (failed(funcOp)) { if (failed(funcOp)) {
signalPassFailure(); signalPassFailure();
return; return;
@@ -698,7 +725,7 @@ struct PimOneShotBufferizationPass
StringRef getArgument() const override { return "pim-one-shot-bufferization"; } StringRef getArgument() const override { return "pim-one-shot-bufferization"; }
StringRef getDescription() const override { StringRef getDescription() const override {
return "Run one-shot bufferization for PIM and Spatial tensors."; return "Run one-shot bufferization for Pim and Spatial tensors.";
} }
void runOnOperation() final { void runOnOperation() final {
@@ -713,12 +740,12 @@ struct PimMemoryNormalizationPass
StringRef getArgument() const override { return "pim-memory-normalization"; } StringRef getArgument() const override { return "pim-memory-normalization"; }
StringRef getDescription() const override { StringRef getDescription() const override {
return "Normalize PIM memory copies and verify addressable operands."; return "Normalize Pim memory copies and verify addressable operands.";
} }
void runOnOperation() final { void runOnOperation() final {
ModuleOp moduleOp = getOperation(); ModuleOp moduleOp = getOperation();
auto funcOp = requirePimEntryFunc(moduleOp, "PIM memory normalization"); auto funcOp = requirePimEntryFunc(moduleOp, "Pim memory normalization");
if (failed(funcOp)) { if (failed(funcOp)) {
signalPassFailure(); signalPassFailure();
return; return;
@@ -734,20 +761,20 @@ static LogicalResult verifyNoTensorValues(ModuleOp moduleOp) {
if (failureCount >= 8) if (failureCount >= 8)
return; return;
if (op->getDialect()->getNamespace() == "tensor") { if (op->getDialect()->getNamespace() == "tensor") {
op->emitOpError("tensor operation remains after PIM bufferization"); op->emitOpError("tensor operation remains after Pim bufferization");
++failureCount; ++failureCount;
return; return;
} }
for (Value value : op->getOperands()) { for (Value value : op->getOperands()) {
if (isa<TensorType>(value.getType())) { if (isa<TensorType>(value.getType())) {
op->emitOpError("tensor operand remains after PIM bufferization"); op->emitOpError("tensor operand remains after Pim bufferization");
++failureCount; ++failureCount;
return; return;
} }
} }
for (Value value : op->getResults()) { for (Value value : op->getResults()) {
if (isa<TensorType>(value.getType())) { if (isa<TensorType>(value.getType())) {
op->emitOpError("tensor result remains after PIM bufferization"); op->emitOpError("tensor result remains after Pim bufferization");
++failureCount; ++failureCount;
return; return;
} }
@@ -755,7 +782,7 @@ static LogicalResult verifyNoTensorValues(ModuleOp moduleOp) {
}); });
if (failureCount != 0) if (failureCount != 0)
moduleOp.emitError() << "found " << failureCount moduleOp.emitError() << "found " << failureCount
<< " tensor value(s) after PIM bufferization" << " tensor value(s) after Pim bufferization"
<< (failureCount == 8 ? " (first 8 reported)" : ""); << (failureCount == 8 ? " (first 8 reported)" : "");
return success(failureCount == 0); return success(failureCount == 0);
} }
@@ -766,7 +793,7 @@ struct PimBufferizationVerificationPass
StringRef getArgument() const override { return "pim-bufferization-verification"; } StringRef getArgument() const override { return "pim-bufferization-verification"; }
StringRef getDescription() const override { StringRef getDescription() const override {
return "Verify tensor elimination, contiguity, and PIM copy address spaces."; return "Verify tensor elimination, contiguity, and Pim copy address spaces.";
} }
void runOnOperation() final { void runOnOperation() final {
@@ -16,7 +16,7 @@ struct HostConstantFoldingPass : PassWrapper<HostConstantFoldingPass, OperationP
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(HostConstantFoldingPass) MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(HostConstantFoldingPass)
StringRef getArgument() const override { return "pim-host-constant-folding-pass"; } StringRef getArgument() const override { return "pim-host-constant-folding-pass"; }
StringRef getDescription() const override { return "Fold host-side constant expressions before PIM verification"; } StringRef getDescription() const override { return "Fold host-side constant expressions before Pim verification"; }
LogicalResult initialize(MLIRContext* context) override { LogicalResult initialize(MLIRContext* context) override {
RewritePatternSet owningPatterns(context); RewritePatternSet owningPatterns(context);
@@ -38,7 +38,7 @@ struct HostConstantFoldingPass : PassWrapper<HostConstantFoldingPass, OperationP
GreedyRewriteConfig config; GreedyRewriteConfig config;
config.enableFolding(); config.enableFolding();
if (failed(applyPatternsGreedily(moduleOp, *patterns, config))) { if (failed(applyPatternsGreedily(moduleOp, *patterns, config))) {
moduleOp.emitError("PIM host constant folding failed in the greedy rewrite driver"); moduleOp.emitError("Pim host constant folding failed in the greedy rewrite driver");
signalPassFailure(); signalPassFailure();
return; return;
} }
@@ -472,7 +472,7 @@ struct FoldConstantHostCopyPattern final : OpRewritePattern<memref::CopyOp> {
} }
}; };
// Converts PIM copies from dense globals into direct folded globals before codegen. // Converts Pim copies from dense globals into direct folded globals before codegen.
struct FoldConstantMemCpPattern final : OpRewritePattern<pim::PimMemCopyOp> { struct FoldConstantMemCpPattern final : OpRewritePattern<pim::PimMemCopyOp> {
using OpRewritePattern::OpRewritePattern; using OpRewritePattern::OpRewritePattern;
@@ -40,7 +40,7 @@ struct LowerTransposePattern final : OpRewritePattern<pim::PimTransposeOp> {
auto sourceType = dyn_cast<MemRefType>(op.getInput().getType()); auto sourceType = dyn_cast<MemRefType>(op.getInput().getType());
auto targetType = dyn_cast<MemRefType>(op.getOutputBuffer().getType()); auto targetType = dyn_cast<MemRefType>(op.getOutputBuffer().getType());
if (!sourceType || !targetType || !sourceType.hasStaticShape() || !targetType.hasStaticShape()) if (!sourceType || !targetType || !sourceType.hasStaticShape() || !targetType.hasStaticShape())
return op.emitOpError("requires static memref operands before PIM instruction selection"); return op.emitOpError("requires static memref operands before Pim instruction selection");
ArrayRef<int64_t> sourceShape = sourceType.getShape(); ArrayRef<int64_t> sourceShape = sourceType.getShape();
size_t rank = sourceShape.size(); size_t rank = sourceShape.size();
@@ -147,7 +147,7 @@ struct InstructionSelectionPass : PassWrapper<InstructionSelectionPass, Operatio
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(InstructionSelectionPass) MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(InstructionSelectionPass)
StringRef getArgument() const override { return "pim-instruction-selection"; } StringRef getArgument() const override { return "pim-instruction-selection"; }
StringRef getDescription() const override { return "Select explicit PIM ISA operations"; } StringRef getDescription() const override { return "Select explicit Pim ISA operations"; }
void runOnOperation() override { void runOnOperation() override {
RewritePatternSet patterns(&getContext()); RewritePatternSet patterns(&getContext());
@@ -36,7 +36,7 @@ struct PimLocalMemoryPlanningPass : PassWrapper<PimLocalMemoryPlanningPass, Oper
StringRef getArgument() const override { return "pim-local-memory-planning"; } StringRef getArgument() const override { return "pim-local-memory-planning"; }
StringRef getDescription() const override { StringRef getDescription() const override {
return "Plan liveness-based addresses for PIM core-local memory"; return "Plan liveness-based addresses for Pim core-local memory";
} }
void runOnOperation() override { void runOnOperation() override {
@@ -149,14 +149,14 @@ FailureOr<CoreMemoryPlan> buildCoreMemoryPlan(Operation* coreLikeOp) {
plan.intervals = std::move(*intervals); plan.intervals = std::move(*intervals);
auto placements = planLocalMemoryPlacements(plan.intervals, kPimLocalMemoryAddressLimit); auto placements = planLocalMemoryPlacements(plan.intervals, kPimLocalMemoryAddressLimit);
if (failed(placements)) { if (failed(placements)) {
coreLikeOp->emitError("PIM local-memory plan exceeds the signed int32 address range"); coreLikeOp->emitError("Pim local-memory plan exceeds the signed int32 address range");
return failure(); return failure();
} }
plan.placements = std::move(*placements); plan.placements = std::move(*placements);
for (const LocalMemoryPlacement& placement : plan.placements) { for (const LocalMemoryPlacement& placement : plan.placements) {
auto end = alignedEnd(placement.address, placement.size, kPimLocalMemoryAddressLimit); auto end = alignedEnd(placement.address, placement.size, kPimLocalMemoryAddressLimit);
if (failed(end)) { if (failed(end)) {
coreLikeOp->emitError("PIM local-memory plan has invalid address arithmetic"); coreLikeOp->emitError("Pim local-memory plan has invalid address arithmetic");
return failure(); return failure();
} }
plan.arenaSize = std::max(plan.arenaSize, *end); plan.arenaSize = std::max(plan.arenaSize, *end);
@@ -117,7 +117,7 @@ verifyLocalMemoryPlan(Operation* coreLikeOp, pim::CappedDiagnosticReporter& diag
for (StringRef name : kRemovedLocalMemoryPlanAttrNames) for (StringRef name : kRemovedLocalMemoryPlanAttrNames)
if (coreLikeOp->hasAttr(name)) { if (coreLikeOp->hasAttr(name)) {
diagnostics.report(coreLikeOp, [name](Operation* op) { diagnostics.report(coreLikeOp, [name](Operation* op) {
op->emitError() << "contains removed PIM local-memory planning attribute '" << name << "'"; op->emitError() << "contains removed Pim local-memory planning attribute '" << name << "'";
}); });
hasFailure = true; hasFailure = true;
} }
@@ -137,7 +137,7 @@ verifyLocalMemoryPlan(Operation* coreLikeOp, pim::CappedDiagnosticReporter& diag
auto analyzed = pim::analyzeLocalMemoryLifetimes(coreLikeOp); auto analyzed = pim::analyzeLocalMemoryLifetimes(coreLikeOp);
if (failed(analyzed)) { if (failed(analyzed)) {
diagnostics.report(coreLikeOp, [](Operation* op) { diagnostics.report(coreLikeOp, [](Operation* op) {
op->emitError("cannot analyze PIM local-memory lifetimes for plan verification"); op->emitError("cannot analyze Pim local-memory lifetimes for plan verification");
}); });
return failure(); return failure();
} }
@@ -156,7 +156,7 @@ verifyLocalMemoryPlan(Operation* coreLikeOp, pim::CappedDiagnosticReporter& diag
for (StringRef name : kRemovedLocalMemoryPlanAttrNames) for (StringRef name : kRemovedLocalMemoryPlanAttrNames)
if (allocation->hasAttr(name)) { if (allocation->hasAttr(name)) {
diagnostics.report(allocation, [name](Operation* op) { diagnostics.report(allocation, [name](Operation* op) {
op->emitOpError() << "contains removed PIM local-memory planning attribute '" << name << "'"; op->emitOpError() << "contains removed Pim local-memory planning attribute '" << name << "'";
}); });
hasFailure = true; hasFailure = true;
} }
@@ -171,7 +171,7 @@ verifyLocalMemoryPlan(Operation* coreLikeOp, pim::CappedDiagnosticReporter& diag
uint64_t address = static_cast<uint64_t>(addressAttr.getInt()); uint64_t address = static_cast<uint64_t>(addressAttr.getInt());
if (address % 4 != 0 || address > arenaSize || interval.size > arenaSize - address) { if (address % 4 != 0 || address > arenaSize || interval.size > arenaSize - address) {
diagnostics.report(allocation, [&](Operation* op) { diagnostics.report(allocation, [&](Operation* op) {
op->emitOpError() << "has invalid PIM local-memory range [" << address << ", " op->emitOpError() << "has invalid Pim local-memory range [" << address << ", "
<< (address <= arenaSize && interval.size <= arenaSize - address << (address <= arenaSize && interval.size <= arenaSize - address
? address + interval.size ? address + interval.size
: arenaSize) : arenaSize)
@@ -221,7 +221,7 @@ verifyLocalMemoryPlan(Operation* coreLikeOp, pim::CappedDiagnosticReporter& diag
memref::AllocOp otherAllocation = other.allocation; memref::AllocOp otherAllocation = other.allocation;
diagnostics.report(allocation, [&](Operation*) { diagnostics.report(allocation, [&](Operation*) {
auto diagnostic = allocation.emitOpError() auto diagnostic = allocation.emitOpError()
<< "PIM local-memory plan assigns simultaneously live allocations to overlapping ranges; first range [" << "Pim local-memory plan assigns simultaneously live allocations to overlapping ranges; first range ["
<< conflicting->first << ", " << conflicting->first + other.size << "), second range [" << address << conflicting->first << ", " << conflicting->first + other.size << "), second range [" << address
<< ", " << address + interval.size << "), live positions overlap at [" << ", " << address + interval.size << "), live positions overlap at ["
<< std::max(interval.start, other.start) << ", " << std::min(interval.end, other.end) << "]"; << std::max(interval.start, other.start) << ", " << std::min(interval.end, other.end) << "]";
@@ -471,7 +471,7 @@ static LogicalResult appendCoreCommunicationEvents(Block& block,
auto targetCoreId = resolveIndexValue(sendOp.getTargetCoreId(), knowledge); auto targetCoreId = resolveIndexValue(sendOp.getTargetCoreId(), knowledge);
if (failed(targetCoreId)) { if (failed(targetCoreId)) {
diagnostics.report(&op, [](Operation* illegalOp) { diagnostics.report(&op, [](Operation* illegalOp) {
illegalOp->emitOpError("cannot statically resolve send target core for PIM communication deadlock check"); illegalOp->emitOpError("cannot statically resolve send target core for Pim communication deadlock check");
}); });
return failure(); return failure();
} }
@@ -490,7 +490,7 @@ static LogicalResult appendCoreCommunicationEvents(Block& block,
if (failed(sourceCoreId)) { if (failed(sourceCoreId)) {
diagnostics.report(&op, [](Operation* illegalOp) { diagnostics.report(&op, [](Operation* illegalOp) {
illegalOp->emitOpError( illegalOp->emitOpError(
"cannot statically resolve receive source core for PIM communication deadlock check"); "cannot statically resolve receive source core for Pim communication deadlock check");
}); });
return failure(); return failure();
} }
@@ -530,7 +530,7 @@ static void printCommunicationWindow(llvm::raw_ostream& os,
static void printCommunicationDeadlockReport(const DenseMap<int64_t, CommunicationEventVector>& coreEvents, static void printCommunicationDeadlockReport(const DenseMap<int64_t, CommunicationEventVector>& coreEvents,
const DenseMap<int64_t, size_t>& programCounters, const DenseMap<int64_t, size_t>& programCounters,
ArrayRef<int64_t> cycle) { ArrayRef<int64_t> cycle) {
llvm::errs() << "\n=== PIM static communication deadlock report ===\n"; llvm::errs() << "\n=== Pim static communication deadlock report ===\n";
llvm::errs() << "wait cycle:"; llvm::errs() << "wait cycle:";
for (int64_t coreId : cycle) for (int64_t coreId : cycle)
llvm::errs() << " " << coreId; llvm::errs() << " " << coreId;
@@ -565,7 +565,7 @@ static void printCommunicationDeadlockReport(const DenseMap<int64_t, Communicati
continue; continue;
printCommunicationWindow(llvm::errs(), coreEvents, coreId, pcIt->second); printCommunicationWindow(llvm::errs(), coreEvents, coreId, pcIt->second);
} }
llvm::errs() << "=== end PIM static communication deadlock report ===\n\n"; llvm::errs() << "=== end Pim static communication deadlock report ===\n\n";
} }
static void emitCommunicationDeadlockCycle(ModuleOp moduleOp, static void emitCommunicationDeadlockCycle(ModuleOp moduleOp,
@@ -576,8 +576,8 @@ static void emitCommunicationDeadlockCycle(ModuleOp moduleOp,
auto diagnostic = auto diagnostic =
moduleOp.emitError() moduleOp.emitError()
<< "PIM communication deadlock check found a blocking send/receive cycle while statically simulating the " << "Pim communication deadlock check found a blocking send/receive cycle while statically simulating the "
"expanded per-core communication streams; see the PIM static communication deadlock report above"; "expanded per-core communication streams; see the Pim static communication deadlock report above";
for (int64_t coreId : cycle) { for (int64_t coreId : cycle) {
auto eventsIt = coreEvents.find(coreId); auto eventsIt = coreEvents.find(coreId);
@@ -728,7 +728,7 @@ static LogicalResult verifyNoStaticCommunicationDeadlock(ModuleOp moduleOp,
auto diagnostic = auto diagnostic =
moduleOp.emitError() moduleOp.emitError()
<< "PIM communication deadlock check stalled without finding a closed wait cycle; this usually means a " << "Pim communication deadlock check stalled without finding a closed wait cycle; this usually means a "
"send/receive peer is missing or ordered after a finished core"; "send/receive peer is missing or ordered after a finished core";
for (const auto& [coreId, events] : coreEvents) { for (const auto& [coreId, events] : coreEvents) {
size_t pc = programCounters[coreId]; size_t pc = programCounters[coreId];
@@ -746,7 +746,7 @@ struct VerificationPass : PassWrapper<VerificationPass, OperationPass<ModuleOp>>
StringRef getArgument() const override { return "verify-pim-pass"; } StringRef getArgument() const override { return "verify-pim-pass"; }
StringRef getDescription() const override { StringRef getDescription() const override {
return "Verify that bufferized PIM IR contains only explicit host/device transfers"; return "Verify that bufferized Pim IR contains only explicit host/device transfers";
} }
VerificationPass() {} VerificationPass() {}
@@ -763,7 +763,7 @@ struct VerificationPass : PassWrapper<VerificationPass, OperationPass<ModuleOp>>
pim::CappedDiagnosticReporter diagnostics; pim::CappedDiagnosticReporter diagnostics;
if (!hasTarget || failed(targetResources.verify())) { if (!hasTarget || failed(targetResources.verify())) {
moduleOp.emitError("PIM codegen verification requires valid injected target resources"); moduleOp.emitError("Pim codegen verification requires valid injected target resources");
signalPassFailure(); signalPassFailure();
return; return;
} }
@@ -792,7 +792,7 @@ struct VerificationPass : PassWrapper<VerificationPass, OperationPass<ModuleOp>>
return; return;
diagnostics.report(op, [](Operation* illegalOp) { diagnostics.report(op, [](Operation* illegalOp) {
illegalOp->emitError("illegal Spatial operation reached PIM codegen verification"); illegalOp->emitError("illegal Spatial operation reached Pim codegen verification");
}); });
}); });
@@ -833,7 +833,7 @@ struct VerificationPass : PassWrapper<VerificationPass, OperationPass<ModuleOp>>
if (!isAddressOnlyHostOp(&op)) { if (!isAddressOnlyHostOp(&op)) {
diagnostics.report(&op, [](Operation* illegalOp) { diagnostics.report(&op, [](Operation* illegalOp) {
illegalOp->emitOpError("illegal host-side runtime op remains after PIM bufferization; " illegalOp->emitOpError("illegal host-side runtime op remains after Pim bufferization; "
"fold it to constants or lower it into pim.core"); "fold it to constants or lower it into pim.core");
}); });
continue; continue;
@@ -849,7 +849,7 @@ struct VerificationPass : PassWrapper<VerificationPass, OperationPass<ModuleOp>>
if (diagnostics.hasFailure()) { if (diagnostics.hasFailure()) {
diagnostics.emitSuppressedSummary(moduleOp, "verification failures"); diagnostics.emitSuppressedSummary(moduleOp, "verification failures");
moduleOp.emitError("PIM codegen verification failed; see diagnostics above"); moduleOp.emitError("Pim codegen verification failed; see diagnostics above");
hasFailure = true; hasFailure = true;
} }
@@ -928,7 +928,7 @@ private:
bool hasFailure = false; bool hasFailure = false;
if (!isSupportedCoreInstructionOp(&op)) { if (!isSupportedCoreInstructionOp(&op)) {
diagnostics.report(&op, [](Operation* illegalOp) { diagnostics.report(&op, [](Operation* illegalOp) {
illegalOp->emitOpError("unsupported executable op reached PIM codegen verification"); illegalOp->emitOpError("unsupported executable op reached Pim codegen verification");
}); });
hasFailure = true; hasFailure = true;
} }
@@ -990,7 +990,7 @@ private:
if (failed(resolveIndexValue(storeOp.getHostTargetOffset(), knowledge)) if (failed(resolveIndexValue(storeOp.getHostTargetOffset(), knowledge))
|| failed(resolveIndexValue(storeOp.getDeviceSourceOffset(), knowledge))) { || failed(resolveIndexValue(storeOp.getDeviceSourceOffset(), knowledge))) {
diagnostics.report(&op, [](Operation* illegalOp) { diagnostics.report(&op, [](Operation* illegalOp) {
illegalOp->emitOpError("offset operands must be statically evaluable for PIM codegen"); illegalOp->emitOpError("offset operands must be statically evaluable for Pim codegen");
}); });
hasFailure = true; hasFailure = true;
} }
@@ -1006,7 +1006,7 @@ private:
if (failed(resolveIndexValue(loadOp.getDeviceTargetOffset(), knowledge)) if (failed(resolveIndexValue(loadOp.getDeviceTargetOffset(), knowledge))
|| failed(resolveIndexValue(loadOp.getHostSourceOffset(), knowledge))) { || failed(resolveIndexValue(loadOp.getHostSourceOffset(), knowledge))) {
diagnostics.report(&op, [](Operation* illegalOp) { diagnostics.report(&op, [](Operation* illegalOp) {
illegalOp->emitOpError("offset operands must be statically evaluable for PIM codegen"); illegalOp->emitOpError("offset operands must be statically evaluable for Pim codegen");
}); });
hasFailure = true; hasFailure = true;
} }
@@ -1022,7 +1022,7 @@ private:
if (failed(resolveIndexValue(copyOp.getTargetOffset(), knowledge)) if (failed(resolveIndexValue(copyOp.getTargetOffset(), knowledge))
|| failed(resolveIndexValue(copyOp.getSourceOffset(), knowledge))) { || failed(resolveIndexValue(copyOp.getSourceOffset(), knowledge))) {
diagnostics.report(&op, [](Operation* illegalOp) { diagnostics.report(&op, [](Operation* illegalOp) {
illegalOp->emitOpError("offset operands must be statically evaluable for PIM codegen"); illegalOp->emitOpError("offset operands must be statically evaluable for Pim codegen");
}); });
hasFailure = true; hasFailure = true;
} }
@@ -1032,7 +1032,7 @@ private:
&& failed(resolveIndexValue(receiveOp.getOutputOffset(), knowledge))) { && failed(resolveIndexValue(receiveOp.getOutputOffset(), knowledge))) {
diagnostics.report(&op, [](Operation* illegalOp) { diagnostics.report(&op, [](Operation* illegalOp) {
illegalOp->emitOpError( illegalOp->emitOpError(
"output offset must be statically evaluable for PIM codegen"); "output offset must be statically evaluable for Pim codegen");
}); });
hasFailure = true; hasFailure = true;
} }
+3 -3
View File
@@ -11,7 +11,7 @@ include "mlir/Dialect/Bufferization/IR/BufferViewFlowOpInterface.td"
def PimDialect : Dialect { def PimDialect : Dialect {
let name = "pim"; let name = "pim";
let summary = "A low-level dialect for the PIM coprocessors on ReRAM crossbars"; let summary = "A low-level dialect for the Pim coprocessors on ReRAM crossbars";
let cppNamespace = "::onnx_mlir::pim"; let cppNamespace = "::onnx_mlir::pim";
} }
@@ -27,7 +27,7 @@ def PimTensor :
def PimCoreOp : PimOp<"core", [SingleBlock, def PimCoreOp : PimOp<"core", [SingleBlock,
DeclareOpInterfaceMethods<OpAsmOpInterface, ["getAsmBlockArgumentNames"]>]> { DeclareOpInterfaceMethods<OpAsmOpInterface, ["getAsmBlockArgumentNames"]>]> {
let summary = "Execute a block on a PIM core"; let summary = "Execute a block on a Pim core";
let regions = (region SizedRegion<1>:$body); let regions = (region SizedRegion<1>:$body);
@@ -136,7 +136,7 @@ def PimWaitOp : PimOp<"wait", []> {
let arguments = (ins let arguments = (ins
Index:$eventRegister, Index:$eventRegister,
I32Attr:$waitValue Index:$waitValue
); );
let assemblyFormat = [{ let assemblyFormat = [{
@@ -242,11 +242,146 @@ static void appendReceive(BoundaryProgram &boundary,
target.collection, {slice}, {0, 1}, {target.position}, {lanes}, lanes}); target.collection, {slice}, {0, 1}, {target.position}, {lanes}, lanes});
} }
struct HostTransferRef {
ExternalTransferFamily *family = nullptr;
size_t index = 0;
};
static unsigned getBarrierRoundCount(size_t coreCount) {
unsigned rounds = 0;
for (size_t distance = 1; distance < coreCount; distance *= 2)
++rounds;
return rounds;
}
static LogicalResult assignPipelineSynchronization(
DeferredTransferPlan &transfers,
ArrayRef<BoundaryProgram> boundaries,
size_t synchronizationRegisterCount) {
bool pipelined = false;
for (ScheduledInfo &scheduled : transfers.scheduled) {
if (scheduled.pipelineStages.empty())
continue;
pipelined = true;
llvm::append_range(transfers.downstreamCores, scheduled.cores);
for (auto [core, stage] :
llvm::zip_equal(scheduled.cores, scheduled.pipelineStages))
if (stage == 0)
transfers.stageZeroCores.push_back(core);
}
if (!pipelined)
return success();
transfers.synchronizationRegisterCount = synchronizationRegisterCount;
llvm::sort(transfers.stageZeroCores);
transfers.stageZeroCores.erase(
llvm::unique(transfers.stageZeroCores), transfers.stageZeroCores.end());
llvm::sort(transfers.downstreamCores);
transfers.downstreamCores.erase(
llvm::unique(transfers.downstreamCores),
transfers.downstreamCores.end());
llvm::erase_if(transfers.downstreamCores, [&](int64_t core) {
return llvm::is_contained(transfers.stageZeroCores, core);
});
DenseMap<int64_t, SmallVector<HostTransferRef>> incomingByCore;
DenseMap<ExternalTransferFamily *, SmallVector<int64_t>> eventRegisters;
DenseMap<ExternalTransferFamily *, SmallVector<int64_t>> waitValues;
DenseMap<ExternalTransferFamily *, SmallVector<int64_t>> acknowledgementRegisters;
auto initialize = [&](ExternalTransferFamily &family) {
size_t count = family.targetCores.size();
eventRegisters.try_emplace(&family, count, 0);
waitValues.try_emplace(&family, count, 0);
acknowledgementRegisters.try_emplace(&family, count, 0);
};
for (const BoundaryProgram &boundary : boundaries)
for (const BoundaryInstruction &instruction : boundary.instructions) {
auto *receive = std::get_if<EmitReceiveAssemblyRun>(&instruction);
if (!receive || receive->slices.empty()
|| !receive->slices.front().family->hostRouted)
continue;
for (const ScheduledTransferSlice &slice : receive->slices) {
ExternalTransferFamily &family = *slice.family;
initialize(family);
for (size_t offset = 0; offset < slice.transferCount; ++offset) {
size_t index = slice.familyOffset + offset;
int64_t source = family.sourceCores.valueAt(index);
int64_t target = family.targetCores.valueAt(index);
incomingByCore[target].push_back({&family, index});
++transfers.hostAcknowledgementCounts[source];
}
}
}
unsigned barrierRounds = getBarrierRoundCount(
transfers.stageZeroCores.size());
bool stageZeroNeedsAcknowledgements = llvm::any_of(
transfers.stageZeroCores, [&](int64_t core) {
return transfers.hostAcknowledgementCounts.contains(core);
});
for (auto &[target, incoming] : incomingByCore) {
bool needsAcknowledgementRegister =
transfers.hostAcknowledgementCounts.contains(target);
bool stageZero = llvm::is_contained(transfers.stageZeroCores, target);
size_t reserved = stageZero
? barrierRounds + (stageZeroNeedsAcknowledgements ? 1 : 0)
: 1 + (needsAcknowledgementRegister ? 1 : 0);
if (reserved >= synchronizationRegisterCount) {
incoming.front().family->requirement->exchange->deferred.emitOpError(
"pipeline synchronization leaves no event register for incoming host transfers");
return failure();
}
size_t groupCount = std::min(
incoming.size(), synchronizationRegisterCount - reserved);
// One wait consumes a complete consecutive group of producer signals.
SmallVector<size_t> groupSizes(groupCount);
for (size_t ordinal = 0; ordinal < incoming.size(); ++ordinal)
++groupSizes[ordinal * groupCount / incoming.size()];
SmallVector<bool> first(groupCount, true);
for (size_t ordinal = 0; ordinal < incoming.size(); ++ordinal) {
size_t group = ordinal * groupCount / incoming.size();
HostTransferRef transfer = incoming[ordinal];
eventRegisters[transfer.family][transfer.index] = group;
acknowledgementRegisters[transfer.family][transfer.index] =
synchronizationRegisterCount - 1;
if (first[group]) {
waitValues[transfer.family][transfer.index] = groupSizes[group];
first[group] = false;
}
}
}
for (auto &[family, values] : eventRegisters) {
family->eventRegisters = StaticIntSequence::fromValues(values);
family->waitValues = StaticIntSequence::fromValues(waitValues[family]);
family->acknowledgementEventRegisters =
StaticIntSequence::fromValues(acknowledgementRegisters[family]);
}
if (!transfers.stageZeroCores.empty()) {
size_t reserved = barrierRounds
+ (stageZeroNeedsAcknowledgements ? 1 : 0);
if (reserved > synchronizationRegisterCount)
return transfers.scheduled.front().op->emitOpError(
"pipeline stage-zero barrier requires more synchronization registers than the target provides");
}
if (!transfers.downstreamCores.empty()) {
bool needsAcknowledgements = llvm::any_of(
transfers.downstreamCores, [&](int64_t core) {
return transfers.hostAcknowledgementCounts.contains(core);
});
if (1 + (needsAcknowledgements ? 1 : 0)
> synchronizationRegisterCount)
return transfers.scheduled.front().op->emitOpError(
"pipeline stage-zero release requires more synchronization registers than the target provides");
}
return success();
}
} // namespace } // namespace
FailureOr<DeferredBoundaryPlan> buildDeferredBoundaryPlan( FailureOr<DeferredBoundaryPlan> buildDeferredBoundaryPlan(
DeferredTransferPlan &transfers, DeferredTransferPlan &transfers,
const ScheduledCommunicationPlan &schedule) { const ScheduledCommunicationPlan &schedule,
size_t synchronizationRegisterCount) {
DeferredBoundaryPlan result; DeferredBoundaryPlan result;
SmallVector<BoundaryProgram> boundaries; SmallVector<BoundaryProgram> boundaries;
DenseMap<BoundaryKey, unsigned> indices; DenseMap<BoundaryKey, unsigned> indices;
@@ -373,6 +508,9 @@ FailureOr<DeferredBoundaryPlan> buildDeferredBoundaryPlan(
return std::tie(scheduledOrder[lhs.key.first], lhs.key.second) return std::tie(scheduledOrder[lhs.key.first], lhs.key.second)
< std::tie(scheduledOrder[rhs.key.first], rhs.key.second); < std::tie(scheduledOrder[rhs.key.first], rhs.key.second);
}); });
if (failed(assignPipelineSynchronization(
transfers, boundaries, synchronizationRegisterCount)))
return failure();
result.boundaries = std::move(boundaries); result.boundaries = std::move(boundaries);
return result; return result;
} }
@@ -53,6 +53,7 @@ struct DeferredBoundaryPlan {
}; };
mlir::FailureOr<DeferredBoundaryPlan> buildDeferredBoundaryPlan(DeferredTransferPlan& transfers, mlir::FailureOr<DeferredBoundaryPlan> buildDeferredBoundaryPlan(DeferredTransferPlan& transfers,
const ScheduledCommunicationPlan& schedule); const ScheduledCommunicationPlan& schedule,
size_t synchronizationRegisterCount);
} // namespace onnx_mlir::spatial } // namespace onnx_mlir::spatial
@@ -4,6 +4,7 @@
#include "DeferredBoundaryRealization.hpp" #include "DeferredBoundaryRealization.hpp"
#include "DeferredProjectionAnalysis.hpp" #include "DeferredProjectionAnalysis.hpp"
#include "DeferredResultRealization.hpp" #include "DeferredResultRealization.hpp"
#include "DeferredTransferPlanning.hpp"
#include "src/Accelerators/PIM/Common/IR/LoopUtils.hpp" #include "src/Accelerators/PIM/Common/IR/LoopUtils.hpp"
#include "src/Accelerators/PIM/Common/IR/StaticIntGrid.hpp" #include "src/Accelerators/PIM/Common/IR/StaticIntGrid.hpp"
#include "src/Accelerators/PIM/Common/IR/StaticIntSequence.hpp" #include "src/Accelerators/PIM/Common/IR/StaticIntSequence.hpp"
@@ -21,6 +22,8 @@ struct LogicalTransferMetadataView {
StaticIntSequenceChain targetCores; StaticIntSequenceChain targetCores;
StaticIntSequenceChain hostOffsets; StaticIntSequenceChain hostOffsets;
StaticIntSequenceChain eventRegisters; StaticIntSequenceChain eventRegisters;
StaticIntSequenceChain waitValues;
StaticIntSequenceChain acknowledgementEventRegisters;
StaticIntSequenceChain targetLanes; StaticIntSequenceChain targetLanes;
StaticIntSequenceChain localOffsets; StaticIntSequenceChain localOffsets;
SmallVector<StaticIntSequenceChain> projectionOffsets; SmallVector<StaticIntSequenceChain> projectionOffsets;
@@ -91,6 +94,9 @@ static void appendMetadata(const ScheduledTransferSlice &slice, LogicalTransferM
metadata.hostOffsets.append(family.hostOffsets, familyIndex, count); metadata.hostOffsets.append(family.hostOffsets, familyIndex, count);
metadata.eventRegisters.append( metadata.eventRegisters.append(
family.eventRegisters, familyIndex, count); family.eventRegisters, familyIndex, count);
metadata.waitValues.append(family.waitValues, familyIndex, count);
metadata.acknowledgementEventRegisters.append(
family.acknowledgementEventRegisters, familyIndex, count);
} }
metadata.targetLanes.append(StaticIntSequence::affine(targetLane, 1, count)); metadata.targetLanes.append(StaticIntSequence::affine(targetLane, 1, count));
if (family.requirement->producerLocalOffsets) if (family.requirement->producerLocalOffsets)
@@ -296,13 +302,21 @@ static FailureOr<Value> emitReceiveValue(ArrayRef<ScheduledTransferSlice> slices
if (failed(grids)) return failure(); if (failed(grids)) return failure();
std::optional<StaticIntGrid> hostOffsets; std::optional<StaticIntGrid> hostOffsets;
std::optional<StaticIntGrid> eventRegisters; std::optional<StaticIntGrid> eventRegisters;
std::optional<StaticIntGrid> waitValues;
std::optional<StaticIntGrid> acknowledgementEventRegisters;
if (slices.front().family->hostRouted) { if (slices.front().family->hostRouted) {
auto offsets = buildGrid(metadata.hostOffsets); auto offsets = buildGrid(metadata.hostOffsets);
auto events = buildGrid(metadata.eventRegisters); auto events = buildGrid(metadata.eventRegisters);
if (failed(offsets) || failed(events)) auto waits = buildGrid(metadata.waitValues);
auto acknowledgements = buildGrid(
metadata.acknowledgementEventRegisters);
if (failed(offsets) || failed(events) || failed(waits)
|| failed(acknowledgements))
return failure(); return failure();
hostOffsets = std::move(*offsets); hostOffsets = std::move(*offsets);
eventRegisters = std::move(*events); eventRegisters = std::move(*events);
waitValues = std::move(*waits);
acknowledgementEventRegisters = std::move(*acknowledgements);
} }
Value position = lane ? lane : context.constants.getIndex(0); Value position = lane ? lane : context.constants.getIndex(0);
Value row = context.constants.getIndex(0); Value row = context.constants.getIndex(0);
@@ -319,6 +333,10 @@ static FailureOr<Value> emitReceiveValue(ArrayRef<ScheduledTransferSlice> slices
hostOffsets->emitLookup( hostOffsets->emitLookup(
row, position, anchor, context.constants, context.rewriter, anchor->getLoc()), row, position, anchor, context.constants, context.rewriter, anchor->getLoc()),
eventRegisters->emitLookup( eventRegisters->emitLookup(
row, position, anchor, context.constants, context.rewriter, anchor->getLoc()),
waitValues->emitLookup(
row, position, anchor, context.constants, context.rewriter, anchor->getLoc()),
acknowledgementEventRegisters->emitLookup(
row, position, anchor, context.constants, context.rewriter, anchor->getLoc())); row, position, anchor, context.constants, context.rewriter, anchor->getLoc()));
receive = op; receive = op;
output = op.getOutput(); output = op.getOutput();
@@ -387,6 +405,8 @@ static FailureOr<Value> emitReceiveAssembly(const EmitReceiveAssemblyRun &run, V
std::optional<StaticIntGrid> positions; std::optional<StaticIntGrid> positions;
std::optional<StaticIntGrid> hostOffsets; std::optional<StaticIntGrid> hostOffsets;
std::optional<StaticIntGrid> eventRegisters; std::optional<StaticIntGrid> eventRegisters;
std::optional<StaticIntGrid> waitValues;
std::optional<StaticIntGrid> acknowledgementEventRegisters;
bool hostRouted = run.slices.front().family->hostRouted; bool hostRouted = run.slices.front().family->hostRouted;
auto metadataByEntry = buildRectangularReceiveMetadata(run, laneCount); auto metadataByEntry = buildRectangularReceiveMetadata(run, laneCount);
if (succeeded(metadataByEntry)) { if (succeeded(metadataByEntry)) {
@@ -402,10 +422,17 @@ static FailureOr<Value> emitReceiveAssembly(const EmitReceiveAssemblyRun &run, V
&LogicalTransferMetadataView::hostOffsets); &LogicalTransferMetadataView::hostOffsets);
auto events = buildRows( auto events = buildRows(
&LogicalTransferMetadataView::eventRegisters); &LogicalTransferMetadataView::eventRegisters);
if (failed(offsets) || failed(events)) auto waits = buildRows(
&LogicalTransferMetadataView::waitValues);
auto acknowledgements = buildRows(
&LogicalTransferMetadataView::acknowledgementEventRegisters);
if (failed(offsets) || failed(events) || failed(waits)
|| failed(acknowledgements))
return failure(); return failure();
hostOffsets = std::move(*offsets); hostOffsets = std::move(*offsets);
eventRegisters = std::move(*events); eventRegisters = std::move(*events);
waitValues = std::move(*waits);
acknowledgementEventRegisters = std::move(*acknowledgements);
} }
SmallVector<StaticIntSequence> positionRows; SmallVector<StaticIntSequence> positionRows;
for (unsigned position : run.positions) for (unsigned position : run.positions)
@@ -456,10 +483,17 @@ static FailureOr<Value> emitReceiveAssembly(const EmitReceiveAssemblyRun &run, V
&LogicalTransferMetadataView::hostOffsets); &LogicalTransferMetadataView::hostOffsets);
auto events = buildGrid( auto events = buildGrid(
&LogicalTransferMetadataView::eventRegisters); &LogicalTransferMetadataView::eventRegisters);
if (failed(offsets) || failed(events)) auto waits = buildGrid(
&LogicalTransferMetadataView::waitValues);
auto acknowledgements = buildGrid(
&LogicalTransferMetadataView::acknowledgementEventRegisters);
if (failed(offsets) || failed(events) || failed(waits)
|| failed(acknowledgements))
return failure(); return failure();
hostOffsets = std::move(*offsets); hostOffsets = std::move(*offsets);
eventRegisters = std::move(*events); eventRegisters = std::move(*events);
waitValues = std::move(*waits);
acknowledgementEventRegisters = std::move(*acknowledgements);
} }
SmallVector<StaticIntSequence> positionColumns; SmallVector<StaticIntSequence> positionColumns;
for (const StaticIntSequenceChain &values : positionsByLane) for (const StaticIntSequenceChain &values : positionsByLane)
@@ -491,6 +525,10 @@ static FailureOr<Value> emitReceiveAssembly(const EmitReceiveAssemblyRun &run, V
hostOffsets->emitLookup( hostOffsets->emitLookup(
entry, runtimeLane, anchor, context.constants, context.rewriter, loc), entry, runtimeLane, anchor, context.constants, context.rewriter, loc),
eventRegisters->emitLookup( eventRegisters->emitLookup(
entry, runtimeLane, anchor, context.constants, context.rewriter, loc),
waitValues->emitLookup(
entry, runtimeLane, anchor, context.constants, context.rewriter, loc),
acknowledgementEventRegisters->emitLookup(
entry, runtimeLane, anchor, context.constants, context.rewriter, loc)); entry, runtimeLane, anchor, context.constants, context.rewriter, loc));
receive = op; receive = op;
output = op.getOutput(); output = op.getOutput();
@@ -1158,9 +1196,199 @@ static LogicalResult emitBoundary(const BoundaryProgram &boundary, ArrayRef<Defe
return failed(values) ? failure() : replaceResults(exchanges, *values, replacements); return failed(values) ? failure() : replaceResults(exchanges, *values, replacements);
} }
static unsigned getBarrierRoundCount(size_t coreCount) {
unsigned rounds = 0;
for (size_t distance = 1; distance < coreCount; distance *= 2)
++rounds;
return rounds;
}
static LogicalResult emitCompletionSynchronization(
DeferredTransferPlan &transfers, DeferredEmissionContext &context) {
if (transfers.synchronizationRegisterCount == 0)
return success();
size_t acknowledgementRegister =
transfers.synchronizationRegisterCount - 1;
unsigned barrierRounds = getBarrierRoundCount(
transfers.stageZeroCores.size());
bool stageZeroNeedsAcknowledgements = llvm::any_of(
transfers.stageZeroCores, [&](int64_t core) {
return transfers.hostAcknowledgementCounts.contains(core);
});
size_t firstBarrierRegister = acknowledgementRegister
- (stageZeroNeedsAcknowledgements ? 1 : 0);
DenseMap<int64_t, unsigned> stageZeroRank;
for (auto [rank, core] : llvm::enumerate(transfers.stageZeroCores))
stageZeroRank[core] = rank;
DenseMap<int64_t, unsigned> downstreamRank;
for (auto [rank, core] : llvm::enumerate(transfers.downstreamCores))
downstreamRank[core] = rank;
auto getReleaseRegister = [&](int64_t core) {
return acknowledgementRegister
- (transfers.hostAcknowledgementCounts.contains(core) ? 1 : 0);
};
for (ScheduledInfo &scheduled : transfers.scheduled) {
Block *block = scheduled.blocks.front();
context.rewriter.setInsertionPoint(block->getTerminator());
Location loc = scheduled.op->getLoc();
Value lane;
if (auto batch = dyn_cast<SpatScheduledComputeBatch>(scheduled.op))
lane = *batch.getLaneArgument();
SmallVector<int64_t> acknowledgementCounts, releaseRegisters;
SmallVector<int64_t> releaseWaitValues, leftTargets, leftRegisters;
SmallVector<int64_t> rightTargets, rightRegisters;
LaneSet barrierLanes, leaderLanes, leftLanes, rightLanes;
for (auto [index, core] : llvm::enumerate(scheduled.cores)) {
acknowledgementCounts.push_back(
transfers.hostAcknowledgementCounts.lookup(core));
if (stageZeroRank.contains(core))
barrierLanes = barrierLanes.unite(
LaneSet::range(index, index + 1));
if (!transfers.stageZeroCores.empty()
&& core == transfers.stageZeroCores.front())
leaderLanes = leaderLanes.unite(LaneSet::range(index, index + 1));
auto rank = downstreamRank.find(core);
if (rank == downstreamRank.end()) {
releaseRegisters.push_back(0);
releaseWaitValues.push_back(0);
leftTargets.push_back(core);
leftRegisters.push_back(0);
rightTargets.push_back(core);
rightRegisters.push_back(0);
continue;
}
releaseRegisters.push_back(getReleaseRegister(core));
releaseWaitValues.push_back(1);
size_t left = 2 * rank->second + 1;
size_t right = left + 1;
if (left < transfers.downstreamCores.size()) {
int64_t child = transfers.downstreamCores[left];
leftTargets.push_back(child);
leftRegisters.push_back(getReleaseRegister(child));
leftLanes = leftLanes.unite(LaneSet::range(index, index + 1));
} else {
leftTargets.push_back(core);
leftRegisters.push_back(0);
}
if (right < transfers.downstreamCores.size()) {
int64_t child = transfers.downstreamCores[right];
rightTargets.push_back(child);
rightRegisters.push_back(getReleaseRegister(child));
rightLanes = rightLanes.unite(LaneSet::range(index, index + 1));
} else {
rightTargets.push_back(core);
rightRegisters.push_back(0);
}
}
Value runtimeLane = lane ? lane : context.constants.getIndex(0);
auto emitForLanes = [&](const LaneSet &active, auto emit) -> LogicalResult {
if (active.empty())
return success();
if (!lane) {
if (active.contains(0))
emit();
return success();
}
auto condition = emitLaneCondition(
active, lane, scheduled.cores.size(), scheduled.op, context, loc);
if (failed(condition))
return failure();
auto conditional = scf::IfOp::create(
context.rewriter, loc, TypeRange {}, *condition, false);
OpBuilder::InsertionGuard guard(context.rewriter);
context.rewriter.setInsertionPoint(
conditional.getThenRegion().front().getTerminator());
emit();
return success();
};
Value acknowledgementCount = emitStaticIntLookup(
StaticIntSequence::fromValues(acknowledgementCounts),
runtimeLane, scheduled.op,
context.constants, context.rewriter, loc);
SpatWaitOp::create(
context.rewriter, loc,
context.constants.getIndex(acknowledgementRegister),
acknowledgementCount);
// Dissemination barrier: every round doubles the covered stage-zero peers.
auto emitBarrier = [&]() {
for (unsigned round = 0; round < barrierRounds; ++round) {
SmallVector<int64_t> targets;
targets.reserve(scheduled.cores.size());
size_t distance = size_t {1} << round;
for (int64_t core : scheduled.cores) {
auto rank = stageZeroRank.find(core);
targets.push_back(rank == stageZeroRank.end()
? core
: transfers.stageZeroCores[
(rank->second + distance)
% transfers.stageZeroCores.size()]);
}
Value target = emitStaticIntLookup(
StaticIntSequence::fromValues(targets),
runtimeLane, scheduled.op,
context.constants, context.rewriter, loc);
Value eventRegister = context.constants.getIndex(
firstBarrierRegister - round);
SpatSyncOp::create(
context.rewriter, loc, target, eventRegister);
SpatWaitOp::create(
context.rewriter, loc, eventRegister,
context.constants.getIndex(1));
}
};
if (barrierRounds > 0
&& failed(emitForLanes(barrierLanes, emitBarrier)))
return failure();
// Gate downstream restarts so no core advances the simulator input
// iteration ahead of stage zero.
if (!transfers.downstreamCores.empty()
&& failed(emitForLanes(leaderLanes, [&]() {
int64_t root = transfers.downstreamCores.front();
SpatSyncOp::create(
context.rewriter, loc, context.constants.getIndex(root),
context.constants.getIndex(getReleaseRegister(root)));
})))
return failure();
Value releaseRegister = emitStaticIntLookup(
StaticIntSequence::fromValues(releaseRegisters), runtimeLane,
scheduled.op, context.constants, context.rewriter, loc);
Value releaseWaitValue = emitStaticIntLookup(
StaticIntSequence::fromValues(releaseWaitValues), runtimeLane,
scheduled.op, context.constants, context.rewriter, loc);
SpatWaitOp::create(
context.rewriter, loc, releaseRegister, releaseWaitValue);
auto emitChild = [&](ArrayRef<int64_t> targets,
ArrayRef<int64_t> registers) {
Value target = emitStaticIntLookup(
StaticIntSequence::fromValues(targets), runtimeLane, scheduled.op,
context.constants, context.rewriter, loc);
Value eventRegister = emitStaticIntLookup(
StaticIntSequence::fromValues(registers), runtimeLane, scheduled.op,
context.constants, context.rewriter, loc);
SpatSyncOp::create(context.rewriter, loc, target, eventRegister);
};
if (failed(emitForLanes(leftLanes, [&]() {
emitChild(leftTargets, leftRegisters);
}))
|| failed(emitForLanes(rightLanes, [&]() {
emitChild(rightTargets, rightRegisters);
})))
return failure();
}
return success();
}
} // namespace } // namespace
LogicalResult realizeDeferredBoundaries(ArrayRef<BoundaryProgram> boundaries, ArrayRef<DeferredResultPlan> results, DeferredEmissionContext &context, LogicalResult realizeDeferredBoundaries(ArrayRef<BoundaryProgram> boundaries, ArrayRef<DeferredResultPlan> results,
DeferredTransferPlan &transfers, DeferredEmissionContext &context,
DeferredReplacementMap &replacements) { DeferredReplacementMap &replacements) {
ScheduledInfo *scheduled = nullptr; ScheduledInfo *scheduled = nullptr;
for (const BoundaryProgram &boundary : boundaries) { for (const BoundaryProgram &boundary : boundaries) {
@@ -1171,7 +1399,7 @@ LogicalResult realizeDeferredBoundaries(ArrayRef<BoundaryProgram> boundaries, Ar
if (failed(emitBoundary(boundary, results, context, replacements))) if (failed(emitBoundary(boundary, results, context, replacements)))
return boundary.key.first->op->emitOpError("phase 2 failed to realize a communication boundary"); return boundary.key.first->op->emitOpError("phase 2 failed to realize a communication boundary");
} }
return success(); return emitCompletionSynchronization(transfers, context);
} }
} // namespace onnx_mlir::spatial } // namespace onnx_mlir::spatial
@@ -42,6 +42,7 @@ using DeferredReplacementMap =
mlir::LogicalResult realizeDeferredBoundaries(mlir::ArrayRef<BoundaryProgram> boundaries, mlir::LogicalResult realizeDeferredBoundaries(mlir::ArrayRef<BoundaryProgram> boundaries,
mlir::ArrayRef<DeferredResultPlan> results, mlir::ArrayRef<DeferredResultPlan> results,
DeferredTransferPlan& transfers,
DeferredEmissionContext& context, DeferredEmissionContext& context,
DeferredReplacementMap& replacements); DeferredReplacementMap& replacements);
@@ -236,6 +236,9 @@ struct ExternalTransferFamily {
StaticIntSequence channelIds = StaticIntSequence::uniform(0, 1); StaticIntSequence channelIds = StaticIntSequence::uniform(0, 1);
StaticIntSequence hostOffsets = StaticIntSequence::uniform(0, 1); StaticIntSequence hostOffsets = StaticIntSequence::uniform(0, 1);
StaticIntSequence eventRegisters = StaticIntSequence::uniform(0, 1); StaticIntSequence eventRegisters = StaticIntSequence::uniform(0, 1);
StaticIntSequence waitValues = StaticIntSequence::uniform(1, 1);
StaticIntSequence acknowledgementEventRegisters =
StaticIntSequence::uniform(0, 1);
bool hostRouted = false; bool hostRouted = false;
}; };
@@ -34,7 +34,9 @@ static LogicalResult verifyNoEscapingRegionValues(Operation* owner, StringRef ph
<< escapingUser->getName() << " at " << escapingUser->getLoc(); << escapingUser->getName() << " at " << escapingUser->getLoc();
} }
static LogicalResult placeLogicalProcessorsOnPhysicalCores(DeferredTransferPlan& plan, const SchedulingTarget& target) { static LogicalResult placeLogicalProcessorsOnPhysicalCores(
DeferredTransferPlan& plan, const SchedulingTarget& target,
size_t pipelineStages) {
std::vector<Cost> logicalTrafficFlits(target.processorCount * target.processorCount, 0); std::vector<Cost> logicalTrafficFlits(target.processorCount * target.processorCount, 0);
for (const std::unique_ptr<DeferredExchangePlan>& exchange : plan.exchanges) for (const std::unique_ptr<DeferredExchangePlan>& exchange : plan.exchanges)
for (const ExternalTransferFamily& transfer : exchange->external) { for (const ExternalTransferFamily& transfer : exchange->external) {
@@ -55,8 +57,15 @@ static LogicalResult placeLogicalProcessorsOnPhysicalCores(DeferredTransferPlan&
} }
} }
std::vector<size_t> placementGroups;
if (pipelineStages > 1) {
if (plan.processorStages.size() != target.processorCount)
return failure();
placementGroups = plan.processorStages;
}
std::vector<size_t> physicalCoreForLogicalProcessor = std::vector<size_t> physicalCoreForLogicalProcessor =
mapLogicalProcessorsToPhysicalCores(logicalTrafficFlits, target); mapLogicalProcessorsToPhysicalCores(
logicalTrafficFlits, target, placementGroups);
auto getPhysicalCore = [&](int64_t logicalProcessor) { auto getPhysicalCore = [&](int64_t logicalProcessor) {
assert(logicalProcessor >= 0 && static_cast<size_t>(logicalProcessor) < physicalCoreForLogicalProcessor.size() assert(logicalProcessor >= 0 && static_cast<size_t>(logicalProcessor) < physicalCoreForLogicalProcessor.size()
&& "logical processor is outside the scheduling target"); && "logical processor is outside the scheduling target");
@@ -218,7 +227,8 @@ LogicalResult realizeDeferredCommunication(func::FuncOp funcOp,
funcOp, materialization, pipelineStages, target.processorCount); funcOp, materialization, pipelineStages, target.processorCount);
if (failed(transfers)) if (failed(transfers))
return funcOp.emitOpError("phase 2 failed to build symbolic transfer families"); return funcOp.emitOpError("phase 2 failed to build symbolic transfer families");
if (failed(placeLogicalProcessorsOnPhysicalCores(*transfers, target))) if (failed(placeLogicalProcessorsOnPhysicalCores(
*transfers, target, pipelineStages)))
return failure(); return failure();
if (transfers->pipelineHostBufferBytes != 0) { if (transfers->pipelineHostBufferBytes != 0) {
auto bytes = pim::checkedCast<int64_t>( auto bytes = pim::checkedCast<int64_t>(
@@ -232,7 +242,8 @@ LogicalResult realizeDeferredCommunication(func::FuncOp funcOp,
auto schedule = scheduleDeferredCommunication(funcOp, *transfers); auto schedule = scheduleDeferredCommunication(funcOp, *transfers);
if (failed(schedule) || failed(verifyPlannedCommunicationDeadlockFree(funcOp, transfers->stepCounts, *schedule))) if (failed(schedule) || failed(verifyPlannedCommunicationDeadlockFree(funcOp, transfers->stepCounts, *schedule)))
return funcOp.emitOpError("phase 2 failed to schedule symbolic communication"); return funcOp.emitOpError("phase 2 failed to schedule symbolic communication");
auto boundaries = buildDeferredBoundaryPlan(*transfers, *schedule); auto boundaries = buildDeferredBoundaryPlan(
*transfers, *schedule, target.synchronizationRegisterCount);
if (failed(boundaries)) if (failed(boundaries))
return funcOp.emitOpError("phase 2 failed to build sparse boundary programs"); return funcOp.emitOpError("phase 2 failed to build sparse boundary programs");
@@ -242,7 +253,9 @@ LogicalResult realizeDeferredCommunication(func::FuncOp funcOp,
ConstantPool constants(funcOp, rewriter); ConstantPool constants(funcOp, rewriter);
DeferredEmissionContext context(rewriter, constants); DeferredEmissionContext context(rewriter, constants);
DeferredReplacementMap replacements; DeferredReplacementMap replacements;
if (failed(realizeDeferredBoundaries(boundaries->boundaries, boundaries->results, context, replacements))) if (failed(realizeDeferredBoundaries(
boundaries->boundaries, boundaries->results, *transfers,
context, replacements)))
return failure(); return failure();
for (auto [op, replacement] : replacements) { for (auto [op, replacement] : replacements) {
if (op->getResult(0) == replacement) if (op->getResult(0) == replacement)
@@ -32,9 +32,11 @@ static LogicalResult collectScheduledOperations(
DeferredTransferPlan &plan, DeferredTransferPlan &plan,
size_t pipelineStageCount, size_t pipelineStageCount,
size_t processorCount) { size_t processorCount) {
if (pipelineStageCount == 0 || processorCount % pipelineStageCount != 0) if (pipelineStageCount == 0
|| (pipelineStageCount > 1
&& materialization.processorStages.size() != processorCount))
return failure(); return failure();
size_t stageSize = processorCount / pipelineStageCount; plan.processorStages = materialization.processorStages;
unsigned nextStream = 0; unsigned nextStream = 0;
for (const ScheduledMaterializationRecord &record : for (const ScheduledMaterializationRecord &record :
materialization.materializedSchedules) { materialization.materializedSchedules) {
@@ -56,8 +58,12 @@ static LogicalResult collectScheduledOperations(
if (core >= processorCount) if (core >= processorCount)
return op.emitOpError("phase 2 scheduled core is outside the target"); return op.emitOpError("phase 2 scheduled core is outside the target");
info.cores.push_back(core); info.cores.push_back(core);
if (pipelineStageCount > 1) if (pipelineStageCount > 1) {
info.pipelineStages.push_back(core / stageSize); size_t stage = materialization.processorStages[core];
if (stage >= pipelineStageCount)
return op.emitOpError("phase 2 scheduled core has an invalid pipeline stage");
info.pipelineStages.push_back(stage);
}
} }
for (size_t lane = 0; lane < info.cores.size(); ++lane) for (size_t lane = 0; lane < info.cores.size(); ++lane)
info.streamIds.push_back(nextStream++); info.streamIds.push_back(nextStream++);
@@ -322,8 +328,7 @@ static LogicalResult buildRequirementFamilies(DeferredTransferPlan& plan,
static LogicalResult buildAvailabilityFamilies( static LogicalResult buildAvailabilityFamilies(
DeferredTransferPlan &plan, DeferredTransferPlan &plan,
DeferredExchangePlan& exchange, DeferredExchangePlan& exchange,
uint64_t& nextChannel, uint64_t& nextChannel) {
DenseMap<int64_t, DenseMap<int64_t, unsigned>>& eventRegistersByTarget) {
enum class Availability { Local, Direct, Host }; enum class Availability { Local, Direct, Host };
for (RequirementFamily& requirement : exchange.requirements) { for (RequirementFamily& requirement : exchange.requirements) {
for (LaneInterval interval : requirement.targetLanes.intervals()) { for (LaneInterval interval : requirement.targetLanes.intervals()) {
@@ -357,18 +362,6 @@ static LogicalResult buildAvailabilityFamilies(
family.channelIds = StaticIntSequence::affine(nextChannel, 1, count); family.channelIds = StaticIntSequence::affine(nextChannel, 1, count);
family.hostRouted = runAvailability == Availability::Host; family.hostRouted = runAvailability == Availability::Host;
if (family.hostRouted) { if (family.hostRouted) {
SmallVector<int64_t> eventRegisters;
for (int64_t targetCore : targetCores) {
auto &registers = eventRegistersByTarget[targetCore];
auto it = registers.try_emplace(
requirement.producer->core, registers.size()).first;
if (it->second >= kPimEventRegisterCount)
return exchange.deferred.emitOpError(
"pipeline host transfer requires more event registers than the target core provides");
eventRegisters.push_back(it->second);
}
family.eventRegisters = StaticIntSequence::fromValues(
eventRegisters);
auto fragmentType = dyn_cast<ShapedType>( auto fragmentType = dyn_cast<ShapedType>(
requirement.publicationFragmentType); requirement.publicationFragmentType);
auto fragmentBytes = fragmentType auto fragmentBytes = fragmentType
@@ -431,7 +424,6 @@ static LogicalResult buildExchanges(func::FuncOp funcOp, DeferredTransferPlan& p
funcOp.walk([&](SpatDeferredCommunicationOp op) { deferredOps.push_back(op); }); funcOp.walk([&](SpatDeferredCommunicationOp op) { deferredOps.push_back(op); });
GraphBatchPublicationCache publicationCache; GraphBatchPublicationCache publicationCache;
uint64_t nextChannel = 0; uint64_t nextChannel = 0;
DenseMap<int64_t, DenseMap<int64_t, unsigned>> eventRegistersByTarget;
for (SpatDeferredCommunicationOp deferred : deferredOps) { for (SpatDeferredCommunicationOp deferred : deferredOps) {
Operation* targetOp = deferred->getParentOfType<SpatScheduledCompute>(); Operation* targetOp = deferred->getParentOfType<SpatScheduledCompute>();
if (!targetOp) if (!targetOp)
@@ -451,8 +443,7 @@ static LogicalResult buildExchanges(func::FuncOp funcOp, DeferredTransferPlan& p
exchange->program = std::move(*program); exchange->program = std::move(*program);
if (failed(buildRequirementFamilies(plan, *exchange, publicationCache))) if (failed(buildRequirementFamilies(plan, *exchange, publicationCache)))
return failure(); return failure();
if (failed(buildAvailabilityFamilies( if (failed(buildAvailabilityFamilies(plan, *exchange, nextChannel)))
plan, *exchange, nextChannel, eventRegistersByTarget)))
return failure(); return failure();
plan.exchanges.push_back(std::move(exchange)); plan.exchanges.push_back(std::move(exchange));
} }
@@ -8,11 +8,16 @@
namespace onnx_mlir::spatial { namespace onnx_mlir::spatial {
struct DeferredTransferPlan { struct DeferredTransferPlan {
std::vector<size_t> processorStages;
llvm::SmallVector<ScheduledInfo, 0> scheduled; llvm::SmallVector<ScheduledInfo, 0> scheduled;
llvm::SmallVector<std::unique_ptr<ProducedValue>> producedStorage; llvm::SmallVector<std::unique_ptr<ProducedValue>> producedStorage;
llvm::DenseMap<int64_t, llvm::SmallVector<ProducedValue*>> producedByGraph; llvm::DenseMap<int64_t, llvm::SmallVector<ProducedValue*>> producedByGraph;
llvm::SmallVector<std::unique_ptr<DeferredExchangePlan>> exchanges; llvm::SmallVector<std::unique_ptr<DeferredExchangePlan>> exchanges;
llvm::SmallVector<unsigned> stepCounts; llvm::SmallVector<unsigned> stepCounts;
llvm::DenseMap<int64_t, unsigned> hostAcknowledgementCounts;
llvm::SmallVector<int64_t> stageZeroCores;
llvm::SmallVector<int64_t> downstreamCores;
size_t synchronizationRegisterCount = 0;
size_t pipelineHostBufferBytes = 0; size_t pipelineHostBufferBytes = 0;
}; };
@@ -808,7 +808,9 @@ materializeScheduledCompute(func::FuncOp funcOp,
} }
} }
return ScheduledComputeMaterializationResult {std::move(peftClassPlans), std::move(materializedSchedules), std::move(graphComputeToBlockMap)}; return ScheduledComputeMaterializationResult {
std::move(peftClassPlans), std::move(materializedSchedules),
std::move(graphComputeToBlockMap), schedule.processorStages};
} }
@@ -14,6 +14,7 @@ struct ScheduledComputeMaterializationResult {
llvm::MapVector<size_t, PeftClassPlan> peftClassPlans; llvm::MapVector<size_t, PeftClassPlan> peftClassPlans;
std::vector<ScheduledMaterializationRecord> materializedSchedules; std::vector<ScheduledMaterializationRecord> materializedSchedules;
DenseMap<GraphComputeBlockKey, Block *> graphComputeToBlockMap; DenseMap<GraphComputeBlockKey, Block *> graphComputeToBlockMap;
std::vector<size_t> processorStages;
}; };
FailureOr<BatchFragmentSpec> FailureOr<BatchFragmentSpec>
@@ -28,7 +28,7 @@ static bool hasValidTarget(const SchedulingTarget& target) {
static FailureOr<func::FuncOp> requireEntry(ModuleOp moduleOp) { static FailureOr<func::FuncOp> requireEntry(ModuleOp moduleOp) {
auto entry = getPimEntryFunc(moduleOp); auto entry = getPimEntryFunc(moduleOp);
if (failed(entry)) { if (failed(entry)) {
moduleOp.emitError("failed to locate the PIM entry function during Spatial scheduling and realization"); moduleOp.emitError("failed to locate the Pim entry function during Spatial scheduling and realization");
return failure(); return failure();
} }
return *entry; return *entry;
@@ -39,8 +39,9 @@ static SchedulingTarget getPipelineSchedulingTarget(
if (pipelineStages == 1) if (pipelineStages == 1)
return physicalTarget; return physicalTarget;
PipelineCoreLayout layout(physicalTarget.processorCount, pipelineStages);
SchedulingTarget schedulingTarget = physicalTarget; SchedulingTarget schedulingTarget = physicalTarget;
schedulingTarget.processorCount = physicalTarget.processorCount / pipelineStages; schedulingTarget.processorCount = layout.getLogicalProcessorCount();
schedulingTarget.residentWeightCapacity = checkedMultiply( schedulingTarget.residentWeightCapacity = checkedMultiply(
physicalTarget.residentWeightCapacity, pipelineStages); physicalTarget.residentWeightCapacity, pipelineStages);
schedulingTarget.interProcessorLatencyNs.assign( schedulingTarget.interProcessorLatencyNs.assign(
@@ -88,7 +89,10 @@ struct ScheduleAndRealizeSpatialPass final
signalPassFailure(); signalPassFailure();
return; return;
} }
if (pipelineStages == 0 || target.processorCount % pipelineStages != 0 PipelineCoreLayout pipelineLayout(target.processorCount, pipelineStages);
if (!pipelineLayout.isValid()
|| (pipelineStages > 1
&& target.synchronizationRegisterCount == 0)
|| target.residentWeightCapacity || target.residentWeightCapacity
> std::numeric_limits<size_t>::max() / pipelineStages) { > std::numeric_limits<size_t>::max() / pipelineStages) {
moduleOp.emitError("ScheduleAndRealizeSpatial requires valid pipeline stages and resource counts"); moduleOp.emitError("ScheduleAndRealizeSpatial requires valid pipeline stages and resource counts");
@@ -113,19 +117,24 @@ struct ScheduleAndRealizeSpatialPass final
scheduledGraph = analysis.getGraph(); scheduledGraph = analysis.getGraph();
schedule = std::move(analysis.getResult()); schedule = std::move(analysis.getResult());
std::string pipelineError; std::string pipelineError;
if (pipelineStages > 1) {
FailureOr<PipelineWorkloadPreparation> preparation =
preparePipelineWorkload(
scheduledGraph, schedule, pipelineStages, target, pipelineError);
if (failed(preparation)) {
moduleOp.emitError() << pipelineError;
signalPassFailure();
return;
}
if (*preparation == PipelineWorkloadPreparation::Changed)
continue;
}
if (succeeded(applyPipelineScheduling( if (succeeded(applyPipelineScheduling(
scheduledGraph, schedule, pipelineStages, target, pipelineError))) scheduledGraph, schedule, pipelineStages, target, pipelineError)))
break; break;
std::string splitError; moduleOp.emitError() << pipelineError;
if (pipelineStages == 1 signalPassFailure();
|| failed(splitPipelineWorkload( return;
scheduledGraph, schedule, pipelineStages, target, splitError))) {
if (!splitError.empty())
pipelineError = splitError;
moduleOp.emitError() << pipelineError;
signalPassFailure();
return;
}
} }
PatternRewriter rewriter(moduleOp.getContext()); PatternRewriter rewriter(moduleOp.getContext());
FailureOr<ScheduledComputeMaterializationResult> materialization = FailureOr<ScheduledComputeMaterializationResult> materialization =
@@ -1,16 +0,0 @@
#pragma once
#include "ScheduledComputeMaterialization.hpp"
#include "Scheduling/MergeSchedulingAnalysis.hpp"
#include <memory>
#include <optional>
namespace onnx_mlir::spatial {
struct ScheduledSpatialState {
std::optional<MergeScheduleResult> logicalSchedule;
std::optional<ScheduledComputeMaterializationResult> materialization;
};
} // namespace onnx_mlir::spatial
@@ -14,6 +14,7 @@ namespace spatial {
struct MergeScheduleResult { struct MergeScheduleResult {
size_t processorCount = 0; size_t processorCount = 0;
std::vector<size_t> processorStages;
std::vector<ComputeInstance> dominanceOrderCompute; std::vector<ComputeInstance> dominanceOrderCompute;
llvm::DenseMap<ComputeInstance, size_t> computeToCpuMap; llvm::DenseMap<ComputeInstance, size_t> computeToCpuMap;
llvm::DenseMap<ComputeInstance, size_t> computeToCpuSlotMap; llvm::DenseMap<ComputeInstance, size_t> computeToCpuSlotMap;
@@ -244,11 +244,13 @@ FailureOr<LanePublicationSignatures> buildLanePublicationSignatures(SpatComputeB
} // namespace } // namespace
std::vector<size_t> mapLogicalProcessorsToPhysicalCores(ArrayRef<Cost> logicalTrafficFlits, std::vector<size_t> mapLogicalProcessorsToPhysicalCores(ArrayRef<Cost> logicalTrafficFlits,
const SchedulingTarget& target) { const SchedulingTarget& target,
ArrayRef<size_t> placementGroups) {
const size_t processorCount = target.processorCount; const size_t processorCount = target.processorCount;
assert(logicalTrafficFlits.size() == processorCount * processorCount assert(logicalTrafficFlits.size() == processorCount * processorCount
&& "logical traffic matrix must cover every processor pair"); && "logical traffic matrix must cover every processor pair");
assert((placementGroups.empty() || placementGroups.size() == processorCount)
&& "physical placement groups must cover every processor");
std::vector<size_t> physicalCoreForLogicalProcessor(processorCount); std::vector<size_t> physicalCoreForLogicalProcessor(processorCount);
std::iota(physicalCoreForLogicalProcessor.begin(), physicalCoreForLogicalProcessor.end(), 0); std::iota(physicalCoreForLogicalProcessor.begin(), physicalCoreForLogicalProcessor.end(), 0);
@@ -266,6 +268,10 @@ std::vector<size_t> mapLogicalProcessorsToPhysicalCores(ArrayRef<Cost> logicalTr
for (size_t peerLogicalProcessor = 0; peerLogicalProcessor < processorCount; ++peerLogicalProcessor) { for (size_t peerLogicalProcessor = 0; peerLogicalProcessor < processorCount; ++peerLogicalProcessor) {
if (peerLogicalProcessor == logicalProcessor) if (peerLogicalProcessor == logicalProcessor)
continue; continue;
if (!placementGroups.empty()
&& placementGroups[peerLogicalProcessor]
!= placementGroups[logicalProcessor])
continue;
size_t physicalCore = physicalCoreForLogicalProcessor[logicalProcessor]; size_t physicalCore = physicalCoreForLogicalProcessor[logicalProcessor];
size_t peerPhysicalCore = physicalCoreForLogicalProcessor[peerLogicalProcessor]; size_t peerPhysicalCore = physicalCoreForLogicalProcessor[peerLogicalProcessor];
Cost currentCost = 0; Cost currentCost = 0;
@@ -29,7 +29,8 @@ inline Time getPeftTransferTime(const TransferCost& transferCost,
MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftScheduleOptions& options); MergeScheduleResult runPeftScheduler(const ComputeGraph& graph, const PeftScheduleOptions& options);
std::vector<size_t> mapLogicalProcessorsToPhysicalCores(llvm::ArrayRef<Cost> logicalTrafficFlits, std::vector<size_t> mapLogicalProcessorsToPhysicalCores(llvm::ArrayRef<Cost> logicalTrafficFlits,
const SchedulingTarget& target); const SchedulingTarget& target,
llvm::ArrayRef<size_t> placementGroups = {});
} // namespace spatial } // namespace spatial
} // namespace onnx_mlir } // namespace onnx_mlir
@@ -1,7 +1,10 @@
#include "llvm/ADT/DenseMap.h"
#include "llvm/ADT/SmallBitVector.h"
#include "llvm/ADT/STLExtras.h" #include "llvm/ADT/STLExtras.h"
#include <algorithm> #include <algorithm>
#include <limits> #include <limits>
#include <numeric>
#include <optional> #include <optional>
#include <queue> #include <queue>
#include <tuple> #include <tuple>
@@ -556,6 +559,10 @@ static LogicalResult splitBatchCompute(SpatGraphComputeBatch batch,
return failure(); return failure();
if (failed(verifySplittableVmmUses(splitBody->vmms, error))) if (failed(verifySplittableVmmUses(splitBody->vmms, error)))
return failure(); return failure();
if (!batch->hasAttr("pipeline.stage_group"))
batch->setAttr(
"pipeline.stage_group",
DistinctAttr::create(UnitAttr::get(batch.getContext())));
size_t partCount = std::min(pipelineStages, splitBody->vmms.size()); size_t partCount = std::min(pipelineStages, splitBody->vmms.size());
SmallVector<SmallVector<unsigned, 8>, 4> partitions = SmallVector<SmallVector<unsigned, 8>, 4> partitions =
@@ -601,11 +608,10 @@ static LogicalResult splitBatchCompute(SpatGraphComputeBatch batch,
return success(); return success();
} }
static LogicalResult splitPipelineWorkloadImpl(const ComputeGraph &graph, static FailureOr<PipelineWorkloadPreparation> preparePipelineWorkloadImpl(
const MergeScheduleResult &schedule, const ComputeGraph &graph, const MergeScheduleResult &schedule,
size_t pipelineStages, size_t pipelineStages, const SchedulingTarget &physicalTarget,
const SchedulingTarget &physicalTarget, std::string &error) {
std::string &error) {
size_t groupSize = schedule.processorCount; size_t groupSize = schedule.processorCount;
std::vector<TaskList> tasksByCpu(groupSize); std::vector<TaskList> tasksByCpu(groupSize);
for (size_t task = 0; task < graph.nodes.size(); ++task) { for (size_t task = 0; task < graph.nodes.size(); ++task) {
@@ -654,7 +660,7 @@ static LogicalResult splitPipelineWorkloadImpl(const ComputeGraph &graph,
std::string currentError; std::string currentError;
if (succeeded(splitBatchCompute( if (succeeded(splitBatchCompute(
batch, pipelineStages, physicalTarget, currentError))) batch, pipelineStages, physicalTarget, currentError)))
return success(); return PipelineWorkloadPreparation::Changed;
if (!currentError.empty()) if (!currentError.empty())
candidateError = currentError; candidateError = currentError;
} }
@@ -663,7 +669,8 @@ static LogicalResult splitPipelineWorkloadImpl(const ComputeGraph &graph,
: candidateError; : candidateError;
return failure(); return failure();
} }
return failure();
return PipelineWorkloadPreparation::Ready;
} }
bool fits(const ComputeGraph& graph, bool fits(const ComputeGraph& graph,
@@ -719,193 +726,585 @@ Cost findMaximumPackCost(const ComputeGraph& graph,
return low; return low;
} }
static Cost getCoreCost(const TaskList &tasks, const TaskCosts &taskCosts) { static Cost findMaximumIndexedPackCost(
Cost cost = 0; const TaskCosts &taskCosts,
for (size_t task : tasks) const std::vector<TaskList> &taskWeightIds, size_t weightCount,
cost = checkedAdd(cost, taskCosts[task]); const TaskList &tasks, size_t residentWeightCapacity,
return cost; size_t maximumPacks) {
} Cost low = 0;
Cost high = 0;
static Cost getStageMaximumAssemblyCost( for (size_t task : tasks) {
const std::vector<TaskList> &tasksByCpu, low = std::max(low, taskCosts[task]);
const TaskCosts &assemblyCosts, size_t groupSize, size_t stage) { high = checkedAdd(high, taskCosts[task]);
Cost maximum = 0;
for (size_t cpu = stage * groupSize;
cpu < (stage + 1) * groupSize; ++cpu)
maximum = std::max(
maximum, getCoreCost(tasksByCpu[cpu], assemblyCosts));
return maximum;
}
static bool fitsResidentWeights(const ComputeGraph &graph,
const TaskList &tasks, size_t candidate,
size_t residentWeightCapacity) {
ResidentWeightSet weights;
for (size_t task : tasks)
insertResidentWeights(weights, graph.nodes[task].residentWeights);
return getResidentWeightUnionSize(
weights, graph.nodes[candidate].residentWeights)
<= residentWeightCapacity;
}
static void repackPipelineStage(
const ComputeGraph &graph, const TaskCosts &schedulingCosts,
const TaskCosts &assemblyCosts,
std::vector<TaskList> &tasksByCpu, size_t groupSize, size_t stage,
size_t residentWeightCapacity) {
TaskList tasks;
Cost originalMaximum = 0;
for (size_t cpu = stage * groupSize;
cpu < (stage + 1) * groupSize; ++cpu) {
llvm::append_range(tasks, tasksByCpu[cpu]);
originalMaximum = std::max(
originalMaximum,
getCoreCost(tasksByCpu[cpu], schedulingCosts));
} }
llvm::sort(tasks, [&](size_t lhs, size_t rhs) { low = std::max(
low, high / maximumPacks + (high % maximumPacks != 0));
std::vector<size_t> seen(weightCount);
size_t generation = 0;
while (low < high) {
Cost middle = low + (high - low) / 2;
size_t packs = 1;
Cost cost = 0;
size_t packWeightCount = 0;
++generation;
bool fits = true;
bool packEmpty = true;
for (size_t task : tasks) {
size_t addedWeights = 0;
for (size_t weight : taskWeightIds[task])
addedWeights += seen[weight] != generation;
Cost taskCost = taskCosts[task];
bool startsNewPack = !packEmpty
&& (cost > middle - taskCost
|| packWeightCount + addedWeights > residentWeightCapacity);
if (startsNewPack) {
if (++packs > maximumPacks) {
fits = false;
break;
}
cost = 0;
packWeightCount = 0;
++generation;
packEmpty = true;
}
cost = checkedAdd(cost, taskCost);
for (size_t weight : taskWeightIds[task])
if (seen[weight] != generation) {
seen[weight] = generation;
++packWeightCount;
}
packEmpty = false;
}
if (fits)
high = middle;
else
low = middle + 1;
}
return low;
}
static size_t findMinimumIndexedPackCount(
const std::vector<TaskList> &taskWeightIds, const TaskList &tasks,
size_t residentWeightCapacity) {
if (tasks.empty())
return 0;
size_t packs = 1;
TaskList weights;
for (size_t task : tasks) {
size_t unionSize = weights.size();
for (size_t weight : taskWeightIds[task])
unionSize += !llvm::is_contained(weights, weight);
if (!weights.empty() && unionSize > residentWeightCapacity) {
++packs;
weights.clear();
}
for (size_t weight : taskWeightIds[task])
if (!llvm::is_contained(weights, weight))
weights.push_back(weight);
}
return packs;
}
struct PipelineGroup {
TaskList tasks;
TaskList successors;
size_t originalOrder = std::numeric_limits<size_t>::max();
TaskList weightIds;
bool consumesPipelineInput = false;
};
struct PipelineStageAssignment {
std::vector<size_t> taskStages;
std::vector<size_t> stageSizes;
};
static const TaskCosts &getPipelineBalanceCosts(
const ComputeGraph &graph, const PipelineTaskModel &model) {
for (size_t task = 0; task < graph.nodes.size(); ++task)
if (graph.nodes[task].instance.op && model.assemblyCosts[task] > 1
&& !model.predecessors[task].empty())
return model.assemblyCosts;
return model.schedulingCosts;
}
static bool consumesPipelineInput(const ComputeGraphNode &node) {
if (!node.instance.op)
return false;
return llvm::any_of(getComputeInstanceInputs(node.instance),
[](Value input) { return isa<BlockArgument>(input); });
}
static FailureOr<PipelineStageAssignment> assignPipelineStages(
const ComputeGraph &graph, const PipelineTaskModel &model,
const PipelineCoreLayout &layout, size_t residentWeightCapacity,
std::string &error) {
if (graph.nodes.empty())
return PipelineStageAssignment {
{}, std::vector<size_t>(layout.getStageSizes())};
std::vector<size_t> tasksByOrder(graph.nodes.size());
std::iota(tasksByOrder.begin(), tasksByOrder.end(), 0);
llvm::sort(tasksByOrder, [&](size_t lhs, size_t rhs) {
return graph.nodes[lhs].originalOrder < graph.nodes[rhs].originalOrder; return graph.nodes[lhs].originalOrder < graph.nodes[rhs].originalOrder;
}); });
std::vector<TaskList> packed(groupSize); std::vector<PipelineGroup> groups;
std::vector<ResidentWeightSet> weights(groupSize); std::vector<size_t> taskToGroup(graph.nodes.size());
TaskCosts loads(groupSize); ResidentWeightSet indexedWeights;
TaskCosts assemblyLoads(groupSize); std::vector<TaskList> taskWeightIds(graph.nodes.size());
for (size_t task : tasksByOrder)
for (const ResidentWeight &weight : graph.nodes[task].residentWeights) {
auto indexed = llvm::find(indexedWeights, weight);
size_t id = indexed - indexedWeights.begin();
if (indexed == indexedWeights.end()) {
id = indexedWeights.size();
indexedWeights.push_back(weight);
}
taskWeightIds[task].push_back(id);
}
llvm::DenseMap<Operation *, size_t> operationToGroup;
llvm::DenseMap<Attribute, size_t> splitOperationToGroup;
for (size_t task : tasksByOrder) {
Operation *operation = graph.nodes[task].instance.op;
Attribute splitGroup = operation
? operation->getAttr("pipeline.stage_group")
: Attribute();
size_t group;
auto existingSplit = splitGroup
? splitOperationToGroup.find(splitGroup)
: splitOperationToGroup.end();
auto existingOperation = operation && !splitGroup
? operationToGroup.find(operation)
: operationToGroup.end();
if (existingSplit != splitOperationToGroup.end()) {
group = existingSplit->second;
} else if (existingOperation != operationToGroup.end()) {
group = existingOperation->second;
} else {
group = groups.size();
groups.emplace_back();
if (splitGroup)
splitOperationToGroup[splitGroup] = group;
else if (operation)
operationToGroup[operation] = group;
}
taskToGroup[task] = group;
PipelineGroup &pipelineGroup = groups[group];
pipelineGroup.tasks.push_back(task);
pipelineGroup.originalOrder = std::min(
pipelineGroup.originalOrder, graph.nodes[task].originalOrder);
pipelineGroup.consumesPipelineInput |=
consumesPipelineInput(graph.nodes[task]);
for (size_t weight : taskWeightIds[task])
if (!llvm::is_contained(pipelineGroup.weightIds, weight))
pipelineGroup.weightIds.push_back(weight);
}
std::vector<size_t> indegree(groups.size());
for (size_t task = 0; task < graph.nodes.size(); ++task)
for (size_t predecessor : model.predecessors[task]) {
size_t source = taskToGroup[predecessor];
size_t target = taskToGroup[task];
if (source == target
|| llvm::is_contained(groups[source].successors, target))
continue;
groups[source].successors.push_back(target);
++indegree[target];
}
auto laterOriginalOrder = [&](size_t lhs, size_t rhs) {
return groups[lhs].originalOrder > groups[rhs].originalOrder;
};
std::priority_queue<size_t, std::vector<size_t>, decltype(laterOriginalOrder)>
ready(laterOriginalOrder);
for (size_t group = 0; group < groups.size(); ++group)
if (indegree[group] == 0)
ready.push(group);
TaskList groupOrder;
while (!ready.empty()) {
size_t group = ready.top();
ready.pop();
groupOrder.push_back(group);
for (size_t successor : groups[group].successors)
if (--indegree[successor] == 0)
ready.push(successor);
}
if (groupOrder.size() != groups.size()) {
error = "pipeline scheduling cannot keep every operation in one stage "
"because the collapsed operation graph is cyclic";
return failure();
}
std::vector<size_t> position(groups.size());
for (auto [index, group] : llvm::enumerate(groupOrder))
position[group] = index;
size_t minimumStageZeroEnd = 0;
for (auto [index, group] : llvm::enumerate(groupOrder))
if (groups[group].consumesPipelineInput)
minimumStageZeroEnd = index + 1;
std::vector<size_t> furthestSuccessorBefore(groups.size() + 1, 0);
bool hasCrossingEdge = false;
size_t furthestSuccessor = 0;
for (size_t cut = 1; cut <= groups.size(); ++cut) {
size_t group = groupOrder[cut - 1];
for (size_t successor : groups[group].successors) {
if (position[successor] <= position[group]) {
error = "pipeline scheduling operation order is not topological";
return failure();
}
furthestSuccessor = std::max(furthestSuccessor, position[successor]);
hasCrossingEdge = true;
}
furthestSuccessorBefore[cut] = furthestSuccessor;
}
auto partitionGroups = [&](ArrayRef<size_t> coreCounts)
-> FailureOr<std::vector<size_t>> {
const Cost infinity = std::numeric_limits<Cost>::max();
const size_t noCut = std::numeric_limits<size_t>::max();
std::vector<std::vector<Cost>> best(
coreCounts.size() + 1,
std::vector<Cost>(groups.size() + 1, infinity));
std::vector<std::vector<size_t>> parent(
coreCounts.size() + 1,
std::vector<size_t>(groups.size() + 1, noCut));
std::vector<size_t> stageCostCache(coreCounts.size());
std::vector<size_t> cachedCoreCounts;
std::vector<std::vector<std::vector<Cost>>> segmentCostCaches;
for (auto [stage, coreCount] : llvm::enumerate(coreCounts)) {
auto cached = llvm::find(cachedCoreCounts, coreCount);
if (cached == cachedCoreCounts.end()) {
stageCostCache[stage] = segmentCostCaches.size();
cachedCoreCounts.push_back(coreCount);
segmentCostCaches.emplace_back(
groups.size() + 1,
std::vector<Cost>(groups.size() + 1, infinity));
} else {
stageCostCache[stage] = cached - cachedCoreCounts.begin();
}
}
best[0][0] = 0;
// ponytail: operation groups are small; replace this quadratic partition
// only if scheduling profiles show it matters.
for (size_t stage = 0; stage < coreCounts.size(); ++stage) {
size_t coreCount = coreCounts[stage];
size_t stageWeightCapacity = checkedMultiply(
coreCount, residentWeightCapacity);
for (size_t start = 0; start < groups.size(); ++start) {
if (best[stage][start] == infinity)
continue;
TaskList segmentTasks;
llvm::SmallBitVector segmentWeights(indexedWeights.size());
size_t segmentWeightCount = 0;
for (size_t end = start + 1; end <= groups.size(); ++end) {
const PipelineGroup &group = groups[groupOrder[end - 1]];
llvm::append_range(segmentTasks, group.tasks);
for (size_t weight : group.weightIds)
if (!segmentWeights.test(weight)) {
segmentWeights.set(weight);
++segmentWeightCount;
}
if (segmentWeightCount > stageWeightCapacity)
break;
if (stage == 0 && end < minimumStageZeroEnd)
continue;
if (stage != 0 && hasCrossingEdge
&& furthestSuccessorBefore[start] >= end)
continue;
Cost &segmentCost =
segmentCostCaches[stageCostCache[stage]][start][end];
if (segmentCost == infinity)
segmentCost = findMaximumIndexedPackCost(
model.schedulingCosts, taskWeightIds, indexedWeights.size(),
segmentTasks, residentWeightCapacity, coreCount);
Cost maximumLoad = std::max(best[stage][start], segmentCost);
if (maximumLoad < best[stage + 1][end]) {
best[stage + 1][end] = maximumLoad;
parent[stage + 1][end] = start;
}
}
}
}
size_t usedStages = 0;
Cost bestLoad = infinity;
for (size_t stages = 1; stages <= coreCounts.size(); ++stages)
if (best[stages][groups.size()] != infinity
&& best[stages][groups.size()] <= bestLoad) {
bestLoad = best[stages][groups.size()];
usedStages = stages;
}
if (usedStages == 0) {
error = "pipeline scheduling cannot split operations into dependency-adjacent "
"stages within the physical crossbar limit";
return failure();
}
std::vector<size_t> groupStages(groups.size());
size_t end = groups.size();
for (size_t stage = usedStages; stage > 0; --stage) {
size_t start = parent[stage][end];
assert(start != noCut && "selected pipeline partition has no parent");
for (size_t position = start; position < end; ++position)
groupStages[groupOrder[position]] = stage - 1;
end = start;
}
return groupStages;
};
FailureOr<std::vector<size_t>> initialGroupStages =
partitionGroups(layout.getStageSizes());
if (failed(initialGroupStages))
return failure();
std::vector<TaskList> tasksByStage(layout.getStageCount());
for (size_t group : groupOrder)
llvm::append_range(
tasksByStage[(*initialGroupStages)[group]], groups[group].tasks);
std::vector<size_t> stageSizes(layout.getStageCount(), 1);
size_t assignedCores = stageSizes.size();
for (size_t stage = 0; stage < stageSizes.size(); ++stage) {
for (size_t task : tasksByStage[stage])
if (taskWeightIds[task].size() > residentWeightCapacity) {
error = "pipeline scheduling cannot fit one compute instance in a "
"physical core's crossbars";
return failure();
}
stageSizes[stage] = std::max(
stageSizes[stage], findMinimumIndexedPackCount(
taskWeightIds, tasksByStage[stage],
residentWeightCapacity));
assignedCores += stageSizes[stage] - 1;
}
if (assignedCores > layout.getProcessorCount()) {
error = "pipeline scheduling cannot fit dependency-adjacent stages "
"within the physical crossbar limit";
return failure();
}
const TaskCosts &balanceCosts = getPipelineBalanceCosts(graph, model);
auto getStageCost = [&](size_t stage, size_t coreCount) {
Cost schedulingCost = findMaximumIndexedPackCost(
model.schedulingCosts, taskWeightIds, indexedWeights.size(),
tasksByStage[stage], residentWeightCapacity, coreCount);
if (&balanceCosts == &model.schedulingCosts)
return schedulingCost;
Cost assemblyCost = findMaximumIndexedPackCost(
balanceCosts, taskWeightIds, indexedWeights.size(),
tasksByStage[stage], residentWeightCapacity, coreCount);
return std::max(schedulingCost, assemblyCost);
};
std::vector<Cost> stageCosts(stageSizes.size());
std::vector<Cost> nextStageCosts(stageSizes.size());
for (size_t stage = 0; stage < stageSizes.size(); ++stage)
stageCosts[stage] = getStageCost(stage, stageSizes[stage]);
for (size_t stage = 0; stage < stageSizes.size(); ++stage)
nextStageCosts[stage] = getStageCost(stage, stageSizes[stage] + 1);
while (assignedCores < layout.getProcessorCount()) {
size_t bestStage = 0;
Cost bestBenefit = 0;
for (size_t stage = 0; stage < stageSizes.size(); ++stage) {
Cost benefit = stageCosts[stage] - nextStageCosts[stage];
if (benefit > bestBenefit
|| (benefit == bestBenefit
&& (stageCosts[stage] > stageCosts[bestStage]
|| (stageCosts[stage] == stageCosts[bestStage]
&& stageSizes[stage] < stageSizes[bestStage])))) {
bestStage = stage;
bestBenefit = benefit;
}
}
++stageSizes[bestStage];
stageCosts[bestStage] = nextStageCosts[bestStage];
nextStageCosts[bestStage] =
getStageCost(bestStage, stageSizes[bestStage] + 1);
++assignedCores;
}
FailureOr<std::vector<size_t>> refinedGroupStages =
partitionGroups(stageSizes);
if (failed(refinedGroupStages))
return failure();
std::vector<size_t> taskStages(graph.nodes.size());
for (size_t task = 0; task < graph.nodes.size(); ++task)
taskStages[task] = (*refinedGroupStages)[taskToGroup[task]];
return PipelineStageAssignment {
std::move(taskStages), std::move(stageSizes)};
}
static bool packPipelineStage(
const ComputeGraph &graph, const TaskCosts &schedulingCosts,
const TaskCosts &assemblyCosts,
std::vector<TaskList> &tasksByCpu, const PipelineCoreLayout &layout,
ArrayRef<size_t> topologicalPosition, size_t stage,
size_t residentWeightCapacity, const SchedulingTarget &target,
std::vector<size_t> &taskCpus) {
PipelineStageRange range = layout.getStageRange(stage);
TaskList tasks;
for (size_t cpu = range.begin; cpu < range.begin + range.size; ++cpu)
llvm::append_range(tasks, tasksByCpu[cpu]);
llvm::sort(tasks, [&](size_t lhs, size_t rhs) {
return topologicalPosition[lhs] < topologicalPosition[rhs];
});
std::vector<TaskList> packed(range.size);
std::vector<ResidentWeightSet> weights(range.size);
TaskCosts loads(range.size);
TaskCosts assemblyLoads(range.size);
for (size_t task : tasks) { for (size_t task : tasks) {
std::optional<size_t> bestCore; std::optional<size_t> bestCore;
std::optional<std::tuple<Cost, Cost, size_t, size_t>> bestScore; using PackScore = std::tuple<Cost, Time, Cost, Cost, size_t>;
for (size_t core = 0; core < groupSize; ++core) { std::optional<PackScore> bestScore;
for (size_t core = 0; core < range.size; ++core) {
size_t unionSize = getResidentWeightUnionSize( size_t unionSize = getResidentWeightUnionSize(
weights[core], graph.nodes[task].residentWeights); weights[core], graph.nodes[task].residentWeights);
if (unionSize > residentWeightCapacity) if (unionSize > residentWeightCapacity)
continue; continue;
size_t addedWeights = unionSize - weights[core].size(); size_t addedWeights = unionSize - weights[core].size();
auto score = std::make_tuple( Cost assemblyLoad = checkedAdd(
checkedAdd(assemblyLoads[core], assemblyCosts[task]), assemblyLoads[core], assemblyCosts[task]);
checkedAdd(loads[core], schedulingCosts[task]), addedWeights, core); Cost schedulingLoad = checkedAdd(
loads[core], schedulingCosts[task]);
Time transferTime = 0;
size_t candidateCpu = range.begin + core;
for (const auto &[predecessor, transferCost] :
graph.predecessors[task])
if (taskCpus[predecessor] < target.processorCount)
transferTime = checkedAdd(
transferTime, getPeftTransferTime(
transferCost, taskCpus[predecessor],
candidateCpu, target));
PackScore score {
assemblyLoad, transferTime, schedulingLoad, addedWeights, core};
if (!bestScore || score < *bestScore) { if (!bestScore || score < *bestScore) {
bestCore = core; bestCore = core;
bestScore = score; bestScore = score;
} }
} }
if (!bestCore) if (!bestCore)
return; return false;
packed[*bestCore].push_back(task); packed[*bestCore].push_back(task);
insertResidentWeights( insertResidentWeights(
weights[*bestCore], graph.nodes[task].residentWeights); weights[*bestCore], graph.nodes[task].residentWeights);
loads[*bestCore] = checkedAdd(loads[*bestCore], schedulingCosts[task]); loads[*bestCore] = checkedAdd(loads[*bestCore], schedulingCosts[task]);
assemblyLoads[*bestCore] = checkedAdd( assemblyLoads[*bestCore] = checkedAdd(
assemblyLoads[*bestCore], assemblyCosts[task]); assemblyLoads[*bestCore], assemblyCosts[task]);
taskCpus[task] = range.begin + *bestCore;
} }
if (*std::max_element(loads.begin(), loads.end()) > originalMaximum) for (size_t core = 0; core < range.size; ++core)
return; tasksByCpu[range.begin + core] = std::move(packed[core]);
for (size_t core = 0; core < groupSize; ++core) return true;
tasksByCpu[stage * groupSize + core] = std::move(packed[core]);
} }
static void rebalancePipelineStages( static LogicalResult packPipelineStages(
const ComputeGraph &graph, const PipelineTaskModel &model, const ComputeGraph &graph, const PipelineTaskModel &model,
std::vector<TaskList> &tasksByCpu, size_t groupSize, std::vector<TaskList> &tasksByCpu, const PipelineCoreLayout &layout,
size_t pipelineStages, size_t residentWeightCapacity) { size_t pipelineStages, size_t residentWeightCapacity,
size_t minimumAssemblyFanIn = std::numeric_limits<size_t>::max(); const SchedulingTarget &target, size_t &failedStage,
std::string &error) {
std::vector<size_t> indegree(graph.nodes.size());
std::vector<TaskList> successors(graph.nodes.size());
for (size_t task = 0; task < graph.nodes.size(); ++task) for (size_t task = 0; task < graph.nodes.size(); ++task)
if (graph.nodes[task].instance.op && model.assemblyCosts[task] > 1 for (size_t predecessor : model.predecessors[task]) {
&& !model.predecessors[task].empty()) successors[predecessor].push_back(task);
minimumAssemblyFanIn = std::min( ++indegree[task];
minimumAssemblyFanIn, model.predecessors[task].size()); }
bool hasAssembly = minimumAssemblyFanIn != std::numeric_limits<size_t>::max(); auto laterOriginalOrder = [&](size_t lhs, size_t rhs) {
if (hasAssembly && groupSize < minimumAssemblyFanIn) return graph.nodes[lhs].originalOrder > graph.nodes[rhs].originalOrder;
return; };
const TaskCosts &balanceCosts = std::priority_queue<size_t, std::vector<size_t>, decltype(laterOriginalOrder)>
hasAssembly ? model.assemblyCosts : model.schedulingCosts; ready(laterOriginalOrder);
for (size_t task = 0; task < graph.nodes.size(); ++task)
if (indegree[task] == 0)
ready.push(task);
std::vector<size_t> topologicalPosition(graph.nodes.size());
size_t position = 0;
while (!ready.empty()) {
size_t task = ready.top();
ready.pop();
topologicalPosition[task] = position++;
for (size_t successor : successors[task])
if (--indegree[successor] == 0)
ready.push(successor);
}
if (position != graph.nodes.size()) {
error = "pipeline rebalancing received a cyclic task graph";
return failure();
}
Cost schedulingLimit = 0; const TaskCosts &balanceCosts = getPipelineBalanceCosts(graph, model);
for (const TaskList &tasks : tasksByCpu) std::vector<size_t> taskCpus(graph.nodes.size(), target.processorCount);
schedulingLimit = std::max(
schedulingLimit, getCoreCost(tasks, model.schedulingCosts));
for (size_t stage = 0; stage < pipelineStages; ++stage) for (size_t stage = 0; stage < pipelineStages; ++stage)
repackPipelineStage( if (!packPipelineStage(
graph, model.schedulingCosts, balanceCosts, tasksByCpu, graph, model.schedulingCosts, balanceCosts, tasksByCpu,
groupSize, stage, layout, topologicalPosition, stage, residentWeightCapacity,
residentWeightCapacity); target, taskCpus)) {
failedStage = stage;
error = "pipeline scheduling cannot pack dependency-monotone stage "
+ std::to_string(stage)
+ " within the physical crossbar limit";
return failure();
}
return success();
}
std::vector<size_t> taskToCpu(graph.nodes.size()); static LogicalResult verifyPipelineStageAssignment(
for (size_t cpu = 0; cpu < tasksByCpu.size(); ++cpu) const ComputeGraph &graph, const PipelineTaskModel &model,
for (size_t task : tasksByCpu[cpu]) const std::vector<TaskList> &tasksByCpu,
taskToCpu[task] = cpu; const PipelineCoreLayout &layout, std::string &error) {
const size_t noStage = std::numeric_limits<size_t>::max();
bool changed; std::vector<size_t> taskStages(graph.nodes.size(), noStage);
do { llvm::DenseMap<Operation *, size_t> operationStages;
changed = false; llvm::DenseMap<Attribute, size_t> splitOperationStages;
for (size_t sourceStage = pipelineStages; sourceStage-- > 1;) { for (size_t cpu = 0; cpu < tasksByCpu.size(); ++cpu) {
size_t targetStage = sourceStage - 1; std::optional<size_t> stage = layout.getStageForCore(cpu);
while (true) { if (!stage) {
Cost sourceMaximum = getStageMaximumAssemblyCost( error = "pipeline scheduling assigned a task outside the stage layout";
tasksByCpu, balanceCosts, groupSize, sourceStage); return failure();
Cost targetMaximum = getStageMaximumAssemblyCost( }
tasksByCpu, balanceCosts, groupSize, targetStage); for (size_t task : tasksByCpu[cpu]) {
if (targetMaximum >= sourceMaximum) if (task >= graph.nodes.size() || taskStages[task] != noStage) {
break; error = "pipeline scheduling did not assign every task exactly once";
return failure();
struct Move { }
size_t sourceCpu; taskStages[task] = *stage;
size_t targetCpu; if (*stage != 0 && consumesPipelineInput(graph.nodes[task])) {
size_t task; error = "pipeline scheduling assigned a direct function-input "
}; "consumer after stage zero";
std::optional<Move> best; return failure();
std::optional<std::tuple<size_t, Cost, Cost, size_t>> bestScore; }
for (size_t sourceCpu = sourceStage * groupSize; Operation *operation = graph.nodes[task].instance.op;
sourceCpu < (sourceStage + 1) * groupSize; ++sourceCpu) { if (!operation)
if (tasksByCpu[sourceCpu].empty()) continue;
continue; Attribute splitGroup = operation->getAttr("pipeline.stage_group");
size_t task = tasksByCpu[sourceCpu].front(); bool consistent;
bool dependenciesReady = llvm::all_of( if (splitGroup) {
model.predecessors[task], [&](size_t predecessor) { auto [entry, inserted] =
return taskToCpu[predecessor] / groupSize <= targetStage; splitOperationStages.try_emplace(splitGroup, *stage);
}); consistent = inserted || entry->second == *stage;
if (!dependenciesReady) } else {
continue; auto [entry, inserted] =
operationStages.try_emplace(operation, *stage);
Cost sourceAfter = getCoreCost( consistent = inserted || entry->second == *stage;
tasksByCpu[sourceCpu], balanceCosts) }
- balanceCosts[task]; if (!consistent) {
for (size_t targetCpu = targetStage * groupSize; error = "pipeline scheduling split one operation across stages";
targetCpu < (targetStage + 1) * groupSize; ++targetCpu) { return failure();
const TaskList &targetTasks = tasksByCpu[targetCpu];
if (!fitsResidentWeights(
graph, targetTasks, task, residentWeightCapacity))
continue;
Cost targetAfter = checkedAdd(
getCoreCost(targetTasks, balanceCosts), balanceCosts[task]);
Cost targetSchedulingAfter = checkedAdd(
getCoreCost(targetTasks, model.schedulingCosts),
model.schedulingCosts[task]);
if (targetAfter >= sourceMaximum
|| targetSchedulingAfter > schedulingLimit)
continue;
auto score = std::make_tuple(
graph.nodes[task].originalOrder,
std::max(sourceAfter, targetAfter), targetAfter, targetCpu);
if (!bestScore || score < *bestScore) {
best = Move {sourceCpu, targetCpu, task};
bestScore = score;
}
}
}
if (!best)
break;
tasksByCpu[best->sourceCpu].erase(
tasksByCpu[best->sourceCpu].begin());
TaskList &targetTasks = tasksByCpu[best->targetCpu];
auto insertion = llvm::find_if(targetTasks, [&](size_t task) {
return graph.nodes[task].originalOrder
> graph.nodes[best->task].originalOrder;
});
targetTasks.insert(insertion, best->task);
taskToCpu[best->task] = best->targetCpu;
changed = true;
} }
} }
} while (changed); }
if (llvm::is_contained(taskStages, noStage)) {
error = "pipeline scheduling did not assign every task exactly once";
return failure();
}
for (size_t task = 0; task < graph.nodes.size(); ++task)
for (size_t predecessor : model.predecessors[task])
if (taskStages[predecessor] > taskStages[task]
|| taskStages[task] - taskStages[predecessor] > 1) {
error = "pipeline scheduling produced a backward or skipped-stage dependency";
return failure();
}
return success();
} }
mlir::LogicalResult assignPipelineCores(const ComputeGraph& graph, mlir::LogicalResult assignPipelineCores(const ComputeGraph& graph,
@@ -914,7 +1313,13 @@ mlir::LogicalResult assignPipelineCores(const ComputeGraph& graph,
const SchedulingTarget& physicalTarget, const SchedulingTarget& physicalTarget,
std::string& error) { std::string& error) {
const size_t groupSize = schedule.processorCount; const size_t groupSize = schedule.processorCount;
std::vector<TaskList> tasksByCpu(groupSize); PipelineCoreLayout balancedLayout(
physicalTarget.processorCount, pipelineStages);
if (!balancedLayout.isValid()
|| groupSize != balancedLayout.getLogicalProcessorCount()) {
error = "pipeline scheduling received an incompatible physical core layout";
return mlir::failure();
}
for (size_t task = 0; task < graph.nodes.size(); ++task) { for (size_t task = 0; task < graph.nodes.size(); ++task) {
const ComputeInstance& instance = graph.nodes[task].instance; const ComputeInstance& instance = graph.nodes[task].instance;
auto cpu = schedule.computeToCpuMap.find(instance); auto cpu = schedule.computeToCpuMap.find(instance);
@@ -924,74 +1329,69 @@ mlir::LogicalResult assignPipelineCores(const ComputeGraph& graph,
error = "pipeline scheduling received an incomplete PEFT schedule"; error = "pipeline scheduling received an incomplete PEFT schedule";
return mlir::failure(); return mlir::failure();
} }
tasksByCpu[cpu->second].push_back(task);
} }
for (TaskList& tasks : tasksByCpu)
llvm::sort(tasks, [&](size_t lhs, size_t rhs) {
return schedule.computeToCpuSlotMap.lookup(graph.nodes[lhs].instance)
< schedule.computeToCpuSlotMap.lookup(graph.nodes[rhs].instance);
});
PipelineTaskModel taskModel = getPipelineTaskModel( PipelineTaskModel taskModel = getPipelineTaskModel(
graph, schedule, physicalTarget); graph, schedule, physicalTarget);
const TaskCosts &taskCosts = taskModel.schedulingCosts; FailureOr<PipelineStageAssignment> assignment = assignPipelineStages(
graph, taskModel, balancedLayout,
physicalTarget.residentWeightCapacity, error);
if (failed(assignment))
return failure();
std::vector<TaskList> tasksByPhysicalCpu(physicalTarget.processorCount); std::vector<TaskList> tasksByPhysicalCpu(physicalTarget.processorCount);
for (size_t sourceCpu = 0; sourceCpu < groupSize; ++sourceCpu) { std::vector<size_t> minimumPackableStageSizes(pipelineStages, 1);
const TaskList& tasks = tasksByCpu[sourceCpu]; std::string packingError;
if (tasks.empty()) bool packed = false;
continue; for (size_t attempt = 0; attempt < physicalTarget.processorCount; ++attempt) {
for (size_t task : tasks) PipelineCoreLayout candidateLayout(assignment->stageSizes);
if (graph.nodes[task].residentWeights.size() > physicalTarget.residentWeightCapacity) { for (TaskList &tasks : tasksByPhysicalCpu)
error = "pipeline scheduling cannot fit one compute instance in a physical core's crossbars"; tasks.clear();
return mlir::failure(); for (size_t task = 0; task < graph.nodes.size(); ++task) {
} PipelineStageRange range =
candidateLayout.getStageRange(assignment->taskStages[task]);
Cost maximumCost = findMaximumPackCost( tasksByPhysicalCpu[range.begin].push_back(task);
graph, taskCosts, tasks, physicalTarget.residentWeightCapacity, pipelineStages);
if (!fits(graph, taskCosts, tasks, maximumCost,
physicalTarget.residentWeightCapacity, pipelineStages)) {
error = "pipeline scheduling cannot partition one PEFT core within the physical crossbar limit";
return mlir::failure();
} }
size_t failedStage = 0;
const size_t desiredPacks = std::min(pipelineStages, tasks.size()); if (succeeded(packPipelineStages(
size_t stage = 0; graph, taskModel, tasksByPhysicalCpu, candidateLayout,
Cost packCost = 0; pipelineStages, physicalTarget.residentWeightCapacity,
ResidentWeightSet packWeights; physicalTarget, failedStage, packingError))) {
bool packEmpty = true; packed = true;
for (size_t index = 0; index < tasks.size(); ++index) { break;
size_t task = tasks[index];
const ComputeGraphNode& node = graph.nodes[task];
Cost taskCost = taskCosts[task];
bool exceedsLimit =
!packEmpty
&& (packCost > maximumCost - taskCost
|| getResidentWeightUnionSize(packWeights, node.residentWeights) > physicalTarget.residentWeightCapacity);
bool reserveOneTaskPerPack = !packEmpty && tasks.size() - index == desiredPacks - stage - 1;
if (exceedsLimit || reserveOneTaskPerPack) {
++stage;
packCost = 0;
packWeights.clear();
packEmpty = true;
}
if (stage >= pipelineStages) {
error = "pipeline scheduling produced too many packs";
return mlir::failure();
}
size_t physicalCpu = sourceCpu + stage * groupSize;
tasksByPhysicalCpu[physicalCpu].push_back(task);
packCost = checkedAdd(packCost, taskCost);
insertResidentWeights(packWeights, node.residentWeights);
packEmpty = false;
} }
minimumPackableStageSizes[failedStage] = std::max(
minimumPackableStageSizes[failedStage],
assignment->stageSizes[failedStage] + 1);
std::optional<size_t> donor;
for (size_t stage = 0; stage < pipelineStages; ++stage)
if (stage != failedStage
&& assignment->stageSizes[stage]
> minimumPackableStageSizes[stage]
&& (!donor
|| assignment->stageSizes[stage]
> assignment->stageSizes[*donor]))
donor = stage;
if (!donor)
break;
--assignment->stageSizes[*donor];
++assignment->stageSizes[failedStage];
} }
if (!packed) {
rebalancePipelineStages( error = packingError;
graph, taskModel, tasksByPhysicalCpu, groupSize, pipelineStages, return failure();
physicalTarget.residentWeightCapacity); }
PipelineCoreLayout pipelineLayout(assignment->stageSizes);
if (failed(verifyPipelineStageAssignment(
graph, taskModel, tasksByPhysicalCpu, pipelineLayout, error)))
return failure();
schedule.computeToCpuMap.clear(); schedule.computeToCpuMap.clear();
schedule.processorCount = physicalTarget.processorCount; schedule.processorCount = physicalTarget.processorCount;
schedule.processorStages.resize(physicalTarget.processorCount);
for (size_t stage = 0; stage < pipelineLayout.getStageCount(); ++stage) {
PipelineStageRange range = pipelineLayout.getStageRange(stage);
std::fill_n(
schedule.processorStages.begin() + range.begin, range.size, stage);
}
schedule.computeToCpuSlotMap.clear(); schedule.computeToCpuSlotMap.clear();
schedule.computeToAestMap.clear(); schedule.computeToAestMap.clear();
schedule.isLastComputeOfCpu.clear(); schedule.isLastComputeOfCpu.clear();
@@ -1071,21 +1471,22 @@ mlir::LogicalResult applyPipelineScheduling(const ComputeGraph& graph,
std::string& error) { std::string& error) {
if (pipelineStages == 1) if (pipelineStages == 1)
return mlir::success(); return mlir::success();
if (pipelineStages == 0 || schedule.processorCount == 0 PipelineCoreLayout pipelineLayout(
|| schedule.processorCount > std::numeric_limits<size_t>::max() / pipelineStages physicalTarget.processorCount, pipelineStages);
|| schedule.processorCount * pipelineStages != physicalTarget.processorCount) { if (!pipelineLayout.isValid() || schedule.processorCount == 0
error = "pipeline scheduling requires physical cores = scheduled cores * pipeline stages"; || schedule.processorCount
!= pipelineLayout.getLogicalProcessorCount()) {
error = "pipeline scheduling requires a valid balanced physical core layout";
return mlir::failure(); return mlir::failure();
} }
return assignPipelineCores(graph, schedule, pipelineStages, physicalTarget, error); return assignPipelineCores(graph, schedule, pipelineStages, physicalTarget, error);
} }
mlir::LogicalResult splitPipelineWorkload(const ComputeGraph &graph, mlir::FailureOr<PipelineWorkloadPreparation> preparePipelineWorkload(
const MergeScheduleResult &schedule, const ComputeGraph &graph, const MergeScheduleResult &schedule,
size_t pipelineStages, size_t pipelineStages, const SchedulingTarget &physicalTarget,
const SchedulingTarget &physicalTarget, std::string &error) {
std::string &error) { return preparePipelineWorkloadImpl(
return splitPipelineWorkloadImpl(
graph, schedule, pipelineStages, physicalTarget, error); graph, schedule, pipelineStages, physicalTarget, error);
} }
@@ -2,8 +2,15 @@
#include "mlir/Support/LogicalResult.h" #include "mlir/Support/LogicalResult.h"
#include "llvm/ADT/ArrayRef.h"
#include "llvm/ADT/STLExtras.h"
#include <algorithm>
#include <cstddef> #include <cstddef>
#include <numeric>
#include <optional>
#include <string> #include <string>
#include <vector>
#include "ComputeGraph.hpp" #include "ComputeGraph.hpp"
#include "MergeSchedule.hpp" #include "MergeSchedule.hpp"
@@ -11,16 +18,82 @@
namespace onnx_mlir::spatial { namespace onnx_mlir::spatial {
struct PipelineStageRange {
size_t begin;
size_t size;
};
class PipelineCoreLayout {
public:
PipelineCoreLayout(size_t processorCount, size_t stageCount)
: processorCount(processorCount), stageSizes(stageCount) {
if (stageCount == 0)
return;
size_t baseSize = processorCount / stageCount;
size_t largerStageCount = processorCount % stageCount;
for (size_t stage = 0; stage < stageCount; ++stage)
stageSizes[stage] = baseSize + (stage < largerStageCount);
}
explicit PipelineCoreLayout(llvm::ArrayRef<size_t> stageSizes)
: processorCount(std::accumulate(
stageSizes.begin(), stageSizes.end(), size_t {0})),
stageSizes(stageSizes.begin(), stageSizes.end()) {}
bool isValid() const {
return !stageSizes.empty()
&& llvm::none_of(stageSizes, [](size_t size) { return size == 0; });
}
size_t getLogicalProcessorCount() const {
return isValid()
? *std::min_element(stageSizes.begin(), stageSizes.end())
: 0;
}
size_t getStageCount() const { return stageSizes.size(); }
size_t getProcessorCount() const { return processorCount; }
llvm::ArrayRef<size_t> getStageSizes() const { return stageSizes; }
PipelineStageRange getStageRange(size_t stage) const {
return {std::accumulate(
stageSizes.begin(), stageSizes.begin() + stage, size_t {0}),
stageSizes[stage]};
}
std::optional<size_t> getStageForCore(size_t core) const {
if (!isValid() || core >= processorCount)
return std::nullopt;
size_t end = 0;
for (auto [stage, size] : llvm::enumerate(stageSizes)) {
end += size;
if (core < end)
return stage;
}
return std::nullopt;
}
private:
size_t processorCount;
std::vector<size_t> stageSizes;
};
mlir::LogicalResult applyPipelineScheduling(const ComputeGraph& graph, mlir::LogicalResult applyPipelineScheduling(const ComputeGraph& graph,
MergeScheduleResult& schedule, MergeScheduleResult& schedule,
size_t pipelineStages, size_t pipelineStages,
const SchedulingTarget& physicalTarget, const SchedulingTarget& physicalTarget,
std::string& error); std::string& error);
mlir::LogicalResult splitPipelineWorkload(const ComputeGraph& graph, enum class PipelineWorkloadPreparation {
const MergeScheduleResult& schedule, Ready,
size_t pipelineStages, Changed,
const SchedulingTarget& physicalTarget, };
std::string& error);
mlir::FailureOr<PipelineWorkloadPreparation> preparePipelineWorkload(
const ComputeGraph& graph, const MergeScheduleResult& schedule,
size_t pipelineStages, const SchedulingTarget& physicalTarget,
std::string& error);
} // namespace onnx_mlir::spatial } // namespace onnx_mlir::spatial
@@ -23,6 +23,7 @@ struct SchedulingTarget {
Cost transferWidthBytes = 8; Cost transferWidthBytes = 8;
Cost vectorWidth = 16; Cost vectorWidth = 16;
Cost vectorLatencyCycles = 4; Cost vectorLatencyCycles = 4;
size_t synchronizationRegisterCount = 0;
Cost matrixRows = 128; Cost matrixRows = 128;
Cost matrixColumns = 128; Cost matrixColumns = 128;
@@ -193,7 +193,7 @@ FailureOr<TopLevelOpInfo> buildTopLevelOpInfo(Operation& op, bool isScheduled, s
if constexpr (std::is_same_v<ComputeOpTy, SpatScheduledCompute>) { if constexpr (std::is_same_v<ComputeOpTy, SpatScheduledCompute>) {
if (auto compute = dyn_cast<ComputeOpTy>(&op)) { if (auto compute = dyn_cast<ComputeOpTy>(&op)) {
auto coreId = getOptionalScheduledCoreId(compute, "spatial dataflow export core id"); auto coreId = getOptionalScheduledCoreId(compute, "Spatial dataflow export core id");
if (failed(coreId)) if (failed(coreId))
return failure(); return failure();
if (*coreId) if (*coreId)
@@ -207,7 +207,7 @@ FailureOr<TopLevelOpInfo> buildTopLevelOpInfo(Operation& op, bool isScheduled, s
template <typename BatchOpTy> template <typename BatchOpTy>
FailureOr<SmallVector<int32_t, 8>> getBatchLaneCoreIds(BatchOpTy batch) { FailureOr<SmallVector<int32_t, 8>> getBatchLaneCoreIds(BatchOpTy batch) {
if constexpr (std::is_same_v<BatchOpTy, SpatScheduledComputeBatch>) { if constexpr (std::is_same_v<BatchOpTy, SpatScheduledComputeBatch>) {
auto coreIds = getOptionalScheduledBatchCoreIds(batch, "spatial dataflow export core ids"); auto coreIds = getOptionalScheduledBatchCoreIds(batch, "Spatial dataflow export core ids");
if (failed(coreIds)) if (failed(coreIds))
return failure(); return failure();
if (!*coreIds) if (!*coreIds)
+33 -4
View File
@@ -13,7 +13,7 @@ include "mlir/Interfaces/SideEffectInterfaces.td"
def SpatialDialect : Dialect { def SpatialDialect : Dialect {
let name = "spat"; let name = "spat";
let summary = "Dialect designed for deep learning computation in a spatial architecture"; let summary = "Dialect designed for deep learning computation in a Spatial architecture";
let cppNamespace = "::onnx_mlir::spatial"; let cppNamespace = "::onnx_mlir::spatial";
let useDefaultAttributePrinterParser = 0; let useDefaultAttributePrinterParser = 0;
let extraClassDeclaration = [{ let extraClassDeclaration = [{
@@ -592,13 +592,15 @@ def SpatHostStoreSyncOp : SpatOp<"host_store_sync", []> {
} }
def SpatHostWaitLoadOp : SpatOp<"host_wait_load", []> { def SpatHostWaitLoadOp : SpatOp<"host_wait_load", []> {
let summary = "Wait for a producer and load its tensor from host memory"; let summary = "Wait for producers, load from host memory, and acknowledge consumption";
let arguments = (ins let arguments = (ins
Index:$sourceCoreId, Index:$sourceCoreId,
Index:$targetCoreId, Index:$targetCoreId,
Index:$hostOffset, Index:$hostOffset,
Index:$eventRegister Index:$eventRegister,
Index:$waitValue,
Index:$acknowledgementEventRegister
); );
let results = (outs let results = (outs
@@ -607,7 +609,34 @@ def SpatHostWaitLoadOp : SpatOp<"host_wait_load", []> {
let assemblyFormat = [{ let assemblyFormat = [{
`from` $sourceCoreId `to` $targetCoreId `from` $sourceCoreId `to` $targetCoreId
`host_offset` $hostOffset `event` $eventRegister attr-dict `:` type($output) `host_offset` $hostOffset `event` $eventRegister `count` $waitValue
`ack` $acknowledgementEventRegister attr-dict `:` type($output)
}];
}
def SpatSyncOp : SpatOp<"sync", []> {
let summary = "Signal a synchronization register on another processor";
let arguments = (ins
Index:$targetCoreId,
Index:$eventRegister
);
let assemblyFormat = [{
$targetCoreId `event` $eventRegister attr-dict
}];
}
def SpatWaitOp : SpatOp<"wait", []> {
let summary = "Wait for a synchronization register value";
let arguments = (ins
Index:$eventRegister,
Index:$waitValue
);
let assemblyFormat = [{
$eventRegister `value` $waitValue attr-dict
}]; }];
} }
@@ -1,24 +0,0 @@
#ifndef SPATIAL_LAYOUT_INTERFACE_TD
#define SPATIAL_LAYOUT_INTERFACE_TD
include "mlir/IR/OpBase.td"
def SpatialLayoutCapabilityInterface : OpInterface<"SpatialLayoutCapabilityInterface"> {
let description = [{
Contract implemented by logical Spatial planning operations that expose
their legal physical layout alternatives to the Spatial planner.
}];
let methods = [
InterfaceMethod<
"Return legal physical layout alternatives for this operation and its current operand layouts.",
"::llvm::SmallVector<::onnx_mlir::spatial::LayoutAlternative>",
"getLayoutAlternatives",
(ins "const ::onnx_mlir::spatial::SpatialTargetInfo &":$target,
"::llvm::ArrayRef<::onnx_mlir::spatial::PhysicalLayout>":$operandLayouts)>
];
let cppNamespace = "::onnx_mlir::spatial";
}
#endif
@@ -1,37 +0,0 @@
#pragma once
#include <cstddef>
#include <cstdint>
namespace onnx_mlir::spatial {
struct MatrixUnitShape {
size_t rows = 128;
size_t columns = 128;
};
enum class ConvLoweringStrategy : uint8_t {
Auto,
Legacy,
Depthwise,
PackedIm2Col,
StreamedPatch,
StreamedPacked,
OutputChannelTiled,
InputKTiled,
Tiled2D,
};
struct SpatialTargetInfo {
MatrixUnitShape matrixShape;
size_t matrixUnitsPerProcessor = 64;
size_t processorCount = 1;
size_t vectorWidth = 16;
uint64_t convIm2colMaxElements = 1ull << 20;
uint64_t convStreamChunkPositions = 1024;
ConvLoweringStrategy convLoweringStrategy = ConvLoweringStrategy::Auto;
bool useExperimentalConvImplementation = false;
};
} // namespace onnx_mlir::spatial
-63
View File
@@ -1,63 +0,0 @@
#pragma once
#include "mlir/Pass/Pass.h"
#include <cstddef>
#include <memory>
#include <string>
namespace onnx_mlir {
namespace spatial {
struct SchedulingTarget;
struct ScheduledSpatialState;
struct SpatialTargetInfo;
std::unique_ptr<mlir::Pass> createScheduleSpatialGraphPass();
std::unique_ptr<mlir::Pass> createScheduleSpatialGraphPass(const SchedulingTarget& target);
std::unique_ptr<mlir::Pass> createScheduleSpatialGraphPass(
const SchedulingTarget& target,
std::shared_ptr<ScheduledSpatialState> state);
std::unique_ptr<mlir::Pass> createVerifyScheduledSpatialPass();
std::unique_ptr<mlir::Pass> createVerifyScheduledSpatialPass(
std::shared_ptr<ScheduledSpatialState> state);
std::unique_ptr<mlir::Pass> createRealizeSpatialCommunicationPass();
std::unique_ptr<mlir::Pass> createRealizeSpatialCommunicationPass(
const SchedulingTarget& target,
std::shared_ptr<ScheduledSpatialState> state);
std::unique_ptr<mlir::Pass> createVerifyRealizedSpatialPass();
std::unique_ptr<mlir::Pass> createVerifyRealizedSpatialPass(
std::shared_ptr<ScheduledSpatialState> state);
}
std::unique_ptr<mlir::Pass> createONNXToSpatialPass();
std::unique_ptr<mlir::Pass> createONNXToSpatialPass(const spatial::SpatialTargetInfo& target);
std::unique_ptr<mlir::Pass> createSpatialLayoutPlanningPass();
std::unique_ptr<mlir::Pass> createSpatialLayoutPlanningPass(const spatial::SpatialTargetInfo& target);
std::unique_ptr<mlir::Pass> createLowerSpatialPlansPass();
std::unique_ptr<mlir::Pass> createLowerSpatialPlansPass(const spatial::SpatialTargetInfo& target);
std::unique_ptr<mlir::Pass> createSpatialToPimPass();
std::unique_ptr<mlir::Pass> createPimBufferizationPreparationPass();
std::unique_ptr<mlir::Pass> createPimOneShotBufferizationPass();
std::unique_ptr<mlir::Pass> createPimMemoryNormalizationPass();
std::unique_ptr<mlir::Pass> createPimBufferizationVerificationPass();
std::unique_ptr<mlir::Pass> createTrivialGraphComputeMergePass();
std::unique_ptr<mlir::Pass> createTrivialGraphComputeMergePass(
size_t residentWeightCapacity);
std::unique_ptr<mlir::Pass> createPimHostConstantFoldingPass();
std::unique_ptr<mlir::Pass> createPimInstructionSelectionPass();
std::unique_ptr<mlir::Pass> createPimLocalMemoryPlanningPass();
std::unique_ptr<mlir::Pass> createPimVerificationPass();
std::unique_ptr<mlir::Pass> createEmitPimCodePass();
std::unique_ptr<mlir::Pass> createMessagePass(std::string message);
} // namespace onnx_mlir
+2 -1
View File
@@ -25,7 +25,8 @@ std::unique_ptr<mlir::Pass> createONNXToSpatialPass(
const spatial::SpatialTargetResources& target, const spatial::SpatialTargetResources& target,
const ONNXToSpatialPlanningOptions& options); const ONNXToSpatialPlanningOptions& options);
std::unique_ptr<mlir::Pass> createSpatialLayoutPlanningPass(); std::unique_ptr<mlir::Pass> createSpatialLayoutPlanningPass();
std::unique_ptr<mlir::Pass> createSpatialLayoutPlanningPass(const spatial::SpatialTargetResources& target); std::unique_ptr<mlir::Pass> createSpatialLayoutPlanningPass(
const spatial::SpatialTargetResources& target, bool selectTrivialPlan = false);
std::unique_ptr<mlir::Pass> createLowerSpatialPlansPass(); std::unique_ptr<mlir::Pass> createLowerSpatialPlansPass();
std::unique_ptr<mlir::Pass> createLowerSpatialPlansPass( std::unique_ptr<mlir::Pass> createLowerSpatialPlansPass(
const spatial::SpatialTargetResources& target, const spatial::SpatialTargetResources& target,
@@ -12,7 +12,7 @@ namespace {
struct EmitPimCodePass : PassWrapper<EmitPimCodePass, OperationPass<ModuleOp>> { struct EmitPimCodePass : PassWrapper<EmitPimCodePass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(EmitPimCodePass); MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(EmitPimCodePass);
StringRef getArgument() const override { return "emit-pim-code-pass"; } StringRef getArgument() const override { return "emit-pim-code-pass"; }
StringRef getDescription() const override { return "Emit PIM simulator code artifacts"; } StringRef getDescription() const override { return "Emit Pim simulator code artifacts"; }
EmitPimCodePass() {} EmitPimCodePass() {}
EmitPimCodePass(const EmitPimCodePass& pass) {} EmitPimCodePass(const EmitPimCodePass& pass) {}
@@ -25,7 +25,7 @@ struct EmitPimCodePass : PassWrapper<EmitPimCodePass, OperationPass<ModuleOp>> {
int compiler_error_code = compileToPimCode(moduleOp, pimDir); int compiler_error_code = compileToPimCode(moduleOp, pimDir);
if (compiler_error_code != CompilerSuccess) { if (compiler_error_code != CompilerSuccess) {
moduleOp.emitError() << "failed to emit PIM simulator code artifacts; compiler error code " moduleOp.emitError() << "failed to emit Pim simulator code artifacts; compiler error code "
<< compiler_error_code; << compiler_error_code;
signalPassFailure(); signalPassFailure();
} }
+97 -18
View File
@@ -9,6 +9,36 @@
using namespace onnx_mlir::spatial; using namespace onnx_mlir::spatial;
int main() { int main() {
PipelineCoreLayout unevenLayout(138, 4);
assert(unevenLayout.isValid());
assert(unevenLayout.getLogicalProcessorCount() == 34);
assert(unevenLayout.getStageRange(0).begin == 0);
assert(unevenLayout.getStageRange(0).size == 35);
assert(unevenLayout.getStageRange(1).begin == 35);
assert(unevenLayout.getStageRange(1).size == 35);
assert(unevenLayout.getStageRange(2).begin == 70);
assert(unevenLayout.getStageRange(2).size == 34);
assert(unevenLayout.getStageRange(3).begin == 104);
assert(unevenLayout.getStageRange(3).size == 34);
assert(unevenLayout.getStageForCore(34) == 0);
assert(unevenLayout.getStageForCore(35) == 1);
assert(unevenLayout.getStageForCore(69) == 1);
assert(unevenLayout.getStageForCore(70) == 2);
assert(unevenLayout.getStageForCore(137) == 3);
assert(!unevenLayout.getStageForCore(138));
PipelineCoreLayout dynamicLayout(std::vector<size_t> {2, 4, 1, 3});
assert(dynamicLayout.isValid());
assert(dynamicLayout.getProcessorCount() == 10);
assert(dynamicLayout.getStageRange(0).begin == 0);
assert(dynamicLayout.getStageRange(1).begin == 2);
assert(dynamicLayout.getStageRange(2).begin == 6);
assert(dynamicLayout.getStageRange(3).begin == 7);
assert(dynamicLayout.getStageForCore(1) == 0);
assert(dynamicLayout.getStageForCore(2) == 1);
assert(dynamicLayout.getStageForCore(6) == 2);
assert(dynamicLayout.getStageForCore(9) == 3);
TransferCost transfer {.fixed = 50, .networkFlits = 4}; TransferCost transfer {.fixed = 50, .networkFlits = 4};
SchedulingTarget fast; SchedulingTarget fast;
@@ -57,6 +87,13 @@ int main() {
}; };
assert(mapLogicalProcessorsToPhysicalCores(logicalTrafficFlits, alreadyPlaced) == std::vector<size_t>({0, 1, 2})); assert(mapLogicalProcessorsToPhysicalCores(logicalTrafficFlits, alreadyPlaced) == std::vector<size_t>({0, 1, 2}));
std::vector<size_t> placementGroups {0, 1, 1};
std::vector<size_t> groupedPlacement = mapLogicalProcessorsToPhysicalCores(
logicalTrafficFlits, line, placementGroups);
for (size_t processor = 0; processor < groupedPlacement.size(); ++processor)
assert(placementGroups[processor]
== placementGroups[groupedPlacement[processor]]);
ComputeGraph graph; ComputeGraph graph;
graph.successors.resize(6); graph.successors.resize(6);
graph.predecessors.resize(6); graph.predecessors.resize(6);
@@ -71,16 +108,16 @@ int main() {
graph.instanceToIndex[instance] = task; graph.instanceToIndex[instance] = task;
} }
MergeScheduleResult pipelineSchedule; MergeScheduleResult logicalSchedule;
pipelineSchedule.processorCount = 2; logicalSchedule.processorCount = 2;
pipelineSchedule.dominanceOrderCompute.reserve(graph.nodes.size()); logicalSchedule.dominanceOrderCompute.reserve(graph.nodes.size());
for (size_t task = 0; task < graph.nodes.size(); ++task) { for (size_t task = 0; task < graph.nodes.size(); ++task) {
const ComputeInstance& instance = graph.nodes[task].instance; const ComputeInstance& instance = graph.nodes[task].instance;
pipelineSchedule.dominanceOrderCompute.push_back(instance); logicalSchedule.dominanceOrderCompute.push_back(instance);
size_t cpu = task < 4 ? 0 : 1; size_t cpu = task < 4 ? 0 : 1;
pipelineSchedule.computeToCpuMap[instance] = cpu; logicalSchedule.computeToCpuMap[instance] = cpu;
pipelineSchedule.computeToCpuSlotMap[instance] = task < 4 ? task : task - 4; logicalSchedule.computeToCpuSlotMap[instance] = task < 4 ? task : task - 4;
pipelineSchedule.computeToAestMap[instance] = task; logicalSchedule.computeToAestMap[instance] = task;
} }
SchedulingTarget physical = fast; SchedulingTarget physical = fast;
@@ -93,20 +130,58 @@ int main() {
3, 3, 3, 0, 3, 3, 3, 0,
}; };
std::string pipelineError; std::string pipelineError;
ComputeGraph emptyGraph;
MergeScheduleResult emptySchedule;
emptySchedule.processorCount = 2;
assert(mlir::succeeded(applyPipelineScheduling(
emptyGraph, emptySchedule, 2, physical, pipelineError)));
assert(emptySchedule.processorCount == physical.processorCount);
assert(emptySchedule.processorStages == std::vector<size_t>({0, 0, 1, 1}));
MergeScheduleResult pipelineSchedule = logicalSchedule;
assert(mlir::succeeded(applyPipelineScheduling( assert(mlir::succeeded(applyPipelineScheduling(
graph, pipelineSchedule, 2, physical, pipelineError))); graph, pipelineSchedule, 2, physical, pipelineError)));
assert(pipelineSchedule.processorCount == 4); assert(pipelineSchedule.processorCount == 4);
assert(pipelineSchedule.computeToCpuMap.lookup(graph.nodes[0].instance) == 0); size_t predecessorCore = pipelineSchedule.computeToCpuMap.lookup(
assert(pipelineSchedule.computeToCpuMap.lookup(graph.nodes[1].instance) == 0); graph.nodes[1].instance);
assert(pipelineSchedule.computeToCpuMap.lookup(graph.nodes[2].instance) == 2); size_t successorCore = pipelineSchedule.computeToCpuMap.lookup(
assert(pipelineSchedule.computeToCpuMap.lookup(graph.nodes[3].instance) == 2); graph.nodes[2].instance);
assert(pipelineSchedule.computeToCpuMap.lookup(graph.nodes[4].instance) == 1);
assert(pipelineSchedule.computeToCpuMap.lookup(graph.nodes[5].instance) == 3);
assert(pipelineSchedule.computeToAestMap.lookup(graph.nodes[2].instance) assert(pipelineSchedule.computeToAestMap.lookup(graph.nodes[2].instance)
>= pipelineSchedule.computeToAestMap.lookup(graph.nodes[1].instance) >= pipelineSchedule.computeToAestMap.lookup(graph.nodes[1].instance)
+ graph.nodes[1].cost + 4); + graph.nodes[1].cost
+ getPeftTransferTime(
TransferCost {.fixed = 1, .networkFlits = 1},
predecessorCore, successorCore, physical));
assert(pipelineSchedule.processorStages[predecessorCore]
<= pipelineSchedule.processorStages[successorCore]);
assert(pipelineSchedule.processorStages[successorCore]
<= pipelineSchedule.processorStages[predecessorCore] + 1);
assert(pipelineSchedule.equivalentClass.empty()); assert(pipelineSchedule.equivalentClass.empty());
graph.successors[0].push_back(
{5, TransferCost {.fixed = 1, .networkFlits = 1}});
graph.predecessors[5].push_back(
{0, TransferCost {.fixed = 1, .networkFlits = 1}});
SchedulingTarget fourStagePhysical = physical;
fourStagePhysical.processorCount = 8;
fourStagePhysical.interProcessorLatencyNs.assign(64, 3);
for (size_t core = 0; core < 8; ++core)
fourStagePhysical.interProcessorLatencyNs[core * 8 + core] = 0;
MergeScheduleResult fourStageSchedule = logicalSchedule;
assert(mlir::succeeded(applyPipelineScheduling(
graph, fourStageSchedule, 4, fourStagePhysical, pipelineError)));
for (size_t task = 0; task < graph.nodes.size(); ++task)
for (const auto &[predecessor, cost] : graph.predecessors[task]) {
(void)cost;
size_t sourceStage = fourStageSchedule.processorStages[
fourStageSchedule.computeToCpuMap.lookup(
graph.nodes[predecessor].instance)];
size_t targetStage = fourStageSchedule.processorStages[
fourStageSchedule.computeToCpuMap.lookup(graph.nodes[task].instance)];
assert(sourceStage <= targetStage);
assert(targetStage <= sourceStage + 1);
}
ComputeGraph communicationGraph; ComputeGraph communicationGraph;
communicationGraph.successors.resize(5); communicationGraph.successors.resize(5);
communicationGraph.predecessors.resize(5); communicationGraph.predecessors.resize(5);
@@ -136,15 +211,19 @@ int main() {
MergeScheduleResult fastCommunicationSchedule = communicationSchedule; MergeScheduleResult fastCommunicationSchedule = communicationSchedule;
assert(mlir::succeeded(applyPipelineScheduling( assert(mlir::succeeded(applyPipelineScheduling(
communicationGraph, fastCommunicationSchedule, 2, fastPipeline, pipelineError))); communicationGraph, fastCommunicationSchedule, 2, fastPipeline, pipelineError)));
assert(fastCommunicationSchedule.computeToCpuMap.lookup(
communicationGraph.nodes[2].instance) == 2);
SchedulingTarget slowPipeline = fastPipeline; SchedulingTarget slowPipeline = fastPipeline;
slowPipeline.averageInterProcessorLatencyNs = 10; slowPipeline.averageInterProcessorLatencyNs = 10;
MergeScheduleResult slowCommunicationSchedule = communicationSchedule; MergeScheduleResult slowCommunicationSchedule = communicationSchedule;
assert(mlir::succeeded(applyPipelineScheduling( assert(mlir::succeeded(applyPipelineScheduling(
communicationGraph, slowCommunicationSchedule, 2, slowPipeline, pipelineError))); communicationGraph, slowCommunicationSchedule, 2, slowPipeline, pipelineError)));
assert(slowCommunicationSchedule.computeToCpuMap.lookup( size_t sourceCore = slowCommunicationSchedule.computeToCpuMap.lookup(
communicationGraph.nodes[2].instance) < 2); communicationGraph.nodes[4].instance);
size_t targetCore = slowCommunicationSchedule.computeToCpuMap.lookup(
communicationGraph.nodes[3].instance);
assert(slowCommunicationSchedule.processorStages[sourceCore]
<= slowCommunicationSchedule.processorStages[targetCore]);
assert(slowCommunicationSchedule.processorStages[targetCore]
<= slowCommunicationSchedule.processorStages[sourceCore] + 1);
return EXIT_SUCCESS; return EXIT_SUCCESS;
} }
+1 -1
View File
@@ -238,7 +238,7 @@ def print_report(path: Path, counts: Counter, groups: dict[tuple[str, str], Chai
def main() -> None: def main() -> None:
parser = argparse.ArgumentParser(description="Analyze repeated Spatial/PIM tensor IR cardinality patterns.") parser = argparse.ArgumentParser(description="Analyze repeated Spatial/Pim tensor IR cardinality patterns.")
parser.add_argument("paths", nargs="+", help="MLIR files to analyze.") parser.add_argument("paths", nargs="+", help="MLIR files to analyze.")
parser.add_argument("--limit", type=int, default=12, help="Maximum number of hot chains to print per file.") parser.add_argument("--limit", type=int, default=12, help="Maximum number of hot chains to print per file.")
args = parser.parse_args() args = parser.parse_args()
@@ -0,0 +1,3 @@
from .cli import main
raise SystemExit(main())
+4 -16
View File
@@ -1,24 +1,12 @@
operations/**/inputs operations/**/artifacts
operations/**/outputs
operations/**/raptor
operations/**/runner
operations/**/simulation
operations/**/*.csv operations/**/*.csv
!operations/validation_results.csv !operations/validation_results.csv
networks/**/inputs networks/**/artifacts
networks/**/outputs
networks/**/raptor
networks/**/pimcomp
networks/**/runner
networks/**/simulation
networks/**/real_image_val
networks/**/*.png networks/**/*.png
networks/**/*.jpg networks/**/*.jpg
networks/**/*.csv networks/**/*.csv
!networks/validation_results.csv !networks/validation_results.csv
!networks/full_net/validation_results.csv !networks/full_net/validation_results.csv
!networks/pimcomp_models/validation_results.csv !networks/pimcomp_models/results_comparison.csv
!networks/pimcomp_models/results.csv !networks/pimcomp_models/results_ablation.csv
!networks/pimcomp_models/validation_results.csv
!operations/validation_results.csv
+66 -52
View File
@@ -1,14 +1,14 @@
# Raptor Validation # Raptor validation
`validate.py` validates every ONNX model below a selected directory. For each `validate.py` validates every ONNX model below a selected directory. For each
model it can: model it can:
1. compile an ONNX-MLIR reference library and runner; 1. compile an ONNX-MLIR reference library and runner;
2. generate deterministic random inputs; 2. generate deterministic random inputs;
3. compile PIM artifacts with Raptor; 3. compile Pim artifacts with Raptor;
4. run the reference implementation and functional PIM simulator; 4. run the reference implementation and functional Pim simulator;
5. compare their outputs; 5. compare their outputs;
6. run `pimsim-nn` to report latency, power, and energy. 6. run Pimsim to report latency, throughput, power, and energy.
Run the script from the repository root with the repository Python environment. Run the script from the repository root with the repository Python environment.
@@ -59,34 +59,23 @@ Validate a network or network slice:
`--operations-dir` may point to any directory tree containing `.onnx` files. `--operations-dir` may point to any directory tree containing `.onnx` files.
The script discovers them recursively and writes `validation_results.csv` in The script discovers them recursively and writes `validation_results.csv` in
that directory while retaining the terminal table. that directory while retaining separate latency and throughput terminal tables.
## Raptor vs PIMCOMP comparison ## Pim validation tools
The PIMCOMP paper-model suite has a one-command Arch-A comparison: - [Pimcomp model suite](networks/pimcomp_models/README.md)
- [Pimcomp model comparison tools](tools/pim/pimcomp/compare/README.md)
```bash - [Pimcomp correctness study](tools/pim/pimcomp/correctness/README.md)
.venv/bin/python validation/tools/pimcomp/run_pimcomp_paper_latency.py - [Raptor compiler ablation study](tools/pim/ablation/README.md)
```
The runner verifies PIMCOMP's population-200, 1000-iteration GA settings,
builds Raptor and the existing `third_party/PIMCOMP-NN/build` tree, and compares
the four paper models one at a time. Each PIMCOMP GA run evaluates candidates
in parallel; set `OMP_NUM_THREADS` to control its worker count. Use
`--models vgg8` for one model or `--dry-run` to print the commands.
Generated artifacts are stored beside each model and ignored by Git. The
comparison reuses the same model-level `inputs/`, `outputs/`, `runner/`,
`raptor/`, and `simulation/` paths as regular validation. PIMCOMP-only
artifacts and `comparison_report.{md,json}` live under `pimcomp/`. See
[`networks/pimcomp_models/README.md`](networks/pimcomp_models/README.md) for
profiles, model provenance, limitations, and remote execution.
## Validation modes ## Validation modes
The default mode performs the complete workflow. The default mode runs latency and throughput in one validation job. Latency
uses one input, while throughput uses `--pipeline=4` with four distinct inputs.
Both modes reuse the generated input batch, native runner, and reference
outputs, and every throughput output is compared with its own reference.
Use `--compile-only` to build the reference runner and PIM artifacts without Use `--compile-only` to build the reference runner and Pim artifacts without
executing either implementation: executing either implementation:
```bash ```bash
@@ -133,22 +122,23 @@ count with `-j` or `--jobs`:
| `--onnx-include-dir PATH` | ONNX-MLIR runtime include directory. Required unless `--clean` is used. | | `--onnx-include-dir PATH` | ONNX-MLIR runtime include directory. Required unless `--clean` is used. |
| `--operations-dir PATH` | Directory tree containing models. Defaults to `validation/operations`. | | `--operations-dir PATH` | Directory tree containing models. Defaults to `validation/operations`. |
| `--simulator-dir PATH` | Functional `pim-simulator` crate directory. Defaults to the in-tree simulator. | | `--simulator-dir PATH` | Functional `pim-simulator` crate directory. Defaults to the in-tree simulator. |
| `--non-functional-simulator-build-dir PATH` | `pimsim-nn` build directory. Defaults to the in-tree build. | | `--non-functional-simulator-build-dir PATH` | Pimsim build directory. Defaults to the in-tree build. |
| `--pimcomp-config {arch-a,arch-b,arch-c}` | Non-functional hardware/timing profile. Defaults to `arch-a`. | | `--pimcomp-config {arch-a,arch-b,arch-c}` | Non-functional hardware/timing profile. Defaults to `arch-a`. |
| `--skip-non-functional-simulation` | Skip `pimsim-nn` latency, power, and energy measurement. | | `--skip-non-functional-simulation` | Skip Pimsim latency, throughput, power, and energy measurement. |
| `--no-fast` | Disable fast throughput convergence for authoritative full-duration Pimsim experiments. |
| `--threshold FLOAT` | Absolute output-comparison tolerance. Defaults to `1e-3`. | | `--threshold FLOAT` | Absolute output-comparison tolerance. Defaults to `1e-3`. |
| `--relative-threshold FLOAT` | Relative output-comparison tolerance. Defaults to `1e-5`. | | `--relative-threshold FLOAT` | Relative output-comparison tolerance. Defaults to `1e-5`. |
| `--seed INT` | Seed for generated inputs. Defaults to `0`. | | `--seed INT` | Seed for generated inputs. Defaults to `0`. |
| `--crossbar-size INT` | Crossbar dimensions passed to Raptor. Defaults to the Arch-A value, `128`. | | `--crossbar-size INT` | Crossbar dimensions passed to Raptor. Defaults to the Arch-A value, `128`. |
| `--crossbar-count INT` | Crossbars per core passed to Raptor. Defaults to the Arch-A value, `96`. | | `--crossbar-count INT` | Crossbars per core passed to Raptor. Defaults to the Arch-A value, `96`. |
| `--core-count INT` | PIM core count passed to Raptor. Defaults to the Arch-A value, `168`. | | `--core-count INT` | Pim core count passed to Raptor. Defaults to the Arch-A value, `168`. |
| `--raptor-extra-arg=ARG` | Additional Raptor compiler argument. Repeat for multiple arguments. | | `--raptor-extra-arg=ARG` | Additional Raptor compiler argument. Repeat for multiple arguments. |
| `--command-timeout-seconds FLOAT` | Timeout for each compiler, runner, and simulator subprocess. Defaults to `1000000.0`. | | `--command-timeout-seconds FLOAT` | Timeout for each compiler, runner, and simulator subprocess. Defaults to `1000000.0`. |
| `-j INT`, `--jobs INT` | Parallel validation workers. Defaults to all available CPUs and must be at least one. | | `-j INT`, `--jobs INT` | Parallel validation workers. Defaults to all available CPUs and must be at least one. |
| `--clean` | Remove generated validation artifacts and exit. | | `--clean` | Remove generated validation artifacts and exit. |
| `--compile-only` | Compile reference and PIM artifacts without execution or comparison. | | `--compile-only` | Compile reference and Pim artifacts without execution or comparison. |
| `--run-only` | Reuse compiled artifacts and perform execution, simulation, and comparison. | | `--run-only` | Reuse compiled artifacts and perform execution, simulation, and comparison. |
| `--verbose` | Print passing per-stage and subprocess logs, plus average PIM pass timings. | | `--verbose` | Print passing per-stage and subprocess logs, plus average Pim pass timings. |
Arguments beginning with `--` that are passed through to Raptor should use the Arguments beginning with `--` that are passed through to Raptor should use the
equals form: equals form:
@@ -159,40 +149,55 @@ equals form:
## Hardware profiles and non-functional simulation ## Hardware profiles and non-functional simulation
The selected PIMCOMP profile must match `--core-count`, `--crossbar-count`, and The selected Pimcomp profile must match `--core-count`, `--crossbar-count`, and
`--crossbar-size`. A mismatch disables only non-functional simulation and `--crossbar-size`. A mismatch disables only non-functional simulation and
prints the incompatible values; functional validation still runs. prints the incompatible values; functional validation still runs.
The checked-in profiles are under The checked-in profiles are under
`validation/pimsim_configs/pimcomp/<profile>/latency_config.json`. `validation/pimsim_configs/pimcomp/<profile>/`; latency uses
`latency_config.json`, while throughput uses
`throughput_config_<time>ms.json` and the mesh beside it.
Use `--skip-non-functional-simulation` when latency, power, and energy are not required. Throughput measurement defaults to `pimsim-nn --fast` with a 1000 ms
convergence deadline. Fast mode compares consecutive two-round windows with a
fixed 1% tolerance and falls back to the legacy full-duration result if it
does not converge. Use `--no-fast` for authoritative experiments.
Use `--skip-non-functional-simulation` when latency, throughput, power, and energy are not required.
The summary reports non-functional results as measured, failed, unsupported, or The summary reports non-functional results as measured, failed, unsupported, or
skipped. skipped.
Overall PASS/FAIL is determined by compilation and functional output Overall PASS/FAIL is determined by compilation and functional output
comparison. A non-functional simulation failure remains visible as `ERROR` in comparison. A non-functional simulation failure remains visible as `ERROR` in
the latency, power, and energy columns but does not change a functional PASS. the corresponding latency, throughput, power, or energy columns but does not
change a functional PASS.
`pimsim-nn` does not currently implement the `vsoftmax` instruction. When its Pimsim does not currently implement the `vsoftmax` instruction. When its
explicit unsupported-op diagnostic is encountered, Softmax validations retain explicit unsupported-op diagnostic is encountered, Softmax validations retain
their functional PASS and show `UNSUPPORTED` in the non-functional columns. their functional PASS and show `UNSUPPORTED` in the non-functional columns.
Other `pimsim-nn` failures remain `ERROR`. Other Pimsim failures remain `ERROR`.
## Generated artifacts ## Generated artifacts
Artifacts are written beside each model: Generated files are grouped below an `artifacts/` directory beside each model
or operation case. This keeps checked-in ONNX files and generated trees
separate and lets `--clean` remove the complete workspace, including stale
validation lock files:
| Path | Contents | | Path | Contents |
|---|---| |---|---|
| `inputs/` | Generated input CSV files. | | `artifacts/inputs.csv` | Generated inputs, one batch entry per line. |
| `outputs/` | ONNX-MLIR reference output CSV files. | | `artifacts/inputs/`, `artifacts/outputs/`, `artifacts/runner/` | Inputs, reference outputs, and the runner shared by latency and throughput validation. |
| `raptor/` | Exported MLIR, dialect snapshots, reports, final FP32 `pim/` artifacts, and the int8-equivalent `pimsim_nn/` latency view. | | `artifacts/raptor/pim/`, `artifacts/simulation/latency/` | Latency Pim artifacts and functional simulator outputs. |
| `runner/` | Generated reference runner source, build tree, and shared library. | | `artifacts/raptor/throughput/pim/`, `artifacts/simulation/throughput/` | Pipeline-4, batch-4 throughput Pim artifacts and functional simulator outputs. |
| `simulation/` | Functional simulator outputs used for numerical comparison. | | `artifacts/common/inputs/` | Shared generated input CSV files. |
| `pimcomp/` | PIMCOMP graph, instruction, simulator, and comparison-report artifacts. | | `artifacts/common/outputs/` | Shared ONNX-MLIR reference output CSV files. |
| `artifacts/common/runner/` | Shared reference runner source, build tree, and library. |
| `artifacts/<arch>/<mode>[/pipelineN]/raptor/` | Architecture- and pipeline-specific Raptor MLIR and Pim artifacts. |
| `artifacts/<arch>/<mode>[/pipelineN]/simulation/` | Functional simulator outputs for that comparison. |
| `artifacts/<arch>/<mode>[/pipelineN]/pimcomp/` | Pimcomp graph, instruction, simulator, and comparison-report artifacts. |
The `raptor/` directory may include `spatial0.mlir`, Each comparison's `raptor/` directory may include `spatial0.mlir`,
`spatial1_graph.mlir`, `spatial2_trivial_merged.mlir`, `spatial1_graph.mlir`, `spatial2_trivial_merged.mlir`,
`spatial3_scheduled_no_comm.mlir`, `spatial4_scheduled.mlir`, `pim0.mlir`, `spatial3_scheduled_no_comm.mlir`, `spatial4_scheduled.mlir`, `pim0.mlir`,
`pim1_buff.mlir`, `pim2_folded.mlir`, and `pim3_memory_planned.mlir`. `pim1_buff.mlir`, `pim2_folded.mlir`, and `pim3_memory_planned.mlir`.
@@ -221,27 +226,36 @@ The generated operation inventory is documented in
## Manual functional simulator tracing ## Manual functional simulator tracing
After validation has produced a `raptor/pim/` directory, rerun the functional After validation has produced an `artifacts/raptor/pim/` directory, rerun the functional
simulator with tracing from its crate directory: simulator with tracing from its crate directory:
```bash ```bash
cd backend-simulators/pim/pim-simulator cd backend-simulators/pim/pim-simulator
cargo run --no-default-features --features tracing --release \ cargo run --no-default-features --features tracing --release \
--package pim-simulator --bin pim-simulator -- \ --package pim-simulator --bin pim-simulator -- \
-f /path/to/workspace/raptor/pim \ -f /path/to/workspace/artifacts/raptor/pim \
-o /path/to/workspace/simulation/out.bin \ -o /path/to/workspace/artifacts/simulation/out.bin \
-d <addr0>,<size0>,<addr1>,<size1>,... -d <addr0>,<size0>,<addr1>,<size1>,... \
--mode latency \
--batch-size 1 \
--input-dir /path/to/workspace/artifacts/simulation/inputs
``` ```
Throughput mode additionally requires `--batch-size N` and at least `N`
`input_<index>.bin` files in `--input-dir`. Each input binary concatenates the model tensors in graph
input order. The comparison validator also writes one native reference and one
`simulation/*_iterations/output_*.bin` dump per batch entry, and checks every
entry rather than only the final output.
Tracing writes `TraceCore0`, `TraceCore1`, and so on beside `out.bin`. The Tracing writes `TraceCore0`, `TraceCore1`, and so on beside `out.bin`. The
validator normally derives the `-d` address and byte ranges from validator normally derives the `-d` address and byte ranges from
`raptor/pim/config.json` and the model output shapes. `raptor/pim/config.json` and the model output shapes.
## Results and exit status ## Results and exit status
The final table reports functional pass/fail state and non-functional latency The final table reports latency and throughput functional pass/fail state plus
and power. The summary includes pass/fail totals, non-functional simulation non-functional latency, throughput, power, and energy. The summary includes pass/fail totals, non-functional simulation
counts, total measured latency, and average PIM pass timings when `--verbose` counts, total measured latency, and average Pim pass timings when `--verbose`
is enabled. is enabled.
- Exit status `0`: all discovered models passed, or cleanup completed. - Exit status `0`: all discovered models passed, or cleanup completed.
+100 -68
View File
@@ -1,34 +1,38 @@
# PIMCOMP comparison models # Pimcomp comparison models
This directory contains the four networks evaluated in This directory contains the four networks evaluated in
[PIMCOMP: An End-to-End DNN Compiler for Processing-In-Memory Accelerators](https://arxiv.org/pdf/2411.09159): [Pimcomp: An End-to-End DNN Compiler for Processing-In-Memory Accelerators](https://arxiv.org/pdf/2411.09159):
VGG-8, ResNet-18, ResNet-34, and GoogLeNet. It also contains YOLO11n as an VGG-8, ResNet-18, ResNet-34, and GoogLeNet. It also contains YOLO11n as an
additional compiler comparison model. additional compiler comparison model.
See the runner-generated [results.csv](results.csv) for the current latency See the runner-generated [results_comparison.csv](results_comparison.csv) for the current comparison
and energy results. results. Rows are retained separately for each model, architecture, mode, and
pipeline. It records separate `PASS`/`FAIL` functional-validation fields for
the Raptor and Pimcomp artifacts; rows without a generated report contain `NA`.
Use `--archs` to select the architecture rows to generate; existing rows for
other architectures remain unchanged.
## Models and provenance ## Models and provenance
| Directory | Model | Input | Provenance | | Directory | Model | Input | Provenance |
|--------------|----------------------|---------------|------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------| |--------------|----------------------|---------------|------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------|
| `resnet18/` | ResNet-18 v1 | `1x3x224x224` | Symlink to the complete [ONNX Model Zoo `resnet18-v1-7`](https://huggingface.co/onnxmodelzoo/resnet18-v1-7) model already present at `../resnet18/depth_68/resnet18_depth_68.onnx`. | | `resnet18/` | ResNet-18 v1 | `1x3x224x224` | Symlink to the complete [ONNX Model Zoo `resnet18-v1-7`](https://huggingface.co/onnxmodelzoo/resnet18-v1-7) model already present at `../resnet18/depth_68/resnet18_depth_68.onnx`. |
| `resnet34/` | ResNet-34 v1 | `1x3x224x224` | [ONNX Model Zoo `resnet34-v1-7`](https://huggingface.co/onnxmodelzoo/resnet34-v1-7), with its symbolic batch fixed to 1 as PIMCOMP's frontend does. | | `resnet34/` | ResNet-34 v1 | `1x3x224x224` | [ONNX Model Zoo `resnet34-v1-7`](https://huggingface.co/onnxmodelzoo/resnet34-v1-7), with its symbolic batch fixed to 1 as Pimcomp's frontend does. |
| `googlenet/` | GoogLeNet | `1x3x224x224` | Unmodified [ONNX Model Zoo `googlenet-12`](https://huggingface.co/onnxmodelzoo/googlenet-12). | | `googlenet/` | GoogLeNet | `1x3x224x224` | Unmodified [ONNX Model Zoo `googlenet-12`](https://huggingface.co/onnxmodelzoo/googlenet-12). |
| `vgg8/` | VGG-8 reconstruction | `1x1x28x28` | Reconstruction of the [PIMCOMP VGG-8 benchmark](https://arxiv.org/html/2411.09159#S8.SS1), with six convolution and two fully connected layers. | | `vgg8/` | VGG-8 reconstruction | `1x1x28x28` | Reconstruction of the [Pimcomp VGG-8 benchmark](https://arxiv.org/html/2411.09159#S8.SS1), with six convolution and two fully connected layers. |
| `yolo11n/` | YOLO11n detection | `1x3x640x640` | Derived from the canonical local model at `../yolo11n/depth_51/yolo11n_depth_51.onnx`, exported from [Ultralytics YOLO11n](https://github.com/ultralytics/ultralytics/blob/main/docs/en/models/yolo11.md). | | `yolo11n/` | YOLO11n detection | `1x3x640x640` | Derived from the canonical local model at `../yolo11n/depth_51/yolo11n_depth_51.onnx`, exported from [Ultralytics YOLO11n](https://github.com/ultralytics/ultralytics/blob/main/docs/en/models/yolo11.md). |
`googlenet/googlenet-12-latency.onnx` is the explicit pimsim-nn-ready GoogLeNet model. `googlenet/googlenet-12-pimsim-nn.onnx` is the explicit Pimsim-ready GoogLeNet model.
It removes the two LRN nodes and terminal Softmax from the original model, It removes the two LRN nodes and terminal Softmax from the original model,
so that the comparison covers only operations scheduled by PIMCOMP and supported by pimsim-nn. so that the comparison covers only operations scheduled by Pimcomp and supported by Pimsim.
`yolo11n/yolo11n-latency.onnx` is the explicit pimsim-nn-ready YOLO11n model. `yolo11n/yolo11n-pimsim-nn.onnx` is the explicit Pimsim-ready YOLO11n model.
It removes the Softmax nodes from the original model, It removes the Softmax nodes from the original model,
so that the compiled artifact can be simulated in pimsim-nn. so that the compiled artifact can be simulated in Pimsim.
## Unsupported and ignored operations ## Unsupported and ignored operations
PIMCOMP's frontend accepts exactly these ONNX operations: Pimcomp's frontend accepts exactly these ONNX operations:
```text ```text
Add, AveragePool, BatchNormalization, Clip, Concat, Conv, Dropout, Flatten, Add, AveragePool, BatchNormalization, Clip, Concat, Conv, Dropout, Flatten,
@@ -36,14 +40,14 @@ Gather, Gemm, GlobalAveragePool, LRN, MatMul, MaxPool, Mul, Pad, Relu, Reshape,
Shape, Sigmoid, Softmax, Squeeze, Sub, Sum, Tanh, Transpose, Unsqueeze Shape, Sigmoid, Softmax, Squeeze, Sub, Sum, Tanh, Transpose, Unsqueeze
``` ```
`Constant` is consumed as frontend metadata rather than emitted as a PIMCOMP `Constant` is consumed as frontend metadata rather than emitted as a Pimcomp
node. Every other ONNX operation is unsupported: the frontend prints node. Every other ONNX operation is unsupported: the frontend prints
`operation: <type> not considered` and stops at the first occurrence. Thus the `operation: <type> not considered` and stops at the first occurrence. Thus the
complete unsupported set is the complement of the allowlist above for the complete unsupported set is the complement of the allowlist above for the
model's ONNX opset. In particular, YOLO11n contains unsupported `Split` and model's ONNX opset. In particular, YOLO11n contains unsupported `Split` and
`Resize` nodes. `Resize` nodes.
PIMCOMP's low-latency scheduler, hierarchy mapper, and genetic algorithm use Pimcomp's low-latency scheduler, hierarchy mapper, and genetic algorithm use
this complete explicit no-consider set: this complete explicit no-consider set:
```text ```text
@@ -59,7 +63,7 @@ specific Shape-Gather-Unsqueeze-Concat shape chain, and merge Pad into its
consumer. These transformations do not make an otherwise standalone ignored consumer. These transformations do not make an otherwise standalone ignored
operation timed. operation timed.
`pimsim-nn` consumes PIM ISA instructions. It supports every named opcode Pimsim consumes Pim ISA instructions. It supports every named opcode
in the shared serialized range except `vsoftmax` (opcode 21), in the shared serialized range except `vsoftmax` (opcode 21),
which is rejected explicitly in both JSON and binary input. It which is rejected explicitly in both JSON and binary input. It
silently ignores no opcode; unknown names and numbers are errors. silently ignores no opcode; unknown names and numbers are errors.
@@ -67,12 +71,12 @@ silently ignores no opcode; unknown names and numbers are errors.
These boundaries explain the dedicated artifacts: These boundaries explain the dedicated artifacts:
- GoogLeNet's two LRN nodes and terminal Softmax perform real computation but - GoogLeNet's two LRN nodes and terminal Softmax perform real computation but
are ignored by PIMCOMP, so the common latency artifact removes them. Its are ignored by Pimcomp, so the common latency artifact removes them. Its
inference Dropout and shape-only Reshape can remain without adding compute. inference Dropout and shape-only Reshape can remain without adding compute.
- YOLO11n's latency artifact bypasses exactly its two Softmax nodes so it can - YOLO11n's latency artifact bypasses exactly its two Softmax nodes so it can
run in `pimsim-nn`. Every other node, including MatMul, Transpose, and the run in Pimsim. Every other node, including MatMul, Transpose, and the
final detection-decoding tail, remains present and timed by Raptor. No final detection-decoding tail, remains present and timed by Raptor. No
PIMCOMP latency is reported because its frontend stops at `Split` and also Pimcomp latency is reported because its frontend stops at `Split` and also
lacks `Resize`; compiling that prefix would not represent YOLO11n. lacks `Resize`; compiling that prefix would not represent YOLO11n.
The authoritative lists are in The authoritative lists are in
@@ -81,16 +85,16 @@ The authoritative lists are in
[`ISA.h`](../../../backend-simulators/pim/pimsim-nn/src/isa/ISA.h), and [`ISA.h`](../../../backend-simulators/pim/pimsim-nn/src/isa/ISA.h), and
[`Instruction.cpp`](../../../backend-simulators/pim/pimsim-nn/src/isa/Instruction.cpp). [`Instruction.cpp`](../../../backend-simulators/pim/pimsim-nn/src/isa/Instruction.cpp).
The PIMCOMP authors did not publish the ONNX checkpoints used by the paper. The Pimcomp authors did not publish the ONNX checkpoints used by the paper.
Running PIMCOMP's frontend on the three Model Zoo files above produces JSON Running Pimcomp's frontend on the three Model Zoo files above produces JSON
graphs exactly equal to PIMCOMP-NN's bundled `resnet18.json`, `resnet34.json`, graphs exactly equal to Pimcomp's bundled `resnet18.json`, `resnet34.json`,
and `googlenet.json`. and `googlenet.json`.
There is no VGG-8 artifact in the ONNX Model Zoo or any PIMCOMP-NN revision. There is no VGG-8 artifact in the ONNX Model Zoo or any Pimcomp revision.
The included VGG-8 therefore has deterministic random weights and is suitable The included VGG-8 therefore has deterministic random weights and is suitable
for compiler and simulator comparison, not paper-accuracy reproduction. The for compiler and simulator comparison, not paper-accuracy reproduction. The
paper also says that VGG-8 and ResNet-18 were trained on MNIST, while the paper also says that VGG-8 and ResNet-18 were trained on MNIST, while the
published PIMCOMP graphs and ResNet Model Zoo artifacts use ImageNet shapes. published Pimcomp graphs and ResNet Model Zoo artifacts use ImageNet shapes.
Current SHA-256 checksums: Current SHA-256 checksums:
@@ -98,9 +102,9 @@ Current SHA-256 checksums:
788088b908e233d924c7c26b997e89ee861290c7bc56783a306e8201d79aac8f resnet18/resnet18-v1-7.onnx 788088b908e233d924c7c26b997e89ee861290c7bc56783a306e8201d79aac8f resnet18/resnet18-v1-7.onnx
c3231061d081bdd47884137b02134f85142752a39e87263c529cd14ed242b096 resnet34/resnet34-v1-7.onnx c3231061d081bdd47884137b02134f85142752a39e87263c529cd14ed242b096 resnet34/resnet34-v1-7.onnx
c99c507058eaf41de8723408fdda7db8325cb57f0a89f2ee07a716d6e963e14e googlenet/googlenet-12.onnx c99c507058eaf41de8723408fdda7db8325cb57f0a89f2ee07a716d6e963e14e googlenet/googlenet-12.onnx
a26f9e33901c573e60c34a3f0abbb4744fff83e4e0f21b18fc66e20395e72982 googlenet/googlenet-12-latency.onnx a26f9e33901c573e60c34a3f0abbb4744fff83e4e0f21b18fc66e20395e72982 googlenet/googlenet-12-pimsim-nn.onnx
396cdea21e5e7d02c3f26f14d22ef20975171702493f5c5e79b8e0d896e541ef vgg8/vgg8-mnist-reconstructed.onnx 396cdea21e5e7d02c3f26f14d22ef20975171702493f5c5e79b8e0d896e541ef vgg8/vgg8-mnist-reconstructed.onnx
229f3975af8933d39aee8d9031d969bff074b69c33e304a78abb35ff0c5f445f yolo11n/yolo11n-latency.onnx 229f3975af8933d39aee8d9031d969bff074b69c33e304a78abb35ff0c5f445f yolo11n/yolo11n-pimsim-nn.onnx
``` ```
## Paper hardware profiles ## Paper hardware profiles
@@ -109,10 +113,10 @@ The files in
[`../../pimsim_configs/pimcomp/`](../../pimsim_configs/pimcomp/) [`../../pimsim_configs/pimcomp/`](../../pimsim_configs/pimcomp/)
encode Table V's explicit resource parameters. encode Table V's explicit resource parameters.
Each profile subdirectory contains pre-generated latency and throughput Each profile subdirectory contains pre-generated latency and throughput
`pimsim-nn` configs plus its matching mesh; comparison and validation select Pimsim configs plus its matching mesh; comparison and validation reference
these checked-in artifacts without generating configs at runtime. these canonical artifacts directly.
| Config | Cores | Crossbars/core | Crossbar | Cell | PIMCOMP layout | | Config | Cores | Crossbars/core | Crossbar | Cell | Pimcomp layout |
|------------------------------|------------------:|---------------:|------------|------:|-----------------| |------------------------------|------------------:|---------------:|------------|------:|-----------------|
| `arch-a/latency_config.json` | 168 | 96 | `128x128` | 2-bit | `12x14` | | `arch-a/latency_config.json` | 168 | 96 | `128x128` | 2-bit | `12x14` |
| `arch-b/latency_config.json` | 138 | 128 | `128x128` | 2-bit | `6x23` | | `arch-b/latency_config.json` | 138 | 128 | `128x128` | 2-bit | `6x23` |
@@ -121,10 +125,10 @@ these checked-in artifacts without generating configs at runtime.
`adc_count` is 16, matching the paper's 16-bit fixed-point weight precision. `adc_count` is 16, matching the paper's 16-bit fixed-point weight precision.
The paper does not give a two-dimensional core topology for Arch-A/B, so the The paper does not give a two-dimensional core topology for Arch-A/B, so the
factorizations above preserve core count but cannot reproduce unpublished NoC factorizations above preserve core count but cannot reproduce unpublished NoC
placement details. Released PIMCOMP-NN has no chip-count field; Arch-C is placement details. Released Pimcomp has no chip-count field; Arch-C is
therefore flattened to 64 cores and does not model chip boundaries. therefore flattened to 64 cores and does not model chip boundaries.
The remaining latency and power values come from PIMCOMP-NN's released default The remaining latency and power values come from Pimcomp's released default
configuration. Consequently, instruction/resource comparisons are configuration. Consequently, instruction/resource comparisons are
reproducible, but absolute paper power and energy numbers are not. reproducible, but absolute paper power and energy numbers are not.
@@ -143,9 +147,9 @@ cmake --build third_party/PIMCOMP-NN/build --target PIMCOMP-NN
Do not build either project with `ninja` directly. Do not build either project with `ninja` directly.
## Compile with PIMCOMP ## Compile with Pimcomp
PIMCOMP-NN reads `third_party/PIMCOMP-NN/config.json` directly. Back it up, Pimcomp reads `third_party/PIMCOMP-NN/config.json` directly. Back it up,
select one paper profile, and restore it when the shell exits: select one paper profile, and restore it when the shell exits:
```bash ```bash
@@ -158,7 +162,7 @@ trap 'cp "$CONFIG_BACKUP" "$PIMCOMP/config.json"' EXIT
cp "$PIMCOMP_CONFIGS/arch-a/latency_config.json" "$PIMCOMP/config.json" cp "$PIMCOMP_CONFIGS/arch-a/latency_config.json" "$PIMCOMP/config.json"
``` ```
The Model Zoo files map exactly to PIMCOMP's bundled model names, so compile The Model Zoo files map exactly to Pimcomp's bundled model names, so compile
them directly: them directly:
```bash ```bash
@@ -175,7 +179,7 @@ cd "$PIMCOMP/build"
./PIMCOMP-NN -m=googlenet -r=balance -p=element -o=YES -v=YES -s=YES ./PIMCOMP-NN -m=googlenet -r=balance -p=element -o=YES -v=YES -s=YES
``` ```
VGG-8 first needs PIMCOMP's JSON frontend. Use a temporary ONNX copy because VGG-8 first needs Pimcomp's JSON frontend. Use a temporary ONNX copy because
the released frontend rewrites the input batch dimension in place: the released frontend rewrites the input batch dimension in place:
```bash ```bash
@@ -197,49 +201,77 @@ random placement code occasionally segfaults; an unchanged retry succeeded in
the observed cases. the observed cases.
The paper's optimizer uses a genetic algorithm with population 200 and up to The paper's optimizer uses a genetic algorithm with population 200 and up to
1000 iterations. The checked-out PIMCOMP submodule already has both paper 1000 iterations. The checked-out Pimcomp submodule already has both paper
settings in `backend/GeneticAlgorithm.h`; select them with `-r=GA`. Fitness settings in `backend/GeneticAlgorithm.h`; select them with `-r=GA`. Fitness
evaluation uses OpenMP and bounded bandwidth timelines. Set `OMP_NUM_THREADS` evaluation uses OpenMP and bounded bandwidth timelines. Set `OMP_NUM_THREADS`
to control its parallelism; otherwise OpenMP uses the available CPUs. The GA to control its parallelism; otherwise OpenMP uses the available CPUs. The GA
uses the fixed seed `1`, so repeated serial and parallel runs are reproducible. uses the fixed seed `1`, so repeated serial and parallel runs are reproducible.
## Compare Raptor and PIMCOMP ## Compare Raptor and Pimcomp
The comparison driver uses one random input and one native ONNX-MLIR reference, The comparison driver uses one random input and one native ONNX-MLIR reference,
compiles both instruction streams, runs both through `pimsim-nn`, runs compiles both instruction streams, runs both through Pimsim, runs
functional validation through `pim-simulator`, and writes Markdown and JSON functional validation through `pim-simulator`, and writes Markdown and JSON
reports. reports.
To reproduce the complete Arch-A latency experiment, use the model-by-model To reproduce the default `arch-a`/`arch-b` architectures and both
runner. It verifies the paper GA settings, builds Raptor and the existing latency/throughput modes, use the model-by-model runner. Use `--archs` to
`third_party/PIMCOMP-NN/build` tree, then runs the `element`/batch-1 comparison specify a different architecture set. It verifies the paper GA settings, expects Raptor
for one model at a time and regenerates `results.csv` from the JSON reports: and the existing `third_party/PIMCOMP-NN/build` tree to already be built, then
runs the comparisons in parallel and regenerates `results_comparison.csv` from the JSON
reports:
```bash ```bash
.venv/bin/python validation/tools/pimcomp/run_pimcomp_paper_latency.py .venv/bin/python validation/tools/pim/pimcomp/compare/run_pimcomp_models.py
``` ```
Each model directory reuses regular validation's ignored `inputs/`, `outputs/`, Use `--archs arch-a --mode latency` for only the Arch-A latency experiment, or
`runner/`, `raptor/`, and `simulation/` paths. PIMCOMP-only artifacts and `--archs arch-a arch-b arch-c` to run all three architectures.
`comparison_report.{md,json}` live under `pimcomp/`. The frontend regenerates
one isolated `models/JSON/` graph because PIMCOMP requires that relative Each model directory has a shared ignored `artifacts/common/` directory containing
`inputs/`, `outputs/`, and the native `runner/`.
Pimsim configs and network meshes remain canonical under
`validation/pimsim_configs/pimcomp/` and are referenced in place.
Architecture- and pipeline-specific `raptor/`, `simulation/`, and Pimcomp
artifacts remain under each `artifacts/<arch>/<mode>[/pipelineN]/` comparison
directory; ablation variants use
`artifacts/<arch>/<mode>[/pipelineN]/ablation/<variant>/`. Pimcomp outputs are prepared
once per model/architecture/mode and linked into the other pipeline directories;
`comparison_report.{md,json}`
live under its `pimcomp/`. The frontend regenerates
one isolated `models/JSON/` graph because Pimcomp requires that relative
layout; it is removed after a successful backend run and the shared submodule layout; it is removed after a successful backend run and the shared submodule
model directory is never modified. Models requiring BatchNormalization folding model directory is never modified. The original ONNX model is passed to the
also receive a prepared ONNX file; other models use the original ONNX directly. frontend unchanged.
PIMCOMP's source tree and build directory remain unchanged at runtime. Use Pimcomp's source tree and build directory remain unchanged at runtime. Use
`--models vgg8` to run one model, `--resume` after an interruption, `--dry-run` `--models vgg8` to run one model, `--mode throughput` to select one mode,
`--pipeline 4` to select one throughput pipeline, `--only raptor` or
`--only pimcomp` to reuse the other compiler's existing artifacts, `--dry-run`
to inspect every command, or `--out-dir PATH` to keep results outside to inspect every command, or `--out-dir PATH` to keep results outside
`validation/`. The runner continues after a failed model so all reports are `validation/`. Use `--clean` to remove generated comparison artifacts and
produced. summaries, including stale reference lock files. Selecting a subset replaces only those comparison rows and
recomputes the aggregate `results_comparison.csv`; missing shared inputs, outputs, or the
reference runner are generated even for an isolated run. Use `--jobs 4` to cap
parallel comparisons. The per-stage timeout is unlimited by default; pass a
positive `--timeout-seconds` value to impose one.
Throughput comparisons default to `pimsim-nn --fast` with a 1000 ms
convergence deadline. Add `--no-fast` for authoritative full-duration runs.
Pimcomp receives the original ONNX model, and its frontend applies native
BatchNormalization fusion when the graph matches its supported Conv/Gemm pattern.
The runner continues after a failed model so all reports are produced.
The known Pimcomp batch-scheduling correctness issue and a reproducible
reference-intermediate prefill experiment are documented in
[`validation/tools/pim/pimcomp/correctness/README.md`](../../tools/pim/pimcomp/correctness/README.md).
Arch-A low-latency example: Arch-A low-latency example:
```bash ```bash
RAPTOR_ROOT=$PWD RAPTOR_ROOT=$PWD
"$RAPTOR_ROOT/.venv/bin/python" "$RAPTOR_ROOT/validation/tools/pimcomp/compare_raptor_pimcomp.py" \ "$RAPTOR_ROOT/.venv/bin/python" "$RAPTOR_ROOT/validation/tools/pim/pimcomp/compare/compare_raptor_pimcomp_model.py" \
--model "$RAPTOR_ROOT/validation/networks/pimcomp_models/resnet34/resnet34-v1-7.onnx" \ --model "$RAPTOR_ROOT/validation/networks/pimcomp_models/resnet34/resnet34-v1-7.onnx" \
--out-dir "$RAPTOR_ROOT/validation/networks/pimcomp_models/resnet34" \ --out-dir "$RAPTOR_ROOT/validation/networks/pimcomp_models/resnet34/artifacts/arch-a/latency" \
--pimcomp-config "$RAPTOR_ROOT/validation/pimsim_configs/pimcomp/arch-a/latency_config.json" \ --pimcomp-config "$RAPTOR_ROOT/validation/pimsim_configs/pimcomp/arch-a/latency_config.json" \
--core-count 168 \ --core-count 168 \
--crossbar-count 96 \ --crossbar-count 96 \
@@ -251,10 +283,10 @@ RAPTOR_ROOT=$PWD
``` ```
Use the same command with Use the same command with
`yolo11n/yolo11n-latency.onnx` to probe YOLO11n. Released PIMCOMP-NN cannot `yolo11n/yolo11n-pimsim-nn.onnx` to probe YOLO11n. Released Pimcomp cannot
compile it: the frontend stops at `/model.2/Split`, and it also has no mapping compile it: the frontend stops at `/model.2/Split`, and it also has no mapping
for YOLO11n's two nearest-neighbor `Resize` nodes. Treating the emitted prefix for YOLO11n's two nearest-neighbor `Resize` nodes. Treating the emitted prefix
as YOLO11n would produce a misleading latency, so no PIMCOMP number is as YOLO11n would produce a misleading latency, so no Pimcomp number is
reported for this model. reported for this model.
For Arch-A high throughput, use `--pimsim-mode throughput For Arch-A high throughput, use `--pimsim-mode throughput
@@ -275,18 +307,18 @@ generated report.
The functional and non-functional simulators intentionally consume different The functional and non-functional simulators intentionally consume different
artifacts: artifacts:
- Raptor and PIMCOMP are validated against the native ONNX-MLIR reference as - Raptor and Pimcomp are validated against the native ONNX-MLIR reference as
FP32 programs in the Rust simulator. Raptor's emitted program is already FP32 programs in the Rust simulator. Raptor's emitted program is already
FP32. The PIMCOMP-to-Rust export expands its element-addressed storage and FP32. The Pimcomp-to-Rust export expands its element-addressed storage and
byte-sized transfers to FP32, emits `setbw 32, 32`, and keeps vector byte-sized transfers to FP32, emits `setbw 32, 32`, and keeps vector
`imm_len` fields as element counts. `imm_len` fields as element counts.
- PIMCOMP's original `SimulationInfo.gz` is copied unchanged for `pimsim-nn`. - Pimcomp's original `SimulationInfo.gz` is copied unchanged for Pimsim.
PIMCOMP hardcodes `setbw 8, 8` and one byte per element without performing Pimcomp hardcodes `setbw 8, 8` and one byte per element without performing
numerical quantization; this artifact is used only for latency estimation. numerical quantization; this artifact is used only for latency estimation.
- Raptor's original FP32 artifact remains unchanged for functional validation. - Raptor's original FP32 artifact remains unchanged for functional validation.
A separate `raptor/pimsim_nn/` view uses `setbw 8, 8` and scales its A separate `raptor/pimsim_nn/` view uses `setbw 8, 8` and scales its
byte-addressed storage and transfer sizes from four bytes to one byte per byte-addressed storage and transfer sizes from four bytes to one byte per
element. Vector `imm_len` fields remain element counts. Like PIMCOMP's element. Vector `imm_len` fields remain element counts. Like Pimcomp's
artifact, this view is not numerically valid and is used only for a fair artifact, this view is not numerically valid and is used only for a fair
non-functional comparison. non-functional comparison.
@@ -297,11 +329,11 @@ either latency-only artifact for semantic validation.
Current Raptor status: Current Raptor status:
- VGG-8, ResNet-18, fixed-batch ResNet-34, and GoogLeNet compile on Arch-A. - VGG-8, ResNet-18, fixed-batch ResNet-34, and GoogLeNet compile on Arch-A.
- Use `googlenet-12-latency.onnx` for the paper-matched latency comparison. - Use `googlenet-12-pimsim-nn.onnx` for the paper-matched latency comparison.
It removes the two LRN nodes and terminal softmax that PIMCOMP does not It removes the two LRN nodes and terminal softmax that Pimcomp does not
schedule. schedule.
- Raptor currently accepts one square `--crossbar-size`; Arch-C's rectangular - Raptor currently accepts one square `--crossbar-size`; Arch-C's rectangular
`512x1024` arrays can therefore be compiled by PIMCOMP but not compared `512x1024` arrays can therefore be compiled by Pimcomp but not compared
exactly with Raptor. exactly with Raptor.
Do not change the hardware profile to bypass either limitation; that would no Do not change the hardware profile to bypass either limitation; that would no
@@ -319,8 +351,8 @@ rsync -azL validation/networks/pimcomp_models/ \
"monolith:$REMOTE_REPO/validation/networks/pimcomp_models/" "monolith:$REMOTE_REPO/validation/networks/pimcomp_models/"
rsync -az validation/pimsim_configs/pimcomp/ \ rsync -az validation/pimsim_configs/pimcomp/ \
"monolith:$REMOTE_REPO/validation/pimsim_configs/pimcomp/" "monolith:$REMOTE_REPO/validation/pimsim_configs/pimcomp/"
rsync -az validation/tools/pimcomp/ \ rsync -az validation/tools/pim/ \
"monolith:$REMOTE_REPO/validation/tools/pimcomp/" "monolith:$REMOTE_REPO/validation/tools/pim/"
rsync -az --exclude=.git --exclude=build --exclude=output \ rsync -az --exclude=.git --exclude=build --exclude=output \
third_party/PIMCOMP-NN/ \ third_party/PIMCOMP-NN/ \
"monolith:$REMOTE_REPO/third_party/PIMCOMP-NN/" "monolith:$REMOTE_REPO/third_party/PIMCOMP-NN/"
@@ -334,10 +366,10 @@ cd /home/gmagnani/Project/Raptor
# One-time setup if the repository virtual environment is absent. # One-time setup if the repository virtual environment is absent.
python3 -m venv .venv python3 -m venv .venv
.venv/bin/python -m pip install numpy onnx onnxruntime onnxsim colorama .venv/bin/python -m pip install numpy onnx onnxruntime colorama
# Run every latency comparison serially. # Run every configured comparison in parallel.
.venv/bin/python validation/tools/pimcomp/run_pimcomp_paper_latency.py .venv/bin/python validation/tools/pim/pimcomp/compare/run_pimcomp_models.py
``` ```
Copy reports back without transferring large compiler artifacts: Copy reports back without transferring large compiler artifacts:
@@ -1,3 +0,0 @@
model,raptor_latency_ms,pimcomp_latency_ms,raptor_energy_pj,pimcomp_energy_pj,faster_compiler,speedup
vgg8,1.521060,7.985074,486309111.040001,1597904071.120000,raptor,5.25
resnet18,33.552733,58.853613,9702508727.119982,13983148468.119974,raptor,1.75
1 model raptor_latency_ms pimcomp_latency_ms raptor_energy_pj pimcomp_energy_pj faster_compiler speedup
2 vgg8 1.521060 7.985074 486309111.040001 1597904071.120000 raptor 5.25
3 resnet18 33.552733 58.853613 9702508727.119982 13983148468.119974 raptor 1.75
@@ -1,6 +0,0 @@
Operation,Result,Compile,Host mem,Cores mem,Cores,Xbars,Latency,Power,Energy
vgg8-mnist-reconstructed,PASS,1.009 s,1.37 MiB,3.14 MiB,141,761,1.465778 ms,325.627854 mW,477298145.040001 pJ
resnet18-v1-7,PASS,11.548 s,9.89 MiB,40.24 MiB,168,7676,28.099952 ms,312.513408 mW,8781611766.119984 pJ
resnet34-v1-7,PASS,28.495 s,9.90 MiB,48.89 MiB,168,15292,45.781486 ms,326.833870 mW,14962940227.679951 pJ
googlenet-12-latency,PASS,6.573 s,10.74 MiB,22.41 MiB,168,7176,13.371204 ms,457.538139 mW,6117835798.919991 pJ
yolo11n-latency,FAIL,58.572 s,82.55 MiB,185.68 MiB,168,6484,885.264931 ms,189.218985 mW,167508931321.001465 pJ
1 Operation Result Compile Host mem Cores mem Cores Xbars Latency Power Energy
2 vgg8-mnist-reconstructed PASS 1.009 s 1.37 MiB 3.14 MiB 141 761 1.465778 ms 325.627854 mW 477298145.040001 pJ
3 resnet18-v1-7 PASS 11.548 s 9.89 MiB 40.24 MiB 168 7676 28.099952 ms 312.513408 mW 8781611766.119984 pJ
4 resnet34-v1-7 PASS 28.495 s 9.90 MiB 48.89 MiB 168 15292 45.781486 ms 326.833870 mW 14962940227.679951 pJ
5 googlenet-12-latency PASS 6.573 s 10.74 MiB 22.41 MiB 168 7176 13.371204 ms 457.538139 mW 6117835798.919991 pJ
6 yolo11n-latency FAIL 58.572 s 82.55 MiB 185.68 MiB 168 6484 885.264931 ms 189.218985 mW 167508931321.001465 pJ
+29 -20
View File
@@ -1,7 +1,7 @@
# Operation Validation Suite # Operation validation suite
This directory contains the ONNX models used by `validation/validate.py` to This directory contains the ONNX models used by `validation/validate.py` to
validate individual operations through compilation, PIM simulation, and validate individual operations through compilation, Pim simulation, and
comparison with the ONNX-MLIR reference runtime. comparison with the ONNX-MLIR reference runtime.
## Naming ## Naming
@@ -38,12 +38,12 @@ Run the complete suite with deadlock detection:
Use `--compile-only` for compiler and deadlock checks, then `--run-only` to Use `--compile-only` for compiler and deadlock checks, then `--run-only` to
reuse those artifacts for reference execution, simulation, and comparison. reuse those artifacts for reference execution, simulation, and comparison.
The validator prints the complete operation results table before its summary The validator prints separate latency and throughput operation tables before
and writes the same rows to `validation_results.csv`. its summary and writes all of their rows to `validation_results.csv`.
## Complete inventory ## Complete inventory
The suite contains 168 models. Tensor shapes, attributes, and constants are The suite contains 177 models. Tensor shapes, attributes, and constants are
defined in `gen_tests.py` and in the checked-in ONNX models. defined in `gen_tests.py` and in the checked-in ONNX models.
### Add (5) ### Add (5)
@@ -64,12 +64,13 @@ defined in `gen_tests.py` and in the checked-in ONNX models.
| `negative_axis` | Concatenates tensors using a negative axis. | | `negative_axis` | Concatenates tensors using a negative axis. |
| `three_inputs_channel_axis` | Concatenates three runtime NCHW tensors along the channel axis. | | `three_inputs_channel_axis` | Concatenates three runtime NCHW tensors along the channel axis. |
### Conv (34) ### Conv (42)
| Case | Description | | Case | Description |
|---|---| |---|---|
| `batch_2` | Batched Conv with SAME_UPPER padding and bias. | | `batch_2` | Batched Conv with SAME_UPPER padding and bias. |
| `batch_4_pointwise` | Pointwise Conv with batch size four. | | `batch_4_pointwise` | Pointwise Conv with batch size four. |
| `input_224_7x7_stride2` | 224x224 RGB Conv with 64 output channels, a 7x7 kernel, stride two, and bias. |
| `depthwise_1024_channels` | Depthwise pointwise Conv with 1024 groups. | | `depthwise_1024_channels` | Depthwise pointwise Conv with 1024 groups. |
| `depthwise_grouped` | Depthwise-style grouped Conv with one input channel per group. | | `depthwise_grouped` | Depthwise-style grouped Conv with one input channel per group. |
| `dilated_3x3` | Conv with a dilated 3x3 kernel. | | `dilated_3x3` | Conv with a dilated 3x3 kernel. |
@@ -88,14 +89,21 @@ defined in `gen_tests.py` and in the checked-in ONNX models.
| `non_square_kernel_1x3` | Conv with a non-square 1x3 kernel. | | `non_square_kernel_1x3` | Conv with a non-square 1x3 kernel. |
| `non_square_kernel_3x1` | Conv with a non-square 3x1 kernel. | | `non_square_kernel_3x1` | Conv with a non-square 3x1 kernel. |
| `non_uniform_stride` | Conv with different height and width strides. | | `non_uniform_stride` | Conv with different height and width strides. |
| `output_channel_grouping_minimal` | Minimal 64-to-256 pointwise Conv for output-channel grouping. |
| `pointwise_1x1` | Basic pointwise channel-mixing Conv. | | `pointwise_1x1` | Basic pointwise channel-mixing Conv. |
| `pointwise_tiled_chain` | Relu and chained pointwise Convs with a tiled intermediate. | | `pointwise_tiled_chain` | Relu and chained pointwise Convs with a tiled intermediate. |
| `real_asymmetric_padding` | Conv with asymmetric explicit padding. | | `real_asymmetric_padding` | Conv with asymmetric explicit padding. |
| `relu_conv_store` | Conv, Relu, and a second Conv to validate an intermediate stored result. | | `relu_conv_store` | Conv, Relu, and a second Conv to validate an intermediate stored result. |
| `same_lower_3x3` | 3x3 Conv with SAME_LOWER padding. | | `same_lower_3x3` | 3x3 Conv with SAME_LOWER padding. |
| `same_padding_3x3` | 3x3 Conv with SAME_UPPER padding. | | `same_padding_3x3` | 3x3 Conv with SAME_UPPER padding. |
| `simple` | Hand-authored basic 2x2 Conv. | | `kernel_2x2` | Hand-authored Conv with a 2x2 kernel. |
| `stride_2` | 3x3 Conv with stride two. | | `stride_2` | 3x3 Conv with stride two. |
| `strategy_depthwise_16` | 16-channel depthwise 3x3 Conv for depthwise lowering. |
| `strategy_input_k_tiled` | 32-channel 3x3 Conv sized to exercise input-K tiling. |
| `strategy_output_channel_tiled` | 8-to-192 3x3 Conv sized to exercise output-channel tiling. |
| `strategy_streamed_packed` | 3-to-16 3x3 Conv on 128x128 input for streamed packed lowering. |
| `strategy_streamed_patch` | 3-to-16 3x3 Conv on 64x64 input for streamed patch lowering. |
| `strategy_tiled_2d` | 32-to-192 3x3 Conv sized to exercise 2D tiling. |
| `with_bias_3x3` | Multi-channel 3x3 Conv with bias. | | `with_bias_3x3` | Multi-channel 3x3 Conv with bias. |
| `with_constant` | Hand-authored SAME_UPPER Conv with constant weight and bias. | | `with_constant` | Hand-authored SAME_UPPER Conv with constant weight and bias. |
| `without_kernel_shape_attr` | Conv whose kernel shape is inferred from its weight tensor. | | `without_kernel_shape_attr` | Conv whose kernel shape is inferred from its weight tensor. |
@@ -135,30 +143,30 @@ defined in `gen_tests.py` and in the checked-in ONNX models.
| `dynamic_beta` | Uses runtime operands and bias with non-default beta scaling. | | `dynamic_beta` | Uses runtime operands and bias with non-default beta scaling. |
| `dynamic_bias` | Uses runtime matrix operands and runtime bias. | | `dynamic_bias` | Uses runtime matrix operands and runtime bias. |
| `dynamic_bias_alpha_beta` | Combines runtime operands and bias with alpha and beta scaling. | | `dynamic_bias_alpha_beta` | Combines runtime operands and bias with alpha and beta scaling. |
| `dynamic_transB` | Transposes a runtime right-hand matrix. | | `dynamic_transpose_b` | Transposes a runtime right-hand matrix. |
| `huge_1024` | Uses 1024-wide inner and output dimensions. | | `huge_1024` | Uses 1024-wide inner and output dimensions. |
| `large` | Exercises larger rectangular matrices. | | `large` | Exercises larger rectangular matrices. |
| `large_k_small_n` | Uses a large reduction dimension and narrow output. | | `large_k_small_n` | Uses a large reduction dimension and narrow output. |
| `non_square` | Uses different reduction and output widths. | | `non_square` | Uses different reduction and output widths. |
| `scalar_bias` | Broadcasts a scalar bias to the full output. | | `scalar_bias` | Broadcasts a scalar bias to the full output. |
| `simple` | Basic Gemm with square weights. | | `square_weights` | Basic Gemm with square weights. |
| `small` | Tiny Gemm for fast focused validation. | | `small` | Tiny Gemm for fast focused validation. |
| `small_k_large_n` | Uses a modest reduction dimension and wide output. | | `small_k_large_n` | Uses a modest reduction dimension and wide output. |
| `transA` | Transposes the left-hand matrix. | | `transpose_a` | Transposes the left-hand matrix. |
| `transA_transB` | Transposes both matrix operands. | | `transpose_a_and_b` | Transposes both matrix operands. |
| `transB` | Transposes the right-hand weight matrix. | | `transpose_b` | Transposes the right-hand weight matrix. |
| `transB_with_bias` | Combines a transposed weight matrix with bias. | | `transpose_b_with_bias` | Combines a transposed weight matrix with bias. |
| `with_bias` | Basic matrix product with vector bias. | | `with_bias` | Basic matrix product with vector bias. |
### Gemv (5) ### Gemv (5)
| Case | Description | | Case | Description |
|---|---| |---|---|
| `constant` | Vector-matrix product with all inputs constant. | | `all_constant` | Vector-matrix product with all inputs constant. |
| `simple` | Basic single-row vector-matrix product. | | `constant_weight` | Basic single-row vector-matrix product with constant weights. |
| `with_heterogeneous_constant` | Adds a non-uniform constant bias pattern. | | `non_uniform_bias` | Adds a non-uniform constant bias pattern. |
| `with_homogeneous_constant` | Adds a constant bias matching the output shape. | | `uniform_bias` | Adds a uniform constant bias pattern. |
| `with_scalar_constant` | Adds a scalar broadcast bias. | | `scalar_bias` | Adds a scalar broadcast bias. |
### MatMul (12) ### MatMul (12)
@@ -177,11 +185,12 @@ defined in `gen_tests.py` and in the checked-in ONNX models.
| `vector_matrix` | Vector-matrix multiplication producing a 1D output. | | `vector_matrix` | Vector-matrix multiplication producing a 1D output. |
| `yolo_attention` | YOLO11n rank-4 dynamic MatMul-scale-transpose-MatMul attention chain. | | `yolo_attention` | YOLO11n rank-4 dynamic MatMul-scale-transpose-MatMul attention chain. |
### Mul (5) ### Mul (6)
| Case | Description | | Case | Description |
|---|---| |---|---|
| `after_conv` | Conv followed by per-channel scaling. | | `after_conv` | Conv followed by per-channel scaling. |
| `after_conv_scalar_constant` | Conv followed by Mul with a scalar constant. |
| `basic` | Elementwise Mul on two inputs with identical shapes. | | `basic` | Elementwise Mul on two inputs with identical shapes. |
| `channel_broadcast_1024` | Mul with NCHW per-channel broadcasting over 1024 channels. | | `channel_broadcast_1024` | Mul with NCHW per-channel broadcasting over 1024 channels. |
| `leading_dimension_broadcast` | Mul with trailing-dimension broadcasting. | | `leading_dimension_broadcast` | Mul with trailing-dimension broadcasting. |
@@ -254,7 +263,7 @@ defined in `gen_tests.py` and in the checked-in ONNX models.
| `height_only` | Nearest-neighbor resize of only the height dimension. | | `height_only` | Nearest-neighbor resize of only the height dimension. |
| `nearest_2x` | Nearest-neighbor upsampling by a factor of two. | | `nearest_2x` | Nearest-neighbor upsampling by a factor of two. |
| `nearest_downsample` | Nearest-neighbor downsampling. | | `nearest_downsample` | Nearest-neighbor downsampling. |
| `non_uniform` | Nearest-neighbor resize with different spatial scales. | | `non_uniform_scales` | Nearest-neighbor resize with different spatial scales. |
| `width_only` | Nearest-neighbor resize of only the width dimension. | | `width_only` | Nearest-neighbor resize of only the width dimension. |
| `with_sizes` | Resize using explicit output sizes instead of scales. | | `with_sizes` | Resize using explicit output sizes instead of scales. |
Binary file not shown.

Some files were not shown because too many files have changed in this diff Show More