Compare commits
231 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 4acd3b0c81 | |||
| 42c236b6a5 | |||
| e2cefd3127 | |||
| 7a3a808ae8 | |||
| 0712c5ba29 | |||
| aeedf2f566 | |||
| a39fdba366 | |||
| a963009855 | |||
| 10b6ee6c32 | |||
| f4a3b012cc | |||
| 9ca1a0ed9f | |||
| a0131c6f7a | |||
| 1e7699f3e5 | |||
| ea8c5386b6 | |||
| 78be3d9613 | |||
| 2daa75a7b2 | |||
| c12f69133d | |||
| 1b4f070bef | |||
| 060a21172e | |||
| 78bfb8a9aa | |||
| 4c6fc1173e | |||
| 2b899b62a8 | |||
| 87bd7b726d | |||
| ab1b3584dd | |||
| e4d8fba579 | |||
| dddbf59241 | |||
| e76c278d1c | |||
| 0fabcbacc2 | |||
| b53586003a | |||
| a73afd62a4 | |||
| 1b4fa949ef | |||
| 0932f75901 | |||
| 7e204e9e63 | |||
| 77cd1157ba | |||
| 2fa9a7c78f | |||
| d367d09cfb | |||
| a877291195 | |||
| f861c6f593 | |||
| e812ffd82f | |||
| 45288635b3 | |||
| e28a2b824d | |||
| 5415e95528 | |||
| 620e381cfb | |||
| 4964e889da | |||
| f3a4e19f7c | |||
| e7611be8e1 | |||
| c59e320efa | |||
| c491078757 | |||
| bb4594fc01 | |||
| 45578ef4c4 | |||
| b491ff77b1 | |||
| 8da4603be3 | |||
| 24c6ea3c6e | |||
| 563921bccd | |||
| b651968f30 | |||
| d88bced271 | |||
| f47fcebb83 | |||
| 3e468b58c8 | |||
| 6bc7709aa4 | |||
| a893d23a74 | |||
| 961fd613bd | |||
| 2d17b50ca8 | |||
| dbb66be93e | |||
| 5c94e00f43 | |||
| 6bad9a8008 | |||
| ab54243fda | |||
| 5f42da36ae | |||
| ae67d720c6 | |||
| d788178749 | |||
| c744f388dc | |||
| 51fdb830e5 | |||
| d1a29ace3c | |||
| 61e3ea9996 | |||
| fed6d343e5 | |||
| 871fcfa832 | |||
| 1f4f58de1c | |||
| 8338caf3f3 | |||
| 47f6715296 | |||
| 2bfc033af9 | |||
| 83a54e28e4 | |||
| cc9b025a35 | |||
| c4dd28a607 | |||
| 8d3eb929f6 | |||
| f5e1c2e706 | |||
| 94c96195b9 | |||
| 645539317b | |||
| 4a98e88e97 | |||
| f492400eda | |||
| e8f09fd67f | |||
| 78e97f9fd8 | |||
| 984f362623 | |||
| 568fd90542 | |||
| be0bcc9dcc | |||
| 62dd40ee89 | |||
| 2b4115699a | |||
| 3a985b3675 | |||
| 4ab24eb288 | |||
| e083c27d80 | |||
| 852bef7605 | |||
| 237654dadf | |||
| 6d69600bc1 | |||
| aec80529ca | |||
| 8ddbbcecfa | |||
| 90c4339808 | |||
| 08870de1a6 | |||
| a34ac223c0 | |||
| 0fa10b4074 | |||
| e166ff7e1d | |||
| a70a8f77cf | |||
| 800c0c4316 | |||
| 1e9e61f5a9 | |||
| 27410207c4 | |||
| cbc9808229 | |||
| 69021d56aa | |||
| dc5edd032c | |||
| e33f517221 | |||
| f94b3d1020 | |||
| 20cf40c9ba | |||
| 37a59054a5 | |||
| 2a8faf9c6b | |||
| 01b9d03fc6 | |||
| 501e6c76f3 | |||
| 3c2667f11e | |||
| 0a5e73c3ea | |||
| 636310d0cb | |||
| 356be6ccc2 | |||
| b678e55d3c | |||
| ab63498f3f | |||
| 7c3943bd06 | |||
| c0238c0d06 | |||
| ff36729140 | |||
| cf93caecd5 | |||
| 2d5b03c08f | |||
| a41f694cf0 | |||
| 8bb0babf1b | |||
| 819d8af0f7 | |||
| 832bd7f1f7 | |||
| 82b44a6387 | |||
| 7fcc765d6e | |||
| f34698a2b6 | |||
| 1ab489fe0a | |||
| cbf7b235f1 | |||
| 00414dd1d9 | |||
| 783dffe553 | |||
| 874a2f53e6 | |||
| 4bdaa57656 | |||
| 1a5d7d2a3f | |||
| 013ae0ac2a | |||
| c6b02af7a9 | |||
| d2048bd394 | |||
| 158f0f0c54 | |||
| 532cac8246 | |||
| d609e84054 | |||
| addfc8a86e | |||
| 0f240af271 | |||
| bdc4ca33f3 | |||
| b79c333c6c | |||
| eea9261c7b | |||
| e8a08f6dd0 | |||
| 4855a2e105 | |||
| 3a7a832198 | |||
| 48ca6bd28d | |||
| f595cc6ffd | |||
| c734f1b37e | |||
| b79ce8eeaa | |||
| 76a37e198f | |||
| 7f3c7464b4 | |||
| c77ffa9c56 | |||
| 495186503c | |||
| 2c1da813b5 | |||
| 8337a11ce9 | |||
| d136136d22 | |||
| 074eb183c7 | |||
| 43ed3914b8 | |||
| 6aaf1c0870 | |||
| fe35b3ed43 | |||
| 90a9339686 | |||
| a50e77ff38 | |||
| f56c4159b5 | |||
| 5637c861b4 | |||
| 94157a8404 | |||
| 68a3521978 | |||
| a103ba328b | |||
| e263e05f56 | |||
| 34c29fdec4 | |||
| aa088e2ba5 | |||
| 2836e759ab | |||
| 8071ebab0b | |||
| f1602c0550 | |||
| de0a2f4561 | |||
| 1c4a5bde76 | |||
| 78242e2887 | |||
| fe244d5aa1 | |||
| d09e76c8f9 | |||
| c5e608fa5b | |||
| 43f3ccdd21 | |||
| 8d95c604a6 | |||
| 55eda487dc | |||
| 061139aefb | |||
| ea61540e08 | |||
| 324178cba8 | |||
| e71ba07cd5 | |||
| 64a3805619 | |||
| 9f9e7c0892 | |||
| 03eab42971 | |||
| c15aba5d96 | |||
| 4821e8a55e | |||
| 88bb223bb1 | |||
| 623ee62a04 | |||
| ad56888b0b | |||
| f993840641 | |||
| 0c7db55a24 | |||
| 41de3cb150 | |||
| 4f3570520c | |||
| 628dc630a4 | |||
| 80a7298552 | |||
| 8ad504fcdf | |||
| e6f442c5d2 | |||
| f6b97b3813 | |||
| 26317ea7d0 | |||
| 909c4acfdd | |||
| feaff820e1 | |||
| 1e279ae9bb | |||
| 57f0cca8c0 | |||
| 5ff364027b | |||
| b1272d2283 | |||
| 58e6587697 | |||
| f6c8cc4aa5 | |||
| 566630b99a | |||
| 74931ad75b | |||
| f2fe147961 |
@@ -0,0 +1,382 @@
|
||||
# Graph Compute Batch Physical-Fragment Invariant
|
||||
|
||||
## Status
|
||||
|
||||
This document is **normative** for Raptor's Spatial graph IR.
|
||||
|
||||
Every developer or coding agent modifying Spatial graph construction, graph
|
||||
verification, Blueprint handling, or `MergeComputeNodes` must read this file
|
||||
after `README.md` and `AGENTS.md`.
|
||||
|
||||
`AGENTS.md` must reference this invariant:
|
||||
|
||||
```text
|
||||
* `.agents/invariants/GRAPH_COMPUTE_BATCH_INVARIANT.md`
|
||||
```
|
||||
|
||||
## Scope
|
||||
|
||||
This invariant applies to:
|
||||
|
||||
- `spat.graph_compute_batch`;
|
||||
- graph-level values produced by it;
|
||||
- `tensor.parallel_insert_slice` operations that publish its lane results;
|
||||
- `spat.blueprint` operations that describe logical reconstruction;
|
||||
- graph analyses and transformations that consume those values;
|
||||
- the graph-to-scheduled transition in `MergeComputeNodes`.
|
||||
|
||||
It does **not** impose the same representation on:
|
||||
|
||||
- `spat.scheduled_compute`;
|
||||
- `spat.scheduled_compute_batch`;
|
||||
- `pim.core` or `pim.core_batch`;
|
||||
- values whose cross-core movement is already represented by explicit
|
||||
`spat.channel_send` and `spat.channel_receive` operations.
|
||||
|
||||
Scheduled IR represents execution on assigned cores. Communication and value
|
||||
availability there are defined by local SSA forwarding and explicit
|
||||
send/receive operations, not by the graph physical-fragment invariant.
|
||||
|
||||
## Core invariant
|
||||
|
||||
For every result of a `spat.graph_compute_batch` with `N` graph lanes:
|
||||
|
||||
1. Every graph lane produces exactly one fragment for that result.
|
||||
2. All lanes produce fragments with the same exact ranked tensor type `F`.
|
||||
3. The graph result is a physical collection of those fragments with type:
|
||||
|
||||
```text
|
||||
tensor<N x shape(F) x element-type(F)>
|
||||
```
|
||||
|
||||
Conceptually, the result is `N × F`: one leading physical fragment-slot
|
||||
dimension followed by the complete per-lane fragment shape.
|
||||
4. Physical slot `i` identifies a fragment publication. It does not, by itself,
|
||||
identify a row, column, channel, tile, or any other logical tensor position.
|
||||
5. The result type carries no logical reconstruction order.
|
||||
|
||||
The leading dimension is therefore a **physical fragment-slot dimension**, not
|
||||
a logical tensor dimension.
|
||||
|
||||
## Per-lane computation is unrestricted
|
||||
|
||||
The invariant constrains the published result representation, not what a lane
|
||||
may compute.
|
||||
|
||||
A graph lane may:
|
||||
|
||||
- read several input slices;
|
||||
- perform reductions;
|
||||
- add or combine multiple columns;
|
||||
- execute matrix/vector operations;
|
||||
- produce a fragment that corresponds to any logical region;
|
||||
- participate in a multi-stage or logarithmic reduction tree implemented by
|
||||
following `spat.graph_compute` or `spat.graph_compute_batch` operations.
|
||||
|
||||
Arithmetic combination is graph computation. `spat.blueprint` is not an
|
||||
arithmetic reduction operation.
|
||||
|
||||
### Example: `16×4 -> 16×2`
|
||||
|
||||
Two graph lanes may compute:
|
||||
|
||||
```text
|
||||
lane 0: input[:, 0] + input[:, 1] -> tensor<16x1>
|
||||
lane 1: input[:, 2] + input[:, 3] -> tensor<16x1>
|
||||
```
|
||||
|
||||
The physical graph result is:
|
||||
|
||||
```text
|
||||
tensor<2x16x1>
|
||||
```
|
||||
|
||||
A Blueprint then maps:
|
||||
|
||||
```text
|
||||
physical slot 0 -> logical output[:, 0:1]
|
||||
physical slot 1 -> logical output[:, 1:2]
|
||||
```
|
||||
|
||||
and describes the logical result `tensor<16x2>`.
|
||||
|
||||
For a larger reduction, following graph compute batches may reduce fragments in
|
||||
`ceil(log2(N))` stages. Every intermediate batch still publishes a physical
|
||||
`batch × fragment` collection.
|
||||
|
||||
## Physical publication inside `spat.graph_compute_batch`
|
||||
|
||||
The batch body must publish each lane's fragment into the physical result.
|
||||
|
||||
For one result with fragment type `F`, the corresponding
|
||||
`tensor.parallel_insert_slice` must insert the fragment into one slot of the
|
||||
physical `N × F` destination:
|
||||
|
||||
```text
|
||||
physical offsets = [slot, 0, 0, ...]
|
||||
physical sizes = [1, shape(F)...]
|
||||
physical strides = [1, 1, 1, ...]
|
||||
```
|
||||
|
||||
The slot may be the graph lane directly or a statically analyzable permutation
|
||||
of it. The insertion describes physical slot placement only. It must not use a
|
||||
logical output dimension as the physical batch dimension.
|
||||
|
||||
For each graph result, the body must contain exactly one physical publication
|
||||
per graph lane. Since the body executes once per lane, this normally means one
|
||||
`tensor.parallel_insert_slice` operation targeting that result.
|
||||
|
||||
## Logical reconstruction
|
||||
|
||||
Logical reconstruction is separate from physical publication.
|
||||
|
||||
The reconstruction descriptor defines, for every physical fragment slot:
|
||||
|
||||
- which physical batch operand owns the fragment;
|
||||
- which physical slot contains it;
|
||||
- its destination offsets in the logical tensor;
|
||||
- its destination sizes;
|
||||
- its destination strides;
|
||||
- coverage and conflict policy where relevant.
|
||||
|
||||
The persistent owner of this information is `spat.blueprint` or an equivalent
|
||||
explicit graph-level reconstruction operation.
|
||||
|
||||
A logical consumer must not infer reconstruction from the physical tensor type
|
||||
or assume that physical slot order equals logical order.
|
||||
|
||||
The logical mapping may be arbitrary. For example:
|
||||
|
||||
```text
|
||||
physical slot 0 -> logical row 13
|
||||
physical slot 1 -> logical row 4
|
||||
physical slot 2 -> logical row 10
|
||||
```
|
||||
|
||||
The physical result remains a regular `batch × fragment` tensor.
|
||||
|
||||
## Relationship between `parallel_insert_slice` and Blueprint
|
||||
|
||||
During graph construction, an algorithm may naturally describe logical
|
||||
placement with `tensor.parallel_insert_slice` geometry. Before the graph is in
|
||||
its canonical form:
|
||||
|
||||
1. that geometry must be separated from physical fragment publication;
|
||||
2. the graph batch result must be normalized to `N × F`;
|
||||
3. the logical insertion geometry must be transferred to a persistent
|
||||
`spat.blueprint` reconstruction descriptor.
|
||||
|
||||
After normalization:
|
||||
|
||||
- `parallel_insert_slice` inside `spat.graph_compute_batch` publishes into
|
||||
physical fragment slots;
|
||||
- `spat.blueprint` describes reconstruction into the logical tensor.
|
||||
|
||||
The original graph operation may be erased only after all reconstruction
|
||||
information needed by later stages has a persistent owner.
|
||||
|
||||
## Blueprint semantics
|
||||
|
||||
Blueprint is placement/reconstruction metadata. It may:
|
||||
|
||||
- concatenate fragments;
|
||||
- reorder fragments;
|
||||
- insert fragments into arbitrary disjoint logical regions;
|
||||
- describe complete or partial logical coverage;
|
||||
- expose a logical tensor view when materialization is required.
|
||||
|
||||
Blueprint must not silently perform arithmetic such as addition, multiplication,
|
||||
maximum, or reduction. Such transformations must be represented by following
|
||||
`spat.graph_compute` or `spat.graph_compute_batch` operations.
|
||||
|
||||
A Blueprint consuming a physical fragment batch must explicitly identify the
|
||||
physical source slot for every logical fragment. It must not derive that slot
|
||||
from operand order unless that convention is explicitly represented and
|
||||
verified.
|
||||
|
||||
## Multiple results
|
||||
|
||||
A `spat.graph_compute_batch` may have several results.
|
||||
|
||||
For each result `r` independently:
|
||||
|
||||
- every lane produces one fragment of type `F_r`;
|
||||
- the graph result type is `N × F_r`;
|
||||
- its physical publication and logical reconstruction descriptor are verified
|
||||
independently.
|
||||
|
||||
Different results may use different fragment shapes.
|
||||
|
||||
## Graph consumers
|
||||
|
||||
A graph consumer of a batch result may:
|
||||
|
||||
1. consume fragments directly as physical fragments;
|
||||
2. select one or more physical slots in a `spat.deferred_communication` body;
|
||||
3. use a Blueprint to obtain or describe a logical reconstruction;
|
||||
4. feed fragments to following graph computes or graph compute batches.
|
||||
|
||||
A consumer must not treat the leading physical slot dimension as a logical
|
||||
model dimension unless an explicit graph operation intentionally performs such
|
||||
an interpretation.
|
||||
|
||||
All constant selection, slicing, reshaping, concatenation, and other
|
||||
compile-time shaping needed for a scheduled consumer must be encoded inside the
|
||||
corresponding `spat.deferred_communication` body. Phase 2 must not recover
|
||||
missing graph semantics by inspecting consumers after the deferred operation.
|
||||
|
||||
### Consumer and producer lanes are many-to-many
|
||||
|
||||
A graph consumer lane may read fragments from several producer lanes. Several
|
||||
consumer lanes may also read the same producer lane. Therefore:
|
||||
|
||||
- consumer and producer batches need not have the same lane count;
|
||||
- consumer lane `i` must not be assumed to depend on producer lane `i`;
|
||||
- scheduling dependencies must follow the statically analyzable physical-slot
|
||||
projections expressed by the consumer body;
|
||||
- an unproven projection must not silently drop producer dependencies or
|
||||
reclassify a produced graph value as a host input.
|
||||
|
||||
For example, one output-row lane of a `3x3` convolution normally reads up to
|
||||
three input-row lanes. Output-channel tiling adds another projection dimension:
|
||||
several output-channel-tile lanes may read the same set of input-row lanes.
|
||||
|
||||
When exact projected dependencies cannot be proven, scheduling may conservatively
|
||||
depend on the complete producer batch, but communication realization must still
|
||||
preserve the consumer body's actual selection semantics.
|
||||
|
||||
## Graph lane, scheduled lane, and physical core are different identities
|
||||
|
||||
These concepts must never be conflated:
|
||||
|
||||
- **graph lane**: the lane of the original `spat.graph_compute_batch`;
|
||||
- **physical fragment slot**: the slot in the graph batch result;
|
||||
- **scheduled lane**: one lane of a `spat.scheduled_compute_batch` equivalence
|
||||
class;
|
||||
- **physical core**: the core selected by PEFT.
|
||||
|
||||
The graph batch body or its Blueprint defines graph-lane-to-fragment-slot and
|
||||
fragment-slot-to-logical-region mappings.
|
||||
|
||||
PEFT defines graph-instance-to-core placement.
|
||||
|
||||
Scheduled communication defines how values move between cores.
|
||||
|
||||
## Scheduled IR exclusion
|
||||
|
||||
Do not add a verifier requiring `spat.scheduled_compute_batch` results to have
|
||||
`laneCount` as their first dimension.
|
||||
|
||||
Do not rewrite scheduled values merely to resemble graph physical fragment
|
||||
collections.
|
||||
|
||||
When lowering graph IR into scheduled IR:
|
||||
|
||||
- resolve graph fragments and reconstruction metadata before erasing their
|
||||
graph owners;
|
||||
- create local forwarding or `spat.channel_send`/`spat.channel_receive` for
|
||||
cross-core dependencies;
|
||||
- allow scheduled result representation to follow the scheduled IR contract;
|
||||
- preserve numerical and deadlock correctness.
|
||||
|
||||
The graph invariant is an input contract for scheduling, not a scheduled-value
|
||||
layout contract.
|
||||
|
||||
## Required verifier properties
|
||||
|
||||
`spat.graph_compute_batch` verification must establish, for every result:
|
||||
|
||||
1. the result is a static or otherwise supported ranked tensor;
|
||||
2. result rank is exactly `fragment rank + 1`;
|
||||
3. result dimension 0 equals `laneCount`;
|
||||
4. every lane publication source has the same exact fragment type;
|
||||
5. the physical insertion targets the corresponding result block argument;
|
||||
6. physical insertion offsets have the fragment slot in dimension 0;
|
||||
7. all remaining physical offsets are zero;
|
||||
8. physical sizes are `[1] + fragment shape`;
|
||||
9. physical strides are unit;
|
||||
10. exactly one publication is defined for each graph result in the per-lane
|
||||
body.
|
||||
|
||||
These checks apply only to `spat.graph_compute_batch`, not to
|
||||
`spat.scheduled_compute_batch`.
|
||||
|
||||
Blueprint verification must establish that every logical reconstruction entry:
|
||||
|
||||
- references an existing physical batch operand;
|
||||
- references a valid physical fragment slot;
|
||||
- maps a fragment compatible with the declared logical slice;
|
||||
- stays within logical bounds;
|
||||
- follows the declared conflict and coverage policies.
|
||||
|
||||
## Invalid representations
|
||||
|
||||
The following are invariant violations.
|
||||
|
||||
### Logical aggregate returned directly by graph batch
|
||||
|
||||
```text
|
||||
laneCount = 16
|
||||
result = tensor<1x4x16x16>
|
||||
```
|
||||
|
||||
with each lane inserting into logical dimension 2.
|
||||
|
||||
This is a logical assembly masquerading as a graph batch result. The graph
|
||||
result must instead be `16 × per-row-fragment`, and a Blueprint must describe
|
||||
placement into `tensor<1x4x16x16>`.
|
||||
|
||||
### Physical storage derived from logical destination shape
|
||||
|
||||
Code equivalent to:
|
||||
|
||||
```cpp
|
||||
shape = logicalDestinationType.getShape();
|
||||
shape[logicalInsertionDimension] = laneCount;
|
||||
```
|
||||
|
||||
is invalid.
|
||||
|
||||
Physical graph storage must be derived from the per-lane fragment type:
|
||||
|
||||
```cpp
|
||||
physicalShape = [laneCount] + fragmentType.getShape();
|
||||
```
|
||||
|
||||
### Reconstruction inferred from result type
|
||||
|
||||
It is invalid to assume that physical slot `i` belongs at logical offset `i`.
|
||||
The Blueprint or another explicit reconstruction descriptor must state the
|
||||
mapping.
|
||||
|
||||
### Blueprint used for arithmetic
|
||||
|
||||
It is invalid to encode `fragment0 + fragment1` as Blueprint reconstruction.
|
||||
Create a following graph compute or graph compute batch for the addition.
|
||||
|
||||
## Ownership
|
||||
|
||||
- ONNX-to-Spatial lowering owns creation of valid graph fragment batches.
|
||||
- Graph canonicalization owns normalization of temporary logical-assembly forms
|
||||
into physical graph batches plus Blueprints.
|
||||
- `spat.graph_compute_batch` verifier rejects invalid physical publications.
|
||||
- `spat.blueprint` owns persistent logical reconstruction metadata.
|
||||
- Deferred communication Phase 1 owns complete consumer-side constant shaping.
|
||||
- Merge scheduling consumes this graph contract and introduces explicit
|
||||
communication.
|
||||
- Scheduled IR verifiers validate scheduled execution and communication, not
|
||||
the graph fragment representation.
|
||||
|
||||
## No repair downstream
|
||||
|
||||
If graph IR violates this invariant, fix the graph producer or graph
|
||||
canonicalization.
|
||||
|
||||
Do not repair an invalid graph batch by:
|
||||
|
||||
- guessing a lane dimension in `MergeComputeNodes`;
|
||||
- deriving physical storage from a logical destination tensor;
|
||||
- inspecting deferred-result users;
|
||||
- reconstructing omitted Blueprint data after graph erasure;
|
||||
- weakening graph verifiers;
|
||||
- imposing the graph representation on scheduled operations.
|
||||
@@ -0,0 +1,105 @@
|
||||
# Compiler Performance Optimization Invariant
|
||||
|
||||
## Scope
|
||||
|
||||
This invariant applies to:
|
||||
|
||||
- ONNX-to-Spatial lowering;
|
||||
- graph and trivial merge;
|
||||
- PEFT scheduling;
|
||||
- scheduled materialization;
|
||||
- deferred communication;
|
||||
- Spatial-to-PIM lowering;
|
||||
- bufferization;
|
||||
- liveness planning;
|
||||
- PIM code generation.
|
||||
|
||||
## Invariant
|
||||
|
||||
A compiler compile-time or memory optimization must not reduce available
|
||||
hardware parallelism or worsen theoretical execution time. In the absence of a
|
||||
precise runtime model, uncertainty is not permission to accept a possible
|
||||
runtime regression. It is forbidden to serialize independent work, reduce
|
||||
the number of simultaneously schedulable compute instances, replace parallel
|
||||
graph or core batches with sequential loops, introduce recomputation, add
|
||||
runtime copies, increase communication or instruction count, or lengthen the
|
||||
schedule critical path merely to reduce compiler time or memory usage.
|
||||
|
||||
Compiler representations may share analysis, planning, or code-generation
|
||||
work only when the emitted execution retains the same independent work,
|
||||
schedule semantics, scheduling granularity, and available parallelism.
|
||||
|
||||
Static runtime proxies are useful for screening changes, but they are not the
|
||||
latency authority. When comparable validation results are available for the
|
||||
same model, inputs, hardware configuration, and non-functional simulation
|
||||
configuration, the `Latency` reported by `validation/validate.py` is
|
||||
authoritative. An optimization is acceptable only when its static runtime
|
||||
proxies are equal or better, unless such a matched validation comparison
|
||||
demonstrates equal or better latency.
|
||||
|
||||
## Asymptotic cost
|
||||
|
||||
For each new or materially changed compiler algorithm, identify the relevant
|
||||
input size and target linear or sublinear time and space. Avoid repeated full-IR
|
||||
walks, nested scans, and per-operation recomputation when indexing, caching, or
|
||||
a single traversal can express the same behavior.
|
||||
|
||||
Bufferization cost scales with both the number of operations and the number of
|
||||
MLIR values it receives. Upstream lowering and scheduling must therefore keep
|
||||
repeated work in compact structured operations, such as statically evaluable
|
||||
loops and batches, until after bufferization. Do not compensate for avoidable
|
||||
pre-bufferization expansion by weakening, partitioning, or special-casing the
|
||||
bufferization analysis; measure both operation and value counts at its input.
|
||||
Compactness must also be preserved after bufferization so that liveness,
|
||||
verification, memory planning, and other downstream passes do not repeat work
|
||||
per logical lane or iteration. Keep structured loops and batches intact until
|
||||
PIM ISA code generation is forced to scalarize them into concrete per-core
|
||||
instructions; earlier expansion requires an explicit semantic necessity and
|
||||
before/after operation and value counts.
|
||||
|
||||
When linear-or-better complexity is not possible, use the lowest justified
|
||||
complexity and report the actual time and space Big-O, the input variable, and
|
||||
why a lower bound is not practical. Include that cost in the final report; do
|
||||
not hide a superlinear path behind small current test sizes.
|
||||
|
||||
## Forbidden optimization trades
|
||||
|
||||
Do not:
|
||||
|
||||
- use fewer parallel lanes or cores to simplify IR;
|
||||
- coarsen lane, core, batch, or scheduling granularity when it may reduce
|
||||
concurrency;
|
||||
- scalarize a valid graph or core batch;
|
||||
- replace independent operations with one sequential loop;
|
||||
- recompute values to reduce storage;
|
||||
- add device/device or host/device copies to reduce peak memory;
|
||||
- change hardware parameters to make a benchmark easier;
|
||||
- reduce communication concurrency by imposing a global serial order;
|
||||
- increase emitted instruction count to reduce compiler analysis;
|
||||
- move PIM work to the host.
|
||||
|
||||
## Required proof
|
||||
|
||||
Every performance change must report before and after:
|
||||
|
||||
- matched validation `Latency` when the changed execution can be simulated;
|
||||
- compiler wall time;
|
||||
- compiler peak RSS;
|
||||
- IR operation and value counts at the changed stage;
|
||||
- logical compute-instance count;
|
||||
- graph compute batch and scheduled batch lane counts;
|
||||
- PEFT class and core assignment;
|
||||
- maximum scheduled critical-path or step count;
|
||||
- MVM/VMM and vector instruction counts;
|
||||
- send and receive counts and bytes;
|
||||
- total emitted instruction count;
|
||||
- maximum instructions assigned to one core;
|
||||
- number and size of runtime memory copies;
|
||||
- final host, weights, logical-core, physical-core, and maximum-core memory;
|
||||
- numerical validation.
|
||||
|
||||
## Stop rule
|
||||
|
||||
If compile time or memory improves but a runtime proxy worsens, stop and reject
|
||||
the change unless a matched validation comparison demonstrates equal or better
|
||||
`Latency`. Proxy improvements do not override a validation latency regression.
|
||||
@@ -0,0 +1,73 @@
|
||||
# Spatial Target Generality Invariant
|
||||
|
||||
## Scope
|
||||
|
||||
This invariant applies to:
|
||||
|
||||
- the Spatial dialect and its verifiers;
|
||||
- ONNX-to-Spatial planning and lowering;
|
||||
- graph transforms and scheduling over Spatial IR;
|
||||
- target information consumed while optimizing Spatial IR;
|
||||
- lowerings from Spatial to target-specific dialects.
|
||||
|
||||
## Invariant
|
||||
|
||||
Spatial represents logical compute, dataflow, layout choices, parallel work,
|
||||
and target-independent resource requirements. It must remain usable by
|
||||
different targets, including PIM and future targets such as PULPIM.
|
||||
|
||||
Raptor may ingest target information and use it to choose Spatial layouts,
|
||||
partitions, placements, or schedules. That information must cross an explicit
|
||||
target interface and be expressed in target-neutral terms at the Spatial
|
||||
layer. The Spatial dialect and scheduler must not parse a simulator-specific
|
||||
configuration, depend on a target dialect, or encode one target's instruction
|
||||
latencies, memory hierarchy, communication protocol, or resource policy.
|
||||
|
||||
Target adapters own translating a target configuration into the neutral
|
||||
information consumed by Spatial. Target-specific dialects and their lowerings
|
||||
own instruction semantics, physical memory details, communication mechanisms,
|
||||
and final legality.
|
||||
|
||||
## Ownership boundary
|
||||
|
||||
- Spatial IR owns logical and physical planning concepts shared across targets.
|
||||
- A target adapter owns configuration parsing and cost-model construction.
|
||||
- The scheduler consumes an injected cost/resource model; it does not infer a
|
||||
target from global PIM options or hardcoded constants.
|
||||
- Spatial-to-target lowering makes the selected representation explicit.
|
||||
- Target dialect verifiers reject target-specific illegal states.
|
||||
|
||||
Target-neutral information may include processor topology, operation and
|
||||
transfer cost queries, available parallel capacity, and opaque resource
|
||||
requirements. Names and APIs at this boundary must describe those concepts,
|
||||
not a particular simulator or target implementation.
|
||||
|
||||
## Forbidden coupling
|
||||
|
||||
Do not:
|
||||
|
||||
- include PIM or PULPIM dialect headers in the Spatial dialect or scheduler;
|
||||
- read PIM compiler globals directly from generic scheduling algorithms;
|
||||
- parse `pimsim-nn`, PIMCOMP, or another simulator's schema in Spatial code;
|
||||
- hardcode Arch-A timing, mesh, crossbar, memory, or vector constants in
|
||||
Spatial cost calculations;
|
||||
- add target-named Spatial operations when an existing logical/layout concept
|
||||
expresses the invariant;
|
||||
- repair target-specific legality in generic Spatial cleanup passes.
|
||||
|
||||
## Required proof
|
||||
|
||||
Changes that use target information in Spatial must show:
|
||||
|
||||
- the target boundary or injected interface used;
|
||||
- that Spatial IR remains valid without target-specific attributes;
|
||||
- that an unknown or unsupported target fails clearly rather than silently
|
||||
using PIM defaults;
|
||||
- unit coverage with at least two distinct target profiles when scheduling or
|
||||
cost decisions change;
|
||||
- target-specific validation after lowering for every implemented target
|
||||
affected by the change.
|
||||
|
||||
If only one target implementation exists, keep the interface narrow and test
|
||||
it with two profiles. Do not add speculative target operations or a framework
|
||||
for unimplemented targets.
|
||||
+6
-5
@@ -4,12 +4,13 @@
|
||||
|
||||
.claude
|
||||
.codex
|
||||
AGENTS.md
|
||||
|
||||
CMakeUserPresets.json
|
||||
|
||||
build
|
||||
cmake-build-debug
|
||||
cmake-build-release
|
||||
build_*
|
||||
compile.sh
|
||||
pimcomp_utils/*
|
||||
|
||||
**/__*
|
||||
*.zip
|
||||
**/*.zip
|
||||
**/__pycache__/
|
||||
|
||||
+4
-1
@@ -3,4 +3,7 @@
|
||||
url = https://github.com/onnx/onnx-mlir.git
|
||||
[submodule "backend-simulators/pim/pimsim-nn"]
|
||||
path = backend-simulators/pim/pimsim-nn
|
||||
url = https://github.com/wangxy-2000/pimsim-nn.git
|
||||
url = https://github.com/HEAPLab/pimsim-nn.git
|
||||
[submodule "third_party/PIMCOMP-NN"]
|
||||
path = third_party/PIMCOMP-NN
|
||||
url = https://github.com/HEAPLab/PIMCOMP-NN.git
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
* Always read the full README.md before doing anything
|
||||
|
||||
# Required project invariants
|
||||
|
||||
Before modifying the relevant subsystem, read:
|
||||
|
||||
* `.agents/invariants/GRAPH_COMPUTE_BATCH_INVARIANT.md`
|
||||
* `.agents/invariants/PERFORMANCE_OPTIMIZATION_INVARIANT.md`
|
||||
* `.agents/invariants/SPATIAL_TARGET_GENERALITY_INVARIANT.md`
|
||||
* Build commands:
|
||||
* `cmake --build ./build_release`
|
||||
* `cmake --build ./build_debug`
|
||||
* Never use `ninja` directly: it bypasses cmake's configuration and invalidates the build cache
|
||||
* Always try the release build first before building with the debug version
|
||||
* Use the debug build only when it is useful to obtain a clear stack trace with symbols, inspect names, place breakpoints, or test a small case interactively
|
||||
* The debug build is very slow, so use it only on small fast tests such as operation validations, not on network validations
|
||||
* Always prepend rtk to shell commands if missing and if rtk is available
|
||||
* Always use the repository Python virtual environment (`.venv/bin/python` and its bundled tools) for Python commands
|
||||
* Run the complete operations validation (`--operations-dir validation/operations`) outside the sandbox so parallel worker processes can start. This does not apply to one-at-a-time network validation
|
||||
|
||||
# Core engineering philosophy
|
||||
|
||||
* Clean architecture matters as much as making the immediate test pass
|
||||
* Prefer fixes that preserve clear ownership boundaries, explicit invariants, and simple dataflow
|
||||
* Do not stack compensating fixes on top of earlier mistakes. If the current approach is becoming messy, stop and explain why
|
||||
* A correct fix should usually make the responsible producer, resolver, verifier, or lowering own the behavior directly
|
||||
* Avoid late repair passes, defensive cleanup, or broad rewrites when a cleaner owner-side fix is possible
|
||||
* Do not hide an upstream modeling bug by normalizing it later in the pipeline. Fix the producer when the producer owns the invariant
|
||||
* Prefer patterns/rewrites for local IR canonicalization. Use module walks only when pass-level structural analysis genuinely requires them
|
||||
* Prefer compact, structured designs over long case-by-case implementations
|
||||
|
||||
# Think before coding
|
||||
|
||||
* State assumptions explicitly before implementing when they affect the design
|
||||
* If multiple interpretations exist, present them instead of silently choosing one
|
||||
* If a simpler approach exists, say so and prefer it unless there is a clear reason not to
|
||||
* If something is unclear, stop, name what is confusing, and ask
|
||||
* If the requested or obvious approach would make the architecture worse, push back and propose a cleaner alternative
|
||||
|
||||
# Code changes
|
||||
|
||||
* Keep changes minimal and localized to the relevant parts of the code
|
||||
* Preserve the existing naming conventions and coding style used in the surrounding code
|
||||
* Keep code easy to read, well organized, and suitable for future extensibility
|
||||
* A function must not exceed roughly 200/250 lines. If a change pushes a function beyond that, extract focused helpers
|
||||
* Prefer clear naming and structure over comments. Add comments only when they materially improve clarity
|
||||
* Do not rename symbols, move files, or restructure modules unless that is necessary for the requested change
|
||||
* Avoid duplicate ad-hoc logic. If the same concept appears in multiple places, consider whether it deserves a shared helper/API
|
||||
* When adding a helper or API, ask:
|
||||
* Could this be useful to another component now
|
||||
* Is another component already implementing the same idea differently
|
||||
* Is this likely to be needed by a future adjacent component
|
||||
* What is the narrowest useful abstraction
|
||||
* What is the correct ownership level for this API
|
||||
* If a shared API is justified, place it at the lowest clean layer that can be used by all relevant consumers without creating dependency cycles or leaking policy across layers
|
||||
* If an existing component should use a newly introduced shared API, refactor that component in the same patch when doing so is directly related and reduces duplication
|
||||
* Do not create broad frameworks just because a helper might someday be useful. Shared APIs should encode a real reusable concept, not speculative generality
|
||||
* If the reusable abstraction is plausible but not clearly needed yet, keep the code local and mention the possible future extraction separately
|
||||
|
||||
# Avoid case-listing designs
|
||||
|
||||
* Avoid solving problems with large chains of `if`/`else`, switches, or repeated special cases that enumerate every possible situation
|
||||
* Long case listings tend to overfit the current tests, grow the codebase, and hide the underlying abstraction
|
||||
* When you see a growing list of special cases, stop and look for the shared concept, data model, interface, or normalization step that would make the cases collapse
|
||||
* Prefer table-driven logic, traits/interfaces, small reusable predicates, structured dispatch, or producer-side normalization when they express the invariant more directly
|
||||
* A few explicit cases are fine when the domain is genuinely small and closed
|
||||
* If the list is likely to grow, refactor toward a cleaner and more compact design instead of adding another branch
|
||||
* When keeping a case list is the pragmatic choice, explain why the domain is closed or why a broader abstraction would be premature
|
||||
|
||||
# Ownership and invariants
|
||||
|
||||
Before implementing, identify the owner of the behavior:
|
||||
|
||||
* A producer should emit IR/data that satisfies the contract of the next stage
|
||||
* A lowering should make representation changes explicit and semantically correct
|
||||
* A resolver should resolve existing structure without silently changing semantics
|
||||
* A verifier should reject invalid states with bounded, actionable diagnostics
|
||||
* Codegen should assume verified invariants and fail clearly if they are violated
|
||||
|
||||
When fixing a bug:
|
||||
|
||||
* State the invariant that was violated
|
||||
* State which component should own that invariant
|
||||
* Fix that component directly
|
||||
* Avoid fixes that merely mask the violation later in the pipeline
|
||||
* Add or preserve verification if the invariant is important enough to regress
|
||||
|
||||
# Refactor and API policy
|
||||
|
||||
You may propose or implement a refactor when:
|
||||
|
||||
* the local fix would duplicate logic
|
||||
* the local fix would violate a layer boundary
|
||||
* the bug exists because responsibility is assigned to the wrong component
|
||||
* multiple components already implement ad-hoc variants of the same concept
|
||||
* a shared helper/API would make the code smaller, clearer, and easier to maintain
|
||||
* existing callers can be migrated cleanly without broad churn
|
||||
* the current implementation is turning into a long list of special cases instead of a structured solution
|
||||
|
||||
When proposing or implementing a refactor:
|
||||
|
||||
* Explain what responsibility is being moved or shared
|
||||
* Justify why the new location is the right ownership level
|
||||
* Keep the API narrow and named after the concept or invariant it represents
|
||||
* Migrate directly related existing users when that improves compactness and consistency
|
||||
* Separate changes required for correctness from optional cleanup
|
||||
* Avoid unrelated renames, formatting changes, or module moves
|
||||
* Do not expand a justified refactor beyond directly related callers
|
||||
|
||||
Do not refactor when:
|
||||
|
||||
* the issue is truly local and a local fix is clearer
|
||||
* the abstraction would have only one user and no clear adjacent use
|
||||
* the abstraction would mix policies from different layers
|
||||
* the refactor would affect unrelated behavior
|
||||
* the refactor is mainly aesthetic
|
||||
|
||||
# Working style
|
||||
|
||||
* Infer style and conventions from the existing code before introducing new patterns
|
||||
* When several implementation options are possible, prefer the simplest one that fits the current architecture and minimizes churn
|
||||
* Push back when the requested or obvious fix would make the architecture worse
|
||||
* If a cleaner fix requires a small refactor or shared helper/API, propose it explicitly and justify it
|
||||
* Avoid broad refactors unless explicitly requested or clearly necessary for correctness and maintainability
|
||||
* When tests fail, bucket failures by likely root cause and separate patch-related failures from pre-existing or out-of-scope failures
|
||||
|
||||
# Simplicity first
|
||||
|
||||
* Minimum code that solves the problem cleanly. Nothing speculative
|
||||
* No features beyond what was asked
|
||||
* No error handling for impossible scenarios
|
||||
* If you write 200 lines and it could be 50, rewrite it
|
||||
* Ask: “Would a senior engineer say this is overcomplicated?” If yes, simplify
|
||||
* Prefer direct, explicit code over generic machinery unless the generic machinery clearly reduces duplication and preserves boundaries
|
||||
|
||||
# Fallbacks and defaults
|
||||
|
||||
* Avoid silent fallback behavior when the semantic category is unknown
|
||||
* Do not treat “unknown” as “safe” unless the codebase already defines that convention
|
||||
* If a value cannot be classified, either preserve the existing behavior deliberately or fail with a clear diagnostic
|
||||
* When adding a fallback, state why it is semantically valid and what invariant makes it safe
|
||||
|
||||
# Surgical changes
|
||||
|
||||
* Touch only what you must
|
||||
* Clean up only the mess introduced by your own change
|
||||
* Do not “improve” adjacent code, comments, or formatting
|
||||
* Match existing style, even if you would personally do it differently
|
||||
* If you notice unrelated dead code, bad abstractions, or fragile design, mention it separately. Do not delete or rewrite it unless asked
|
||||
* When your changes create orphans, remove imports, variables, functions, or files made unused by your change
|
||||
* Every changed line should trace directly to the requested fix, a required cleanup, or a justified reuse/refactor decision
|
||||
|
||||
# Diagnostics and verification
|
||||
|
||||
* Use existing bounded diagnostic mechanisms for pass-level verification or codegen failures
|
||||
* Do not emit unbounded repeated diagnostics from loops or parallel workers
|
||||
* Diagnostics should identify the violated invariant and the relevant value/op when useful
|
||||
* Verifiers should reject invalid states, not repair them
|
||||
* Codegen should not compensate for invalid IR/data unless codegen is the owner of that invariant
|
||||
* Do not make failing tests pass by weakening verifiers, assertions, or diagnostics unless the check itself is proven wrong
|
||||
* If a check is too strict, explain the valid case it rejects and update the invariant accordingly
|
||||
* Prefer fixing invalid IR/data producers over relaxing consumers
|
||||
* If adding diagnostics only for debugging, remove them or cap them before finalizing
|
||||
|
||||
# Temporary debugging code
|
||||
|
||||
* Temporary diagnostics, dumps, assertions, and debug-only helpers must be removed or intentionally converted into bounded permanent diagnostics before finalizing
|
||||
* If debug instrumentation remains, explain why it is useful as permanent infrastructure
|
||||
* Do not leave noisy validation output behind
|
||||
|
||||
# Performance awareness
|
||||
|
||||
* Avoid algorithmic regressions in compiler passes, especially repeated full-module walks, repeated expensive analyses, or per-op recomputation inside nested loops
|
||||
* If a change adds a walk, cache, analysis, or structural traversal, justify why it is needed
|
||||
* For hot paths, prefer preserving existing asymptotic behavior unless a better structure is part of the requested change
|
||||
* If performance may change, mention the expected impact and suggest a targeted timing check
|
||||
* Compile-time and compiler-memory improvements must not trade away hardware
|
||||
parallelism or theoretical runtime. Preserve or improve the static runtime
|
||||
proxies defined in
|
||||
`.agents/invariants/PERFORMANCE_OPTIMIZATION_INVARIANT.md`.
|
||||
|
||||
# Goal-driven execution
|
||||
|
||||
For multi-step tasks, state a brief plan:
|
||||
|
||||
1. [Step] → verify: [check]
|
||||
2. [Step] → verify: [check]
|
||||
3. [Step] → verify: [check]
|
||||
|
||||
Define success criteria before implementing:
|
||||
|
||||
* For bug fixes, success means reproducing or identifying the failure, fixing the responsible owner, and verifying the targeted case
|
||||
* For refactors, success means preserving behavior while making ownership, reuse, or structure cleaner
|
||||
* For validation changes, success means checking both valid and invalid cases when applicable
|
||||
|
||||
Transform tasks into verifiable goals:
|
||||
|
||||
* “Fix the bug” → identify the invariant, reproduce the failure, fix the owner, verify the targeted case
|
||||
* “Add validation” → write or identify tests for invalid inputs, then make them pass/fail as expected
|
||||
* “Refactor X” → preserve behavior before and after, then run relevant tests
|
||||
|
||||
# Final self-review
|
||||
|
||||
Before reporting completion, check:
|
||||
|
||||
* Did I fix the owner of the invariant rather than masking the issue downstream
|
||||
* Did I avoid broad case lists and ad-hoc special handling
|
||||
* Did I introduce a helper/API only at the right ownership level
|
||||
* Did I migrate directly related duplicate logic when doing so improves compactness
|
||||
* Did I avoid weakening verifiers or assertions unnecessarily
|
||||
* Did I remove temporary debugging code or make it bounded and intentional
|
||||
* Did I avoid unrelated formatting, renames, or cleanup
|
||||
* Did I consider performance impact for added walks, analyses, caches, or repeated computations
|
||||
* Did I run the required build/test commands
|
||||
* Did I clearly report remaining failures or risks
|
||||
|
||||
When reporting back:
|
||||
|
||||
* Say what changed
|
||||
* Say what was verified
|
||||
* Say what remains
|
||||
* When showing code in chat, make it easy to copy-paste into the codebase
|
||||
* Keep outputs focused on the changed parts
|
||||
* List bad practices, fragile assumptions, or cleaner alternatives separately
|
||||
* If a change is intentionally pragmatic rather than architecturally ideal, say so and explain the tradeoff
|
||||
+113
-25
@@ -1,33 +1,101 @@
|
||||
# Match the minimum required version of onnx-mlir
|
||||
cmake_minimum_required(VERSION 3.20.0)
|
||||
|
||||
project(raptor)
|
||||
project(Raptor)
|
||||
|
||||
# Add symlink to PIM as accelerator in onnx-mlir
|
||||
function(raptor_ensure_symlink link_path target_path)
|
||||
get_filename_component(link_parent "${link_path}" DIRECTORY)
|
||||
# Materialize a CMake shim directory
|
||||
function(raptor_write_external_cmake_shim shim_dir external_source_dir description)
|
||||
get_filename_component(real_external_source_dir "${external_source_dir}" REALPATH)
|
||||
file(RELATIVE_PATH relative_external_source_dir "${shim_dir}" "${real_external_source_dir}")
|
||||
|
||||
if(NOT EXISTS "${link_parent}")
|
||||
message(FATAL_ERROR "Directory not found: ${link_parent}")
|
||||
endif()
|
||||
|
||||
if(NOT EXISTS "${link_path}")
|
||||
message(STATUS "Creating symlink ${link_path} -> ${target_path}")
|
||||
file(CREATE_LINK
|
||||
"${target_path}"
|
||||
"${link_path}"
|
||||
SYMBOLIC
|
||||
if (NOT EXISTS "${real_external_source_dir}/CMakeLists.txt")
|
||||
message(FATAL_ERROR
|
||||
"External CMake source directory not found or missing CMakeLists.txt:\n"
|
||||
" ${real_external_source_dir}"
|
||||
)
|
||||
endif()
|
||||
endif ()
|
||||
|
||||
if (IS_SYMLINK "${shim_dir}")
|
||||
message(STATUS "Removing old full-directory symlink: ${shim_dir}")
|
||||
file(REMOVE "${shim_dir}")
|
||||
endif ()
|
||||
|
||||
if (EXISTS "${shim_dir}" AND NOT IS_DIRECTORY "${shim_dir}")
|
||||
message(FATAL_ERROR "Expected directory or absent path, got file: ${shim_dir}")
|
||||
endif ()
|
||||
|
||||
file(MAKE_DIRECTORY "${shim_dir}")
|
||||
|
||||
set(shim_file "${shim_dir}/CMakeLists.txt")
|
||||
set(shim_contents
|
||||
"get_filename_component(raptor_external_source_dir
|
||||
\"\${CMAKE_CURRENT_LIST_DIR}/${relative_external_source_dir}\"
|
||||
REALPATH
|
||||
)
|
||||
add_subdirectory(
|
||||
\"\${raptor_external_source_dir}\"
|
||||
\"\${CMAKE_CURRENT_BINARY_DIR}/raptor-external\"
|
||||
)
|
||||
if (DEFINED PIM_ENABLED)
|
||||
set(PIM_ENABLED \"\${PIM_ENABLED}\" PARENT_SCOPE)
|
||||
endif ()
|
||||
"
|
||||
)
|
||||
|
||||
if (EXISTS "${shim_file}")
|
||||
file(READ "${shim_file}" old_contents)
|
||||
else ()
|
||||
set(old_contents "")
|
||||
endif ()
|
||||
|
||||
if (NOT old_contents STREQUAL shim_contents)
|
||||
file(WRITE "${shim_file}" "${shim_contents}")
|
||||
message(STATUS "Wrote CMake shim for ${description}: ${shim_file}")
|
||||
else ()
|
||||
message(STATUS "CMake shim already up to date for ${description}")
|
||||
endif ()
|
||||
|
||||
# Mirror the external tree's first-level entries into the shim directory
|
||||
# so legacy includes like src/Accelerators/PIM/Compiler/... keep working.
|
||||
file(GLOB children RELATIVE "${real_external_source_dir}" "${real_external_source_dir}/*")
|
||||
|
||||
foreach (child IN LISTS children)
|
||||
if (child STREQUAL "CMakeLists.txt")
|
||||
continue()
|
||||
endif ()
|
||||
|
||||
set(real_child "${real_external_source_dir}/${child}")
|
||||
set(shim_child "${shim_dir}/${child}")
|
||||
|
||||
if (IS_SYMLINK "${shim_child}")
|
||||
file(READ_SYMLINK "${shim_child}" existing_link_target)
|
||||
if (existing_link_target STREQUAL real_child)
|
||||
continue()
|
||||
endif ()
|
||||
file(REMOVE_RECURSE "${shim_child}")
|
||||
elseif (EXISTS "${shim_child}")
|
||||
# Do not delete real files/directories. This protects the generated shim.
|
||||
continue()
|
||||
endif ()
|
||||
|
||||
file(CREATE_LINK
|
||||
"${real_child}"
|
||||
"${shim_child}"
|
||||
SYMBOLIC
|
||||
)
|
||||
endforeach ()
|
||||
endfunction()
|
||||
|
||||
raptor_ensure_symlink(
|
||||
"${CMAKE_CURRENT_SOURCE_DIR}/onnx-mlir/src/Accelerators/PIM"
|
||||
"${CMAKE_CURRENT_SOURCE_DIR}/src/PIM"
|
||||
raptor_write_external_cmake_shim(
|
||||
"${CMAKE_CURRENT_SOURCE_DIR}/onnx-mlir/src/Accelerators/PIM"
|
||||
"${CMAKE_CURRENT_SOURCE_DIR}/src/PIM"
|
||||
"PIM accelerator"
|
||||
)
|
||||
raptor_ensure_symlink(
|
||||
"${CMAKE_CURRENT_SOURCE_DIR}/onnx-mlir/test/accelerators/PIM"
|
||||
"${CMAKE_CURRENT_SOURCE_DIR}/test/PIM"
|
||||
|
||||
raptor_write_external_cmake_shim(
|
||||
"${CMAKE_CURRENT_SOURCE_DIR}/onnx-mlir/test/accelerators/PIM"
|
||||
"${CMAKE_CURRENT_SOURCE_DIR}/test/PIM"
|
||||
"PIM accelerator tests"
|
||||
)
|
||||
|
||||
# Patch onnx-mlir sources for PIM accelerator support.
|
||||
@@ -38,21 +106,21 @@ function(raptor_apply_patch file_path anchor replacement description)
|
||||
|
||||
# Already applied – replacement text is present
|
||||
string(FIND "${contents}" "${replacement}" already_applied_pos)
|
||||
if(NOT already_applied_pos EQUAL -1)
|
||||
if (NOT already_applied_pos EQUAL -1)
|
||||
message(STATUS "Patch already applied: ${description}")
|
||||
return()
|
||||
endif()
|
||||
endif ()
|
||||
|
||||
# Anchor must exist for the patch to be applicable
|
||||
string(FIND "${contents}" "${anchor}" anchor_pos)
|
||||
if(anchor_pos EQUAL -1)
|
||||
if (anchor_pos EQUAL -1)
|
||||
message(FATAL_ERROR
|
||||
"Patch anchor not found – onnx-mlir may have changed.\n"
|
||||
" Patch : ${description}\n"
|
||||
" File : ${file_path}\n"
|
||||
" Anchor: ${anchor}"
|
||||
)
|
||||
endif()
|
||||
endif ()
|
||||
|
||||
string(REPLACE "${anchor}" "${replacement}" patched "${contents}")
|
||||
file(WRITE "${file_path}" "${patched}")
|
||||
@@ -77,4 +145,24 @@ raptor_apply_patch(
|
||||
"Skip output emission for PIM accelerator"
|
||||
)
|
||||
|
||||
# Count the PIM status messages instead of the unused standard backend phases
|
||||
raptor_apply_patch(
|
||||
"${ONNX_MLIR_DIR}/src/Compiler/CompilerUtils.cpp"
|
||||
"void showCompilePhase(std::string msg) {"
|
||||
[=[void showCompilePhase(std::string msg) {
|
||||
if (llvm::is_contained(maccel, accel::Accelerator::Kind::PIM)) {
|
||||
if (pimOnlyCodegen)
|
||||
TOTAL_COMPILE_PHASE = 2;
|
||||
else if (pimEmissionTarget == EmitSpatial)
|
||||
TOTAL_COMPILE_PHASE = 3;
|
||||
else if (pimEmissionTarget == EmitPim)
|
||||
TOTAL_COMPILE_PHASE = 4;
|
||||
else if (pimEmissionTarget == EmitPimBufferized)
|
||||
TOTAL_COMPILE_PHASE = 5;
|
||||
else
|
||||
TOTAL_COMPILE_PHASE = 9;
|
||||
}]=]
|
||||
"Use PIM compile phase count"
|
||||
)
|
||||
|
||||
add_subdirectory(onnx-mlir)
|
||||
|
||||
@@ -1,219 +1,289 @@
|
||||
# Raptor
|
||||
|
||||
Raptor is a domain-specific MLIR compiler for neural networks (ONNX format)
|
||||
targeting in-memory computing / processing-in-memory (PIM) architectures.
|
||||
It progressively lowers ONNX-MLIR through a set of MLIR dialects down to
|
||||
target-specific artifacts (currently JSON code for the `pimsim-nn` simulator).
|
||||
Raptor is a domain-specific MLIR compiler for neural networks in ONNX format,
|
||||
targeting in-memory computing / processing-in-memory (PIM) architectures. It
|
||||
extends ONNX-MLIR with a PIM accelerator and progressively lowers ONNX-MLIR
|
||||
through custom MLIR dialects to simulator artifacts.
|
||||
|
||||
The current target is the PIM simulator stack under `backend-simulators/pim`.
|
||||
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
|
||||
instruction files with `--pim-emit-json`.
|
||||
|
||||
## Overview
|
||||
|
||||
PIM architectures perform most of the computation directly in memory.
|
||||
Raptor's first supported target is `pimsim-nn`, which simulates a chip with:
|
||||
- a shared host memory,
|
||||
- a number of cores that do most of the computation directly in their memory
|
||||
(vector ops, vmm/mvm on ReRAM crossbars),
|
||||
- no branching instructions (branchless architecture) and no hardware loop
|
||||
support — any repeated work (e.g. convolutions) must be unrolled into
|
||||
explicit per-iteration instructions.
|
||||
PIM architectures perform most computation directly in memory. The supported
|
||||
target models a chip with:
|
||||
- shared host memory,
|
||||
- multiple PIM cores,
|
||||
- ReRAM crossbars for vector-matrix / matrix-vector work,
|
||||
- explicit communication between cores,
|
||||
- no hardware branch or loop support in emitted simulator code.
|
||||
|
||||
Because of this, the amount of emitted instructions explodes quickly and the
|
||||
compiler must optimize aggressively at every stage to keep compilation
|
||||
tractable.
|
||||
|
||||
A second target, `PulPim`, is planned for an accelerator with RISC-V cores
|
||||
each carrying its own in-memory computing unit and crossbars. It will live in
|
||||
a dedicated dialect (future work).
|
||||
Because repeated work such as convolutions is eventually made explicit, emitted
|
||||
instruction counts can grow quickly. Most compiler work therefore focuses on
|
||||
lowering, scheduling, memory layout, and code-generation optimizations.
|
||||
|
||||
### Targets and simulators
|
||||
|
||||
`pimsim-nn` (under `backend-simulators/pim/pimsim-nn`) is used for
|
||||
**performance** estimates (latency, energy), but does not functionally execute
|
||||
the JSON code it consumes. To validate the numerical correctness of the JSON
|
||||
code produced by Raptor (or, for comparison, by the `pimcomp` compiler), we use
|
||||
a Rust simulator we maintain in-tree at
|
||||
`backend-simulators/pim/pim-simulator`.
|
||||
- `backend-simulators/pim/pim-simulator` is the in-tree Rust functional
|
||||
simulator used by validation. It reads Raptor's `pim/` artifact directory and
|
||||
compares simulator output against native ONNX-MLIR execution.
|
||||
- `backend-simulators/pim/pimsim-nn` is the non-functional simulator submodule
|
||||
used internally by validation for latency, power, and energy.
|
||||
The helper scripts in `pimcomp_utils/` are for comparison with PIMCOMP-NN and
|
||||
contain local paths; treat them as local utilities, not portable workflows.
|
||||
|
||||
## Compilation pipeline
|
||||
|
||||
The PIM-related sources live under `src/PIM` and the tests under `test/PIM`.
|
||||
When working on this codebase, most changes should stay confined to those
|
||||
trees (you only need to look outside, e.g. at `onnx-mlir` or `llvm`, for
|
||||
framework-level details).
|
||||
The PIM sources live under `src/PIM` and tests under `test/PIM`. CMake exposes
|
||||
them to ONNX-MLIR through generated shim directories under
|
||||
`onnx-mlir/src/Accelerators/PIM` and `onnx-mlir/test/accelerators/PIM`.
|
||||
|
||||
High-level lowering flow:
|
||||
|
||||
```
|
||||
ONNX-MLIR ──► Spatial ──► Pim (tensor) ──► Pim (bufferized) ──► PIM JSON
|
||||
ONNX-MLIR -> Spatial -> Pim (tensor) -> Pim (bufferized) -> PIM artifacts
|
||||
```
|
||||
|
||||
1. **ONNX → Spatial** (`src/PIM/Conversion/ONNXToSpatial`).
|
||||
Lowers ONNX ops into the `spat` dialect (`src/PIM/Dialect/Spatial`).
|
||||
Spatial models a high-level spatial in-memory accelerator: vmm/mvm
|
||||
operations are accelerated by storing a constant RHS matrix into a
|
||||
crossbar. Crossbars cannot be re-programmed during execution, have a
|
||||
limited fixed size, and there is a limited number of them per core.
|
||||
Conversion patterns are split by op family under
|
||||
`Conversion/ONNXToSpatial/Patterns/{Math,NN,Tensor}` (Conv, Gemm, MatMul,
|
||||
Elementwise, ReduceMean, Pool, Relu, Sigmoid, Softmax, Concat, Gather,
|
||||
Reshape, Resize, Split).
|
||||
1. **ONNX -> Spatial** (`src/PIM/Conversion/ONNXToSpatial`).
|
||||
Lowers supported ONNX ops into the `spat` dialect
|
||||
(`src/PIM/Dialect/Spatial`). Conversion patterns are split by op family under
|
||||
`Patterns/{Math,NN,Tensor}` and currently cover Conv, Gemm, MatMul,
|
||||
elementwise Add/Mul/Div, ReduceMean, pooling, Relu, Sigmoid, Softmax,
|
||||
Concat, Gather, Reshape, Resize, and Split.
|
||||
The compiler-layer target adapter supplies the target-neutral
|
||||
`SpatialTargetResources`. Layout-aware plan ops advertise typed alternatives
|
||||
through the Spatial layout interface; the layout planner records the
|
||||
selected layout and explicit materialization edges. `LowerSpatialPlans`
|
||||
then pattern-lowers those selected plans. Contraction and Conv lowering
|
||||
keep semantic problems, target-dependent plans, and IR materializers in
|
||||
separate layers. Passes and their invariant/layout analyses live under
|
||||
`Passes/Transforms` and `Passes/Analyses`.
|
||||
|
||||
2. **Spatial → Pim** (`src/PIM/Conversion/SpatialToPim`).
|
||||
Lowers Spatial to the `pim` dialect (`src/PIM/Dialect/Pim`), which
|
||||
materializes PIM cores (`pim.core`), inter-core communication
|
||||
(`pim.send` / `pim.receive`), halts, and crossbar-level operations.
|
||||
2. **Merge, schedule, and realize Spatial communication**
|
||||
(`src/PIM/Dialect/Spatial/Passes/Transforms/MergeComputeNodes`).
|
||||
`TrivialGraphComputeMerge` performs local graph merging. One
|
||||
`ScheduleAndRealizeSpatial` pass then owns scheduling, intermediate
|
||||
verification, communication realization, and final verification. Supporting
|
||||
scheduling code lives under `MergeComputeNodes/Scheduling`.
|
||||
|
||||
3. **Merge compute nodes** (`src/PIM/Dialect/Spatial/Transforms/MergeComputeNodes`).
|
||||
A DCP-inspired heuristic (Dynamic Critical Path — see the original
|
||||
scheduling paper by Kwok & Ahmad,
|
||||
[DCP-eScience2007](https://clouds.cis.unimelb.edu.au/papers/DCP-eScience2007.pdf))
|
||||
that coarsens the virtual node graph and decides how to group compute
|
||||
nodes onto cores. Our implementation is only DCP-*inspired*: it is a
|
||||
heuristic with different assumptions from the paper (different cost
|
||||
model, constraints from crossbar capacity / core resources, and a
|
||||
windowed coarsening loop instead of full-graph reprioritization). The
|
||||
`dcp-critical-window-size` option controls how many lowest-slack virtual
|
||||
nodes each coarsening iteration considers (0 = legacy full-graph
|
||||
analysis). Related sources: `DCPGraph/DCPAnalysis.cpp`, `Graph.cpp/.hpp`,
|
||||
`MergeComputeNodesPass.cpp`.
|
||||
3. **Spatial -> Pim** (`src/PIM/Conversion/SpatialToPim`).
|
||||
Lowers Spatial operations to the `pim` dialect (`src/PIM/Dialect/Pim`),
|
||||
including `pim.core`, `pim.core_batch`, communication, tensor packing, global
|
||||
tensor materialization, and return-path normalization.
|
||||
|
||||
4. **Bufferization** (`src/PIM/Dialect/Pim/Transforms/Bufferization`).
|
||||
Converts tensor-semantics PIM IR into memref-semantics PIM IR using the
|
||||
standard MLIR `BufferizableOpInterface` machinery
|
||||
(`OpBufferizationInterfaces.*`, `PimBufferization.td`).
|
||||
4. **Bufferization** (`src/PIM/Dialect/Pim/Passes/Transforms/Bufferization`).
|
||||
`PimBufferizationPreparation` establishes writable destinations without
|
||||
duplicating the one-shot copy analysis, `PimOneShotBufferization` runs
|
||||
MLIR's one-shot analysis,
|
||||
`PimMemoryNormalization` forwards/removes redundant copies and normalizes
|
||||
addressable accesses, and `PimBufferizationVerification` checks tensor
|
||||
absence, contiguity, and copy address spaces.
|
||||
|
||||
5. **PIM code generation** (`src/PIM/Pass/PimCodegen`):
|
||||
- `HostConstantFolding` — folds host-side constants.
|
||||
- `MaterializeHostConstantsPass` — materializes the remaining host
|
||||
constants for emission.
|
||||
- `VerificationPass` — checks invariants before emission.
|
||||
- `EmitPimJsonPass` — emits the final PIM JSON consumed by `pimsim-nn`
|
||||
and `pim-simulator`.
|
||||
5. **PIM local-memory planning**
|
||||
(`src/PIM/Dialect/Pim/Passes/Transforms/LocalMemoryPlanning`).
|
||||
Computes whole-core lifetimes, reuses addresses for non-overlapping
|
||||
allocations, and records the explicit plan in PIM IR. Reusable lifetime
|
||||
analysis lives under `src/PIM/Dialect/Pim/Passes/Analyses`.
|
||||
6. **PIM verification and code generation** (`src/PIM/Passes/PimCodegen` and
|
||||
`src/PIM/Compiler`).
|
||||
Verifies the memory plan and other PIM invariants, then emits `.pim` core
|
||||
files, weights, and `memory.bin` / `config.json` without rerunning liveness.
|
||||
|
||||
Supporting pieces:
|
||||
- `src/PIM/Compiler` — PIM-specific compiler options (crossbar size/count,
|
||||
core count, DCP window, experimental conv impl, concat error handling, …)
|
||||
and `PimCodeGen` entry points.
|
||||
- `src/PIM/Common` — shared utilities (`PimCommon`, `LabeledList`).
|
||||
- `src/PIM/Pass` — auxiliary passes (`MessagePass`, `CountInstructionPass`)
|
||||
and the `PIMPasses.h` registry used by `PimAccelerator`.
|
||||
- `src/PIM/PimAccelerator.{cpp,hpp}` — accelerator entry point: registers
|
||||
dialects, passes, and plugs Raptor into the ONNX-MLIR driver.
|
||||
- `src/PIM/Common` - shared IR, filesystem, diagnostics, reports, and utility
|
||||
helpers.
|
||||
- `src/PIM/Compiler` - PIM compiler options, planned-address materialization, binary
|
||||
instruction format, artifact writing, weight emission, and codegen entry
|
||||
points.
|
||||
- `src/PIM/Conversion/SpatialToGraphviz` - optional Spatial graphviz conversion
|
||||
pass.
|
||||
- `src/PIM/Passes` - pass registration and auxiliary passes.
|
||||
- `src/PIM/PimAccelerator.{cpp,hpp}` - ONNX-MLIR accelerator entry point.
|
||||
|
||||
## Key compiler options
|
||||
## PIM compiler options
|
||||
|
||||
Pass these on the `onnx-mlir` command line when compiling for PIM:
|
||||
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.
|
||||
|
||||
- `--maccel=PIM` — select the PIM accelerator.
|
||||
- `--EmitSpatial` / `--EmitPim` / `--EmitPimBufferized` / `--EmitPimCodegen`
|
||||
— stop the pipeline at the requested stage (default: `EmitPimCodegen`).
|
||||
- `--pim-only-codegen` — assume the input is already bufferized PIM IR and
|
||||
run only the codegen tail.
|
||||
- `--crossbar-size=<N>` / `--crossbar-count=<N>` — crossbar dimensions and
|
||||
per-core count.
|
||||
- `--core-count=<N>` — number of cores (`-1` picks the minimum).
|
||||
- `--dcp-critical-window-size=<N>` — DCP coarsening window (0 = legacy).
|
||||
- `--use-experimental-conv-impl` — alternative convolution lowering.
|
||||
- `--ignore-concat-error` — soft-fail corner case in `ConcatOp`.
|
||||
- `--maccel=PIM` - select the PIM accelerator.
|
||||
- `--EmitSpatial`, `--EmitPim`, `--EmitPimBufferized`,
|
||||
`--EmitPimCodegen` - stop the PIM pipeline at the requested stage. The PIM
|
||||
default is `--EmitPimCodegen`.
|
||||
- `--core-count=<N>` - required positive core count for PIM compilation.
|
||||
- `--crossbar-size=<N>` - crossbar width/height. Default in code is `128`.
|
||||
- `--crossbar-count=<N>` - crossbars per core. Default in code is `64`.
|
||||
- `--pim-target-config=<PATH>` - optional PIM target configuration used by the
|
||||
target adapter to construct the target-neutral Spatial scheduling cost and
|
||||
topology model. Resource values must match the explicit core/crossbar flags.
|
||||
- `--pim-memory-report=<summary|none>` - emit the concise combined memory report
|
||||
under `reports/memory_report.txt`, or disable it. Default is `summary`.
|
||||
- `--pim-only-codegen` - assume input is already bufferized PIM IR and only run
|
||||
the codegen tail.
|
||||
- `--pim-emit-json` - also emit `core_*.json` instruction files alongside
|
||||
`core_*.pim`.
|
||||
- `--pim-export-spatial-dataflow=<none|spatial1|spatial2|spatial3|spatial4|all>` -
|
||||
control Spatial dataflow CSV reports for the graph, trivially merged graph,
|
||||
scheduled, and realized snapshots under `reports/`. Default is `none`.
|
||||
- `--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`.
|
||||
- `--pim-conv-im2col-max-elements=<N>` - maximum globally materialized im2col
|
||||
elements per convolution before streaming. Default is `1048576`.
|
||||
- `--pim-conv-stream-chunk-positions=<N>` - maximum output positions per
|
||||
streamed convolution chunk. Default is `1024`.
|
||||
- `--pim-detect-communication-deadlock` - statically simulate expanded
|
||||
send/receive ordering and reject blocking deadlocks. Default is off.
|
||||
|
||||
## Standard PIM hardware profile
|
||||
|
||||
Raptor's standard development and YOLO validation profile is:
|
||||
|
||||
| Parameter | Value |
|
||||
| --- | ---: |
|
||||
| Cores | 144 |
|
||||
| Crossbars per core | 64 |
|
||||
| Crossbar size | 128 × 128 |
|
||||
|
||||
Canonical compiler flags:
|
||||
|
||||
`--crossbar-count=64 --crossbar-size=128 --core-count=144`
|
||||
|
||||
`--core-count` remains mandatory and must be passed explicitly to the compiler.
|
||||
|
||||
Example:
|
||||
|
||||
```bash
|
||||
./build_release/Release/bin/onnx-mlir model.onnx -o /tmp/raptor/model \
|
||||
--maccel=PIM --EmitPimCodegen \
|
||||
--crossbar-count=64 --crossbar-size=128 --core-count=144
|
||||
```
|
||||
|
||||
This writes PIM artifacts under `/tmp/raptor/pim/`.
|
||||
|
||||
## Validation
|
||||
|
||||
Functional validation lives in `validation/` and drives the Rust
|
||||
`pim-simulator` to compare Raptor's output against a reference.
|
||||
|
||||
Per-operation validation (from `validation/`):
|
||||
|
||||
```
|
||||
validate.py \
|
||||
--raptor-path ../cmake-build-release/Release/bin/onnx-mlir \
|
||||
--onnx-include-dir ../onnx-mlir/include
|
||||
```
|
||||
|
||||
End-to-end network validation (example: first 4 layers of YOLOv11n):
|
||||
|
||||
```
|
||||
validate.py \
|
||||
--raptor-path ../cmake-build-release/Release/bin/onnx-mlir \
|
||||
--onnx-include-dir ../onnx-mlir/include \
|
||||
--operations-dir ./networks/yolo11n/depth_04 \
|
||||
--crossbar-size 2048 --crossbar-count 256 --core-count 1000
|
||||
```
|
||||
|
||||
Available networks under `validation/networks/`: `vgg16`, `yolo11n`.
|
||||
Available operations under `validation/operations/`: `add`, `conv`, `div`,
|
||||
`gather`, `gemm`, `gemv`, `mul`, `pool`, `reduce_mean`, `relu`, `resize`,
|
||||
`sigmoid`, `softmax`, `split`.
|
||||
|
||||
## Rebuilding
|
||||
|
||||
Release build (fast):
|
||||
|
||||
```
|
||||
cmake --build /home/nico/raptor/raptor/cmake-build-release --target onnx-mlir -j 30
|
||||
```
|
||||
|
||||
A slower debug build is also available — configure it the same way but with
|
||||
`-DCMAKE_BUILD_TYPE=Debug` (see installation instructions below).
|
||||
Functional validation compiles ONNX models, compares native ONNX-MLIR and PIM
|
||||
simulator outputs, and optionally reports latency, power, and energy. See
|
||||
[`validation/README.md`](validation/README.md) for prerequisites, usage,
|
||||
options, artifacts, and results.
|
||||
|
||||
## Build
|
||||
|
||||
Initialize submodules first:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
```
|
||||
|
||||
The project follows ONNX-MLIR's build requirements. The CI workflow documents
|
||||
the currently used versions and setup:
|
||||
- CMake 4.3.0 in CI,
|
||||
- LLVM/MLIR checked out under `onnx-mlir/llvm-project`,
|
||||
- Protobuf `v34.0`,
|
||||
- Rust stable for `pim-simulator`,
|
||||
- Python packages `numpy`, `onnx`, `colorama` for validation.
|
||||
|
||||
### Protobuf
|
||||
|
||||
Use the following commands to install protobuf:
|
||||
```
|
||||
Install Protobuf if your system does not already provide a compatible version:
|
||||
|
||||
```bash
|
||||
git clone --depth 1 --branch v34.0 https://github.com/protocolbuffers/protobuf
|
||||
cd protobuf
|
||||
mkdir build
|
||||
cd build
|
||||
cmake .. -G Ninja -DCMAKE_BUILD_TYPE=Release
|
||||
ninja
|
||||
sudo ninja install
|
||||
cmake -S protobuf -B protobuf/build -G Ninja \
|
||||
-DCMAKE_BUILD_TYPE=Release \
|
||||
-Dprotobuf_BUILD_TESTS=OFF
|
||||
cmake --build protobuf/build
|
||||
sudo cmake --install protobuf/build
|
||||
```
|
||||
|
||||
You can now remove the protobuf repo directory with:
|
||||
```
|
||||
cd ../..
|
||||
You can then remove the temporary checkout:
|
||||
|
||||
```bash
|
||||
rm -rf protobuf
|
||||
```
|
||||
|
||||
### Mlir
|
||||
### MLIR
|
||||
|
||||
Follow the first part of instructions [here](onnx-mlir/docs/BuildOnLinuxOSX.md) to build mlir.
|
||||
Follow the ONNX-MLIR instructions in
|
||||
`onnx-mlir/docs/BuildOnLinuxOSX.md` to build LLVM/MLIR. The local Raptor build
|
||||
expects `MLIR_DIR` to point at the MLIR CMake package, for example:
|
||||
|
||||
Remember to set ```-DCMAKE_BUILD_TYPE=Debug``` for developing on Raptor
|
||||
|
||||
Moreover, if compiling with build type debug, it is also suggested to use
|
||||
mold as linker (you will need to install it if you don't have it already)
|
||||
to reduce memory usage during linking. You can use it by setting the options:
|
||||
```
|
||||
-DLLVM_USE_LINKER=mold
|
||||
```bash
|
||||
MLIR_DIR=$(pwd)/onnx-mlir/llvm-project/build_release/lib/cmake/mlir
|
||||
```
|
||||
|
||||
If your LLVM build directory is named `build` instead of `build_release`, adjust
|
||||
the path accordingly.
|
||||
|
||||
### Raptor
|
||||
|
||||
Use the following commands to build Raptor.
|
||||
Configure a release build:
|
||||
|
||||
Remember to set ```-DCMAKE_BUILD_TYPE=Debug``` for developing on Raptor.
|
||||
|
||||
Also in this case, it is suggested to use mold as linker to reduce link time and memory usage,
|
||||
setting the options:
|
||||
```
|
||||
-DCMAKE_EXE_LINKER_FLAGS="-fuse-ld=mold" \
|
||||
-DCMAKE_SHARED_LINKER_FLAGS="-fuse-ld=mold" \
|
||||
-DCMAKE_MODULE_LINKER_FLAGS="-fuse-ld=mold"
|
||||
```
|
||||
|
||||
```
|
||||
git submodule update --init --recursive
|
||||
|
||||
MLIR_DIR=$(pwd)/onnx-mlir/llvm-project/build/lib/cmake/mlir
|
||||
mkdir build && cd build
|
||||
cmake .. -G Ninja \
|
||||
```bash
|
||||
MLIR_DIR=$(pwd)/onnx-mlir/llvm-project/build_release/lib/cmake/mlir
|
||||
cmake -S . -B build_release -G Ninja \
|
||||
-DCMAKE_BUILD_TYPE=Release \
|
||||
-DONNX_MLIR_ACCELERATORS=PIM \
|
||||
-DLLVM_ENABLE_ASSERTIONS=ON \
|
||||
-DMLIR_DIR=${MLIR_DIR}
|
||||
cmake --build .
|
||||
```
|
||||
|
||||
If the build fails because of protobuf missing uint definitions,
|
||||
just patch the problematic files by adding ```#include <cstdint>``` to their includes.
|
||||
Configure a debug build similarly:
|
||||
|
||||
```bash
|
||||
MLIR_DIR=$(pwd)/onnx-mlir/llvm-project/build_debug/lib/cmake/mlir
|
||||
cmake -S . -B build_debug -G Ninja \
|
||||
-DCMAKE_BUILD_TYPE=Debug \
|
||||
-DONNX_MLIR_ACCELERATORS=PIM \
|
||||
-DLLVM_ENABLE_ASSERTIONS=ON \
|
||||
-DMLIR_DIR=${MLIR_DIR}
|
||||
```
|
||||
|
||||
For debug development, using `mold` can reduce link time and memory use:
|
||||
|
||||
```bash
|
||||
cmake -S . -B build_debug -G Ninja \
|
||||
-DCMAKE_BUILD_TYPE=Debug \
|
||||
-DONNX_MLIR_ACCELERATORS=PIM \
|
||||
-DLLVM_ENABLE_ASSERTIONS=ON \
|
||||
-DMLIR_DIR=${MLIR_DIR} \
|
||||
-DCMAKE_EXE_LINKER_FLAGS="-fuse-ld=mold" \
|
||||
-DCMAKE_SHARED_LINKER_FLAGS="-fuse-ld=mold" \
|
||||
-DCMAKE_MODULE_LINKER_FLAGS="-fuse-ld=mold"
|
||||
```
|
||||
|
||||
Build the compiler with CMake:
|
||||
|
||||
```bash
|
||||
cmake --build ./build_release
|
||||
cmake --build ./build_debug
|
||||
```
|
||||
|
||||
Do not invoke `ninja` directly for this project; use `cmake --build` so CMake's
|
||||
configuration and generated shims stay consistent.
|
||||
|
||||
If a build fails because Protobuf headers are missing fixed-width integer
|
||||
definitions, patch the affected Protobuf-generated files by adding
|
||||
`#include <cstdint>`.
|
||||
|
||||
## Tests
|
||||
|
||||
The Rust simulator has its own tests:
|
||||
|
||||
```bash
|
||||
cd backend-simulators/pim/pim-simulator
|
||||
cargo test
|
||||
```
|
||||
|
||||
## Repository Layout
|
||||
|
||||
- `src/PIM/` - PIM accelerator implementation.
|
||||
- `test/PIM/` - PIM C++ unit tests.
|
||||
- `validation/` - functional validation scripts, ONNX operation tests, network
|
||||
slices, and pimsim config generation.
|
||||
- `backend-simulators/pim/pim-simulator/` - in-tree Rust functional simulator.
|
||||
- `backend-simulators/pim/pimsim-nn/` - non-functional simulator submodule.
|
||||
- `pimcomp_utils/` - local comparison helpers for PIMCOMP-NN.
|
||||
- `.github/actions/` and `.github/workflows/validate_operations.yml` - CI setup
|
||||
for MLIR/Protobuf caching, building Raptor, and validation.
|
||||
|
||||
+2121
-8
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,3 @@
|
||||
|
||||
[package]
|
||||
name = "pim-simulator"
|
||||
version = "0.1.0"
|
||||
@@ -13,8 +12,9 @@ name = "pimcore"
|
||||
path = "src/lib/pimcore.rs"
|
||||
|
||||
[features]
|
||||
default = ["tracing"]
|
||||
default = []
|
||||
tracing = []
|
||||
profile_time = ["dep:plotly", "dep:comfy-table", "dep:statrs"]
|
||||
|
||||
|
||||
|
||||
@@ -27,3 +27,10 @@ hex = "0"
|
||||
paste = "1"
|
||||
serde = { version = "1", features = ["derive"] }
|
||||
serde_json = "1"
|
||||
statrs = {version="0.16", optional=true}
|
||||
comfy-table = {version="7.1", optional=true}
|
||||
plotly = {version="0.8", optional=true}
|
||||
rayon = "1.12.0"
|
||||
faer = "0.24.0"
|
||||
faer-traits = "0.24.0"
|
||||
mimalloc = "0.1.50"
|
||||
|
||||
@@ -1,14 +1,20 @@
|
||||
use mimalloc::MiMalloc;
|
||||
|
||||
#[global_allocator]
|
||||
static GLOBAL: MiMalloc = MiMalloc;
|
||||
|
||||
use anyhow::{Context, Result, bail};
|
||||
use clap::Parser;
|
||||
use glob::glob;
|
||||
use pimcore::binary_to_instruction::binary_to_executor;
|
||||
use pimcore::cpu::crossbar::Crossbar;
|
||||
use pimcore::json_to_instruction::json_to_executor;
|
||||
use pimcore::memory_manager::CoreMemory;
|
||||
use pimcore::tracing::TRACER;
|
||||
use serde_json::Value;
|
||||
use std::collections::HashMap;
|
||||
use std::fs::{self, read_link};
|
||||
use std::io::Write;
|
||||
use std::fs::{self, File};
|
||||
use std::io::{BufReader, Write};
|
||||
use std::path::PathBuf;
|
||||
|
||||
/// Program to simulate core execution configuration
|
||||
@@ -37,25 +43,31 @@ struct Args {
|
||||
|
||||
/// Comma separated list of (address,size) for memory output dump
|
||||
#[arg(short, long, value_delimiter = ',', num_args = 1.., value_name = "ADDR,SIZE")]
|
||||
dump: Vec<i32>,
|
||||
dump: Vec<usize>,
|
||||
}
|
||||
|
||||
fn main() -> Result<()> {
|
||||
let args = Args::parse();
|
||||
|
||||
let config_json = retrive_config(&args)?;
|
||||
let core_jsons = retrive_cores(&args)?;
|
||||
let mut core_inputs = retrive_cores(&args)?;
|
||||
let memory = retrive_memory(&args)?;
|
||||
let global_crossbars = get_crossbars(&config_json, &args).unwrap();
|
||||
let crossbars = map_crossbars_to_cores(&config_json, &args, &global_crossbars);
|
||||
let mut executor =
|
||||
json_to_executor::json_to_executor(config_json, core_jsons.iter(), crossbars);
|
||||
let mut executor = match &mut core_inputs {
|
||||
CoreInputs::Json(core_jsons) => {
|
||||
json_to_executor::json_to_executor(config_json, core_jsons, crossbars)
|
||||
}
|
||||
CoreInputs::Binary(core_bins) => {
|
||||
binary_to_executor(config_json, core_bins.iter(), crossbars)?
|
||||
}
|
||||
};
|
||||
set_memory(&mut executor, memory);
|
||||
TRACER
|
||||
.lock()
|
||||
.unwrap()
|
||||
.init(executor.cpu().num_core(), args.output.clone());
|
||||
executor.execute();
|
||||
executor.execute()?;
|
||||
dump_memory(executor, &args)?;
|
||||
Ok(())
|
||||
}
|
||||
@@ -65,7 +77,7 @@ fn map_crossbars_to_cores<'c>(
|
||||
args: &Args,
|
||||
global_crossbars: &'c HashMap<String, Crossbar>,
|
||||
) -> Vec<Vec<&'c Crossbar>> {
|
||||
let mut res = Vec::new();
|
||||
let mut res = vec![Vec::new()];
|
||||
let num_cores = config.get("core_cnt").unwrap().as_i64().unwrap() as i32;
|
||||
|
||||
if let Some(folder) = args.folder.as_ref() {
|
||||
@@ -98,7 +110,7 @@ fn map_crossbars_to_cores<'c>(
|
||||
sym_link_files.sort_by_key(|&(num, _)| num);
|
||||
|
||||
for (_, symlink) in sym_link_files {
|
||||
let real_path = read_link(symlink).unwrap();
|
||||
let real_path = symlink.canonicalize().unwrap();
|
||||
let path_as_str = real_path.to_str().unwrap();
|
||||
assert!(
|
||||
global_crossbars.contains_key(path_as_str),
|
||||
@@ -140,12 +152,18 @@ 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 mut crossbar =
|
||||
Crossbar::new(column_corssbar * 4, rows_crossbar, CoreMemory::new());
|
||||
let stored_row_bytes = bytes.len() / rows_crossbar;
|
||||
let mut crossbar = Crossbar::new(
|
||||
std::cmp::max(column_corssbar * 4, stored_row_bytes),
|
||||
rows_crossbar,
|
||||
CoreMemory::new(),
|
||||
);
|
||||
crossbar.execute_store(&bytes).unwrap();
|
||||
res.insert(
|
||||
weight_file
|
||||
.path()
|
||||
.canonicalize()
|
||||
.context("Failed to resolve crossbar path")?
|
||||
.to_str()
|
||||
.context("file name not utf-8")?
|
||||
.to_string(),
|
||||
@@ -157,7 +175,7 @@ fn get_crossbars(config: &Value, args: &Args) -> anyhow::Result<HashMap<String,
|
||||
}
|
||||
|
||||
fn dump_memory(mut executor: pimcore::Executable, args: &Args) -> Result<()> {
|
||||
let dumps: Vec<(i32, i32)> = args
|
||||
let dumps: Vec<(usize, usize)> = args
|
||||
.dump
|
||||
.chunks_exact(2)
|
||||
.map(|chunk| (chunk[0], chunk[1]))
|
||||
@@ -214,45 +232,82 @@ fn retrive_memory(args: &Args) -> Result<Vec<u8>> {
|
||||
Ok(memory_vector)
|
||||
}
|
||||
|
||||
fn retrive_cores(args: &Args) -> Result<Vec<Value>, anyhow::Error> {
|
||||
let mut core_jsons: Vec<Value> = Vec::new();
|
||||
if let Some(cores_override) = &args.cores {
|
||||
for core in cores_override {
|
||||
let content = fs::read_to_string(core)
|
||||
.with_context(|| format!("Failed to read core file: {:?}", cores_override))?;
|
||||
let json: Value =
|
||||
serde_json::from_str(&content).context("Failed to parse core json override")?;
|
||||
core_jsons.push(json);
|
||||
}
|
||||
} else if let Some(folder) = args.folder.as_ref() {
|
||||
let pattern = folder.join("core*.json");
|
||||
let pattern_str = pattern.to_str().context("Invalid path encoding")?;
|
||||
let mut paths: Vec<_> = glob(pattern_str)?.map(|x| x.unwrap()).collect();
|
||||
paths.sort_by_cached_key(|x| {
|
||||
let mut x = x
|
||||
.file_stem()
|
||||
.expect("Extracting the stem")
|
||||
.to_str()
|
||||
.expect("File not utf-8");
|
||||
x = &x[5..];
|
||||
x.parse::<i32>().unwrap()
|
||||
});
|
||||
enum CoreInputs {
|
||||
Json(Vec<BufReader<File>>),
|
||||
Binary(Vec<Vec<u8>>),
|
||||
}
|
||||
|
||||
if paths.is_empty() {
|
||||
bail!("No core*.json files found in {:?}", folder);
|
||||
fn retrive_cores(args: &Args) -> Result<CoreInputs, anyhow::Error> {
|
||||
if let Some(cores_override) = &args.cores {
|
||||
let first_extension = cores_override
|
||||
.first()
|
||||
.and_then(|path| path.extension())
|
||||
.and_then(|ext| ext.to_str())
|
||||
.unwrap_or_default();
|
||||
if first_extension == "pim" {
|
||||
let mut core_bins = Vec::with_capacity(cores_override.len());
|
||||
for core in cores_override {
|
||||
core_bins.push(
|
||||
fs::read(core)
|
||||
.with_context(|| format!("Failed to read binary core file: {:?}", core))?,
|
||||
);
|
||||
}
|
||||
return Ok(CoreInputs::Binary(core_bins));
|
||||
}
|
||||
for entry in paths {
|
||||
let path = entry;
|
||||
let content = fs::read_to_string(&path)
|
||||
.with_context(|| format!("Failed to read core file: {:?}", path))?;
|
||||
let json: Value = serde_json::from_str(&content)
|
||||
.with_context(|| format!("Failed to parse JSON in {:?}", path))?;
|
||||
core_jsons.push(json);
|
||||
let mut core_jsons_reader: Vec<BufReader<File>> = Vec::with_capacity(cores_override.len());
|
||||
for core in cores_override {
|
||||
let file = File::open(core)?;
|
||||
let reader = BufReader::new(file);
|
||||
core_jsons_reader.push(reader);
|
||||
}
|
||||
} else {
|
||||
bail!("Either --core or --folder must be provided to find core definitions.");
|
||||
return Ok(CoreInputs::Json(core_jsons_reader));
|
||||
}
|
||||
Ok(core_jsons)
|
||||
|
||||
if let Some(folder) = args.folder.as_ref() {
|
||||
let binary_pattern = folder.join("core*.pim");
|
||||
let binary_pattern_str = binary_pattern.to_str().context("Invalid path encoding")?;
|
||||
let mut binary_paths: Vec<_> = glob(binary_pattern_str)?.map(|x| x.unwrap()).collect();
|
||||
binary_paths.sort_by_cached_key(core_sort_key);
|
||||
if !binary_paths.is_empty() {
|
||||
let mut core_bins = Vec::with_capacity(binary_paths.len());
|
||||
for path in binary_paths {
|
||||
core_bins.push(
|
||||
fs::read(&path)
|
||||
.with_context(|| format!("Failed to read core file: {:?}", path))?,
|
||||
);
|
||||
}
|
||||
return Ok(CoreInputs::Binary(core_bins));
|
||||
}
|
||||
|
||||
let json_pattern = folder.join("core*.json");
|
||||
let json_pattern_str = json_pattern.to_str().context("Invalid path encoding")?;
|
||||
let mut json_paths: Vec<_> = glob(json_pattern_str)?.map(|x| x.unwrap()).collect();
|
||||
json_paths.sort_by_cached_key(core_sort_key);
|
||||
|
||||
if json_paths.is_empty() {
|
||||
bail!("No core*.pim or core*.json files found in {:?}", folder);
|
||||
}
|
||||
|
||||
let mut core_json_reader: Vec<BufReader<File>> = Vec::with_capacity(json_paths.len());
|
||||
for path in json_paths {
|
||||
let file = File::open(path)?;
|
||||
let reader = BufReader::new(file);
|
||||
core_json_reader.push(reader);
|
||||
}
|
||||
return Ok(CoreInputs::Json(core_json_reader));
|
||||
}
|
||||
|
||||
bail!("Either --core or --folder must be provided to find core definitions.");
|
||||
}
|
||||
|
||||
fn core_sort_key(path: &PathBuf) -> i32 {
|
||||
let mut stem = path
|
||||
.file_stem()
|
||||
.expect("Extracting the stem")
|
||||
.to_str()
|
||||
.expect("File not utf-8");
|
||||
stem = &stem[5..];
|
||||
stem.parse::<i32>().unwrap()
|
||||
}
|
||||
|
||||
fn retrive_config(args: &Args) -> Result<Value, anyhow::Error> {
|
||||
|
||||
@@ -0,0 +1,501 @@
|
||||
use crate::{
|
||||
CoreInstructionsBuilder, Executable,
|
||||
cpu::{CPU, crossbar::Crossbar},
|
||||
instruction_set::{InstructionsBuilder, instruction_data::InstructionDataBuilder, isa::*},
|
||||
};
|
||||
use anyhow::{Context, Result, bail, ensure};
|
||||
use serde_json::Value;
|
||||
use std::collections::HashMap;
|
||||
use std::convert::TryFrom;
|
||||
use std::sync::LazyLock;
|
||||
|
||||
const MAGIC: &[u8; 4] = b"PIMB";
|
||||
const VERSION: u32 = 1;
|
||||
const HEADER_SIZE: usize = 12;
|
||||
const RECORD_SIZE: usize = 20;
|
||||
|
||||
macro_rules! add_name {
|
||||
($storage:ident, $opcode:literal, $name:literal) => {
|
||||
$storage.insert($opcode, $name);
|
||||
};
|
||||
}
|
||||
|
||||
static INSTRUCTIONS: LazyLock<HashMap<usize, &'static str>> = LazyLock::new(|| {
|
||||
let mut hash = HashMap::new();
|
||||
add_name!(hash, 0, "nop");
|
||||
add_name!(hash, 1, "sldi");
|
||||
add_name!(hash, 2, "sld");
|
||||
add_name!(hash, 3, "sadd");
|
||||
add_name!(hash, 4, "ssub");
|
||||
add_name!(hash, 5, "smul");
|
||||
add_name!(hash, 6, "saddi");
|
||||
add_name!(hash, 7, "smuli");
|
||||
add_name!(hash, 8, "setbw");
|
||||
add_name!(hash, 9, "mvmul");
|
||||
add_name!(hash, 10, "vvadd");
|
||||
add_name!(hash, 11, "vvsub");
|
||||
add_name!(hash, 12, "vvmul");
|
||||
add_name!(hash, 13, "vvdmul");
|
||||
add_name!(hash, 14, "vvmax");
|
||||
add_name!(hash, 15, "vvsll");
|
||||
add_name!(hash, 16, "vvsra");
|
||||
add_name!(hash, 17, "vavg");
|
||||
add_name!(hash, 18, "vrelu");
|
||||
add_name!(hash, 19, "vtanh");
|
||||
add_name!(hash, 20, "vsigm");
|
||||
add_name!(hash, 21, "vsoftmax");
|
||||
add_name!(hash, 22, "vmv");
|
||||
add_name!(hash, 23, "vrsu");
|
||||
add_name!(hash, 24, "vrsl");
|
||||
add_name!(hash, 25, "ld");
|
||||
add_name!(hash, 26, "st");
|
||||
add_name!(hash, 27, "lldi");
|
||||
add_name!(hash, 28, "lmv");
|
||||
add_name!(hash, 29, "send");
|
||||
add_name!(hash, 30, "recv");
|
||||
add_name!(hash, 31, "wait");
|
||||
add_name!(hash, 32, "sync");
|
||||
hash
|
||||
});
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
struct InstructionRecord {
|
||||
opcode: u8,
|
||||
rd: u8,
|
||||
r1: u8,
|
||||
r2_or_imm: i32,
|
||||
generic1: i32,
|
||||
generic2: i32,
|
||||
generic3: i32,
|
||||
flags: u8,
|
||||
}
|
||||
|
||||
fn read_u32_le(bytes: &[u8], offset: usize) -> u32 {
|
||||
u32::from_le_bytes(bytes[offset..offset + 4].try_into().unwrap())
|
||||
}
|
||||
|
||||
fn read_i32_le(bytes: &[u8], offset: usize) -> i32 {
|
||||
i32::from_le_bytes(bytes[offset..offset + 4].try_into().unwrap())
|
||||
}
|
||||
|
||||
fn parse_binary_records(bytes: &[u8]) -> Result<Vec<InstructionRecord>> {
|
||||
ensure!(bytes.len() >= HEADER_SIZE, "binary core file too small");
|
||||
ensure!(&bytes[0..4] == MAGIC, "invalid PIM binary magic");
|
||||
|
||||
let version = read_u32_le(bytes, 4);
|
||||
ensure!(
|
||||
version == VERSION,
|
||||
"unsupported PIM binary version {version}"
|
||||
);
|
||||
|
||||
let instruction_count = read_u32_le(bytes, 8) as usize;
|
||||
let expected_len = HEADER_SIZE + instruction_count * RECORD_SIZE;
|
||||
ensure!(
|
||||
bytes.len() == expected_len,
|
||||
"PIM binary size mismatch: expected {expected_len} bytes, got {}",
|
||||
bytes.len()
|
||||
);
|
||||
|
||||
let mut records = Vec::with_capacity(instruction_count);
|
||||
for index in 0..instruction_count {
|
||||
let base = HEADER_SIZE + index * RECORD_SIZE;
|
||||
records.push(InstructionRecord {
|
||||
opcode: bytes[base],
|
||||
rd: bytes[base + 1],
|
||||
r1: bytes[base + 2],
|
||||
flags: bytes[base + 3],
|
||||
r2_or_imm: read_i32_le(bytes, base + 4),
|
||||
generic1: read_i32_le(bytes, base + 8),
|
||||
generic2: read_i32_le(bytes, base + 12),
|
||||
generic3: read_i32_le(bytes, base + 16),
|
||||
});
|
||||
}
|
||||
|
||||
Ok(records)
|
||||
}
|
||||
|
||||
fn append_record(
|
||||
inst_builder: &mut InstructionsBuilder,
|
||||
inst_data_builder: &mut InstructionDataBuilder,
|
||||
record: InstructionRecord,
|
||||
) -> Result<()> {
|
||||
let InstructionRecord {
|
||||
opcode,
|
||||
rd,
|
||||
r1,
|
||||
r2_or_imm,
|
||||
generic1,
|
||||
generic2,
|
||||
generic3,
|
||||
flags: _,
|
||||
} = record;
|
||||
|
||||
match opcode {
|
||||
0 => {}
|
||||
1 => {
|
||||
inst_data_builder.set_rd_u8(rd).set_imm(r2_or_imm);
|
||||
inst_builder.make_inst(sldi, inst_data_builder.build());
|
||||
}
|
||||
2 => {
|
||||
inst_data_builder
|
||||
.set_rd_u8(rd)
|
||||
.set_r1_u8(r1)
|
||||
.set_offset_select(generic1)
|
||||
.set_offset_value(generic2);
|
||||
inst_builder.make_inst(sld, inst_data_builder.build());
|
||||
}
|
||||
3 => {
|
||||
inst_data_builder.set_rdr1r2_u8(rd, r1, r2_or_imm);
|
||||
inst_builder.make_inst(sadd, inst_data_builder.build());
|
||||
}
|
||||
4 => {
|
||||
inst_data_builder.set_rdr1r2_u8(rd, r1, r2_or_imm);
|
||||
inst_builder.make_inst(ssub, inst_data_builder.build());
|
||||
}
|
||||
5 => {
|
||||
inst_data_builder.set_rdr1r2_u8(rd, r1, r2_or_imm);
|
||||
inst_builder.make_inst(smul, inst_data_builder.build());
|
||||
}
|
||||
6 => {
|
||||
inst_data_builder.set_rdr1imm_u8(rd, r1, r2_or_imm);
|
||||
inst_builder.make_inst(saddi, inst_data_builder.build());
|
||||
}
|
||||
7 => {
|
||||
inst_data_builder.set_rdr1imm_u8(rd, r1, r2_or_imm);
|
||||
inst_builder.make_inst(smuli, inst_data_builder.build());
|
||||
}
|
||||
8 => {
|
||||
inst_data_builder.set_ibiw_obiw(generic1, generic2);
|
||||
inst_builder.make_inst(setbw, inst_data_builder.build());
|
||||
}
|
||||
9 => {
|
||||
inst_data_builder
|
||||
.set_rd_u8(rd)
|
||||
.set_r1_u8(r1)
|
||||
.set_mbiw_immrelu_immgroup(r2_or_imm, generic1, generic2);
|
||||
inst_builder.make_inst(mvmul, inst_data_builder.build());
|
||||
}
|
||||
10 => {
|
||||
inst_data_builder
|
||||
.set_rdr1r2_u8(rd, r1, r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vvadd, inst_data_builder.build());
|
||||
}
|
||||
11 => {
|
||||
inst_data_builder
|
||||
.set_rdr1r2_u8(rd, r1, r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vvsub, inst_data_builder.build());
|
||||
}
|
||||
12 => {
|
||||
inst_data_builder
|
||||
.set_rdr1r2_u8(rd, r1, r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vvmul, inst_data_builder.build());
|
||||
}
|
||||
13 => {
|
||||
inst_data_builder
|
||||
.set_rdr1r2_u8(rd, r1, r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vvdmul, inst_data_builder.build());
|
||||
}
|
||||
14 => {
|
||||
inst_data_builder
|
||||
.set_rdr1r2_u8(rd, r1, r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vvmax, inst_data_builder.build());
|
||||
}
|
||||
15 => {
|
||||
inst_data_builder
|
||||
.set_rdr1r2_u8(rd, r1, r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vvsll, inst_data_builder.build());
|
||||
}
|
||||
16 => {
|
||||
inst_data_builder
|
||||
.set_rdr1r2_u8(rd, r1, r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vvsra, inst_data_builder.build());
|
||||
}
|
||||
17 => {
|
||||
inst_data_builder
|
||||
.set_rdr1r2_u8(rd, r1, r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vavg, inst_data_builder.build());
|
||||
}
|
||||
18 => {
|
||||
inst_data_builder
|
||||
.set_rdr1_u8(rd, r1)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vrelu, inst_data_builder.build());
|
||||
}
|
||||
19 => {
|
||||
inst_data_builder
|
||||
.set_rdr1_u8(rd, r1)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vtanh, inst_data_builder.build());
|
||||
}
|
||||
20 => {
|
||||
inst_data_builder
|
||||
.set_rdr1_u8(rd, r1)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vsigm, inst_data_builder.build());
|
||||
}
|
||||
21 => {
|
||||
inst_data_builder
|
||||
.set_rdr1_u8(rd, r1)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vsoftmax, inst_data_builder.build());
|
||||
}
|
||||
22 => {
|
||||
inst_data_builder
|
||||
.set_rdr1r2_u8(rd, r1, r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vmv, inst_data_builder.build());
|
||||
}
|
||||
23 => {
|
||||
inst_data_builder
|
||||
.set_rdr1r2_u8(rd, r1, r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vrsu, inst_data_builder.build());
|
||||
}
|
||||
24 => {
|
||||
inst_data_builder
|
||||
.set_rdr1r2_u8(rd, r1, r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(vrsl, inst_data_builder.build());
|
||||
}
|
||||
25 => {
|
||||
inst_data_builder
|
||||
.set_rdr1_u8(rd, r1)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(ld, inst_data_builder.build());
|
||||
}
|
||||
26 => {
|
||||
inst_data_builder
|
||||
.set_rdr1_u8(rd, r1)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(st, inst_data_builder.build());
|
||||
}
|
||||
27 => {
|
||||
inst_data_builder
|
||||
.set_rd_u8(rd)
|
||||
.set_imm(r2_or_imm)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(lldi, inst_data_builder.build());
|
||||
}
|
||||
28 => {
|
||||
inst_data_builder
|
||||
.set_rdr1_u8(rd, r1)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(lmv, inst_data_builder.build());
|
||||
}
|
||||
29 => {
|
||||
inst_data_builder
|
||||
.set_rd_u8(rd)
|
||||
.set_imm_core(r2_or_imm + 1)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(send, inst_data_builder.build());
|
||||
}
|
||||
30 => {
|
||||
inst_data_builder
|
||||
.set_rd_u8(rd)
|
||||
.set_imm_core(r2_or_imm + 1)
|
||||
.set_imm_len(generic3)
|
||||
.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(recv, inst_data_builder.build());
|
||||
}
|
||||
31 => {
|
||||
inst_data_builder.set_offset_select_value(generic1, generic2);
|
||||
inst_builder.make_inst(wait, inst_data_builder.build());
|
||||
}
|
||||
32 => {
|
||||
inst_data_builder
|
||||
.set_imm_core(r2_or_imm + 1)
|
||||
.set_offset_select_value(generic1, 0);
|
||||
inst_builder.make_inst(sync, inst_data_builder.build());
|
||||
}
|
||||
_ => bail!("unsupported PIM binary opcode {opcode}"),
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn binary_to_instructions(
|
||||
core_bytes: &[u8],
|
||||
core_index: i32,
|
||||
) -> Result<Vec<crate::instruction_set::Instruction>> {
|
||||
let records = parse_binary_records(core_bytes)?;
|
||||
let mut insts_builder = InstructionsBuilder::new();
|
||||
let mut inst_data_builder = InstructionDataBuilder::new();
|
||||
inst_data_builder
|
||||
.set_core_indx_u16(u16::try_from(core_index).expect("core index does not fit in u16"))
|
||||
.fix_core_indx();
|
||||
|
||||
for record in records {
|
||||
let opcode = record.opcode;
|
||||
let name = INSTRUCTIONS
|
||||
.get(&(opcode as usize))
|
||||
.copied()
|
||||
.unwrap_or("<unknown>");
|
||||
|
||||
append_record(&mut insts_builder, &mut inst_data_builder, record).with_context(|| {
|
||||
format!(
|
||||
"while decoding binary instruction for core {core_index}: opcode {opcode} ({name})"
|
||||
)
|
||||
})?;
|
||||
}
|
||||
|
||||
Ok(insts_builder.build())
|
||||
}
|
||||
|
||||
pub fn binary_to_executor<'a, 'b>(
|
||||
config: Value,
|
||||
cores: impl Iterator<Item = &'b Vec<u8>>,
|
||||
crossbars: Vec<Vec<&'a Crossbar>>,
|
||||
) -> Result<Executable<'a>> {
|
||||
let core_cnt = config
|
||||
.get("core_cnt")
|
||||
.context("missing core_cnt in config")?
|
||||
.as_i64()
|
||||
.context("core_cnt is not an integer")? as i32;
|
||||
|
||||
let cpu = CPU::new(core_cnt, crossbars);
|
||||
let mut core_insts_builder = CoreInstructionsBuilder::new(core_cnt as usize);
|
||||
for (external_core_indx, core_bytes) in cores.enumerate() {
|
||||
let core_indx = external_core_indx as i32 + 1;
|
||||
let instructions = binary_to_instructions(core_bytes, core_indx)?;
|
||||
core_insts_builder.set_core(core_indx, instructions);
|
||||
}
|
||||
|
||||
Ok(Executable::new(cpu, core_insts_builder.build()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
HEADER_SIZE, InstructionRecord, MAGIC, RECORD_SIZE, VERSION, binary_to_instructions,
|
||||
};
|
||||
use crate::{
|
||||
functor_to_name,
|
||||
instruction_set::{InstructionsBuilder, instruction_data::InstructionDataBuilder},
|
||||
json_to_instruction::json_isa::json_to_instruction,
|
||||
};
|
||||
|
||||
fn encode_record(record: InstructionRecord, dst: &mut Vec<u8>) {
|
||||
dst.push(record.opcode);
|
||||
dst.push(record.rd);
|
||||
dst.push(record.r1);
|
||||
dst.push(record.flags);
|
||||
dst.extend_from_slice(&record.r2_or_imm.to_le_bytes());
|
||||
dst.extend_from_slice(&record.generic1.to_le_bytes());
|
||||
dst.extend_from_slice(&record.generic2.to_le_bytes());
|
||||
dst.extend_from_slice(&record.generic3.to_le_bytes());
|
||||
}
|
||||
|
||||
fn binary_blob(records: &[InstructionRecord]) -> Vec<u8> {
|
||||
let mut blob = Vec::with_capacity(HEADER_SIZE + records.len() * RECORD_SIZE);
|
||||
blob.extend_from_slice(MAGIC);
|
||||
blob.extend_from_slice(&VERSION.to_le_bytes());
|
||||
blob.extend_from_slice(&(records.len() as u32).to_le_bytes());
|
||||
for &record in records {
|
||||
encode_record(record, &mut blob);
|
||||
}
|
||||
blob
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn json_and_binary_decoders_match_for_representative_ops() {
|
||||
let json_program = [
|
||||
r#"{"imm":64,"op":"sldi","rd":0}"#,
|
||||
r#"{"imm":128,"op":"sldi","rd":1}"#,
|
||||
r#"{"len":16,"offset":{"offset_select":0,"offset_value":0},"op":"lmv","rd":0,"rs1":1}"#,
|
||||
r#"{"group":3,"mbiw":8,"op":"mvmul","rd":0,"relu":0,"rs1":1}"#,
|
||||
r#"{"len":16,"offset":{"offset_select":0,"offset_value":0},"op":"vvadd","rd":0,"rs1":1,"rs2":2}"#,
|
||||
r#"{"core":2,"offset":{"offset_select":0,"offset_value":0},"op":"send","rd":0,"size":16}"#,
|
||||
];
|
||||
|
||||
let binary_program = binary_blob(&[
|
||||
InstructionRecord {
|
||||
opcode: 1,
|
||||
rd: 0,
|
||||
r2_or_imm: 64,
|
||||
..Default::default()
|
||||
},
|
||||
InstructionRecord {
|
||||
opcode: 1,
|
||||
rd: 1,
|
||||
r2_or_imm: 128,
|
||||
..Default::default()
|
||||
},
|
||||
InstructionRecord {
|
||||
opcode: 28,
|
||||
rd: 0,
|
||||
r1: 1,
|
||||
generic3: 16,
|
||||
..Default::default()
|
||||
},
|
||||
InstructionRecord {
|
||||
opcode: 9,
|
||||
rd: 0,
|
||||
r1: 1,
|
||||
r2_or_imm: 8,
|
||||
generic2: 3,
|
||||
..Default::default()
|
||||
},
|
||||
InstructionRecord {
|
||||
opcode: 10,
|
||||
rd: 0,
|
||||
r1: 1,
|
||||
r2_or_imm: 2,
|
||||
generic3: 16,
|
||||
..Default::default()
|
||||
},
|
||||
InstructionRecord {
|
||||
opcode: 29,
|
||||
rd: 0,
|
||||
r2_or_imm: 2,
|
||||
generic3: 16,
|
||||
..Default::default()
|
||||
},
|
||||
]);
|
||||
|
||||
let mut json_builder = InstructionsBuilder::new();
|
||||
let mut json_data_builder = InstructionDataBuilder::new();
|
||||
json_data_builder.set_core_indx(1).fix_core_indx();
|
||||
for inst in json_program {
|
||||
let value = serde_json::from_str(inst).unwrap();
|
||||
json_to_instruction(&mut json_builder, &mut json_data_builder, &value);
|
||||
}
|
||||
let json_instructions = json_builder.build();
|
||||
let binary_instructions = binary_to_instructions(&binary_program, 1).unwrap();
|
||||
|
||||
assert_eq!(json_instructions.len(), binary_instructions.len());
|
||||
for (json_inst, binary_inst) in json_instructions.iter().zip(binary_instructions.iter()) {
|
||||
assert_eq!(
|
||||
functor_to_name(json_inst.functor as usize),
|
||||
functor_to_name(binary_inst.functor as usize)
|
||||
);
|
||||
assert_eq!(json_inst.data, binary_inst.data);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,7 +1,6 @@
|
||||
use crate::memory_manager::{CoreMemory, MemoryStorable};
|
||||
use anyhow::{Result, bail, ensure};
|
||||
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct Crossbar {
|
||||
max_width: usize,
|
||||
@@ -12,7 +11,12 @@ pub struct Crossbar {
|
||||
|
||||
impl Crossbar {
|
||||
pub fn new(width: usize, height: usize, memory: CoreMemory) -> Self {
|
||||
Self { max_width: width, max_height: height, memory, stored_bytes:0 }
|
||||
Self {
|
||||
max_width: width,
|
||||
max_height: height,
|
||||
memory,
|
||||
stored_bytes: 0,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn width(&self) -> usize {
|
||||
@@ -27,27 +31,36 @@ impl Crossbar {
|
||||
self.stored_bytes
|
||||
}
|
||||
|
||||
pub fn execute_store<T>(&mut self, element: &[T]) -> Result<()> where
|
||||
T: MemoryStorable, {
|
||||
pub fn execute_store<T>(&mut self, element: &[T]) -> Result<()>
|
||||
where
|
||||
T: MemoryStorable,
|
||||
{
|
||||
self.memory.clear();
|
||||
let total_size = self.max_width * self.max_height;
|
||||
self.memory.set_capacity(total_size);
|
||||
let stored_size = std::mem::size_of_val(element);
|
||||
ensure!(stored_size <= total_size, "Storing more than crossbar can handle");
|
||||
self.stored_bytes=stored_size;
|
||||
ensure!(
|
||||
stored_size <= total_size,
|
||||
"Storing more than crossbar can handle"
|
||||
);
|
||||
self.stored_bytes = stored_size;
|
||||
self.memory.execute_store(0, element)
|
||||
}
|
||||
|
||||
pub fn load<T>(&self, size: usize) -> Result<Vec<&[T]>> where
|
||||
T: MemoryStorable, {
|
||||
if self.memory.get_len() < size
|
||||
//|| self.stored_bytes < size
|
||||
pub fn load<T>(&self, size: usize) -> Result<Vec<&[T]>>
|
||||
where
|
||||
T: MemoryStorable,
|
||||
{
|
||||
if self.memory.get_len() < size
|
||||
//|| self.stored_bytes < size
|
||||
{
|
||||
bail!("Loading outside crossbar boundary [{} {}] < {}", self.stored_bytes, self.memory.get_len() , size);
|
||||
bail!(
|
||||
"Loading outside crossbar boundary [{} {}] < {}",
|
||||
self.stored_bytes,
|
||||
self.memory.get_len(),
|
||||
size
|
||||
);
|
||||
}
|
||||
self.memory.load_const(0, size)
|
||||
}
|
||||
|
||||
|
||||
|
||||
}
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
use std::{collections::HashMap, fmt::Debug};
|
||||
use crate::utility::AddressArg;
|
||||
use anyhow::{Context, Result, ensure};
|
||||
use std::{collections::HashMap, fmt::Debug};
|
||||
|
||||
use crate::{
|
||||
cpu::crossbar::Crossbar,
|
||||
@@ -15,7 +16,7 @@ pub struct CPU<'a> {
|
||||
}
|
||||
|
||||
impl<'a> CPU<'a> {
|
||||
pub fn new(num_cores: impl TryToUsize, crossbars: Vec<Vec<&'a Crossbar>> ) -> Self {
|
||||
pub fn new(num_cores: impl TryToUsize, crossbars: Vec<Vec<&'a Crossbar>>) -> Self {
|
||||
let num_cores = num_cores.try_into().expect("num_cores can not be negative");
|
||||
assert!(crossbars.len() == num_cores + 1);
|
||||
let mut cores = Vec::new();
|
||||
@@ -28,25 +29,31 @@ impl<'a> CPU<'a> {
|
||||
}
|
||||
|
||||
pub fn host<'b>(&'b mut self) -> &'b mut Core<'a>
|
||||
where 'a : 'b
|
||||
where
|
||||
'a: 'b,
|
||||
{
|
||||
& mut self.cores[0]
|
||||
&mut self.cores[0]
|
||||
}
|
||||
|
||||
pub fn core<'b >(&'b mut self, index: impl TryToUsize) -> &'b mut Core<'a>
|
||||
where 'a : 'b
|
||||
pub fn core<'b>(&'b mut self, index: impl TryToUsize) -> &'b mut Core<'a>
|
||||
where
|
||||
'a: 'b,
|
||||
{
|
||||
let index = index.try_into().expect("can not be negative");
|
||||
& mut self.cores[index]
|
||||
&mut self.cores[index]
|
||||
}
|
||||
|
||||
pub fn num_core(&self) -> usize {
|
||||
self.cores.len()
|
||||
}
|
||||
|
||||
pub(crate) fn host_and_cores<'b, 'c >(&'b mut self, core: impl TryToUsize) -> (&'c mut Core<'a>, &'c mut Core<'a>)
|
||||
where 'a: 'b,
|
||||
'b: 'c
|
||||
pub(crate) fn host_and_cores<'b, 'c>(
|
||||
&'b mut self,
|
||||
core: impl TryToUsize,
|
||||
) -> (&'c mut Core<'a>, &'c mut Core<'a>)
|
||||
where
|
||||
'a: 'b,
|
||||
'b: 'c,
|
||||
{
|
||||
let core = core.try_into().expect("core can not be negative");
|
||||
assert_ne!(
|
||||
@@ -61,8 +68,12 @@ impl<'a> CPU<'a> {
|
||||
(host, core)
|
||||
}
|
||||
|
||||
pub fn get_multiple_cores<'b, const N: usize>(&'b mut self, indices: [usize; N]) -> [&'b mut Core<'a>; N]
|
||||
where 'a : 'b
|
||||
pub fn get_multiple_cores<'b, const N: usize>(
|
||||
&'b mut self,
|
||||
indices: [usize; N],
|
||||
) -> [&'b mut Core<'a>; N]
|
||||
where
|
||||
'a: 'b,
|
||||
{
|
||||
self.cores.get_disjoint_mut(indices).unwrap()
|
||||
}
|
||||
@@ -76,7 +87,7 @@ pub struct Core<'a> {
|
||||
}
|
||||
|
||||
impl<'a> Core<'a> {
|
||||
fn new(crossbars : Vec<&'a Crossbar>) -> Self {
|
||||
fn new(crossbars: Vec<&'a Crossbar>) -> Self {
|
||||
Self {
|
||||
crossbars,
|
||||
memory: CoreMemory::new(),
|
||||
@@ -91,30 +102,26 @@ impl<'a> Core<'a> {
|
||||
self.memory.execute_load()
|
||||
}
|
||||
|
||||
pub fn execute_store<T>(&mut self, address: impl TryToUsize, element: &[T]) -> Result<()>
|
||||
pub fn execute_store<T>(&mut self, address: impl AddressArg, element: &[T]) -> Result<()>
|
||||
where
|
||||
T: MemoryStorable,
|
||||
{
|
||||
let address = address.try_into().context("address can not be negative")?;
|
||||
let address = address.to_address_usize()?;
|
||||
self.memory.execute_store(address, element)
|
||||
}
|
||||
|
||||
pub fn reserve_load(
|
||||
&mut self,
|
||||
address: impl TryToUsize,
|
||||
address: impl AddressArg,
|
||||
size: impl TryToUsize,
|
||||
) -> Result<&mut CoreMemory> {
|
||||
let address = address.try_into().context("address can not be negative")?;
|
||||
let address = address.to_address_usize()?;
|
||||
let size = size.try_into().context("size can not be negative")?;
|
||||
self.memory.reserve_load(address, size)
|
||||
}
|
||||
|
||||
pub fn set_register(&mut self, index: impl TryToUsize, value: i32) {
|
||||
let index = index.try_into().expect("index can not be negative");
|
||||
assert!(
|
||||
value >= 0,
|
||||
"Register cannot be negative if happens remove this and go check where it's used as usize"
|
||||
);
|
||||
self.registers[index] = value;
|
||||
}
|
||||
|
||||
@@ -123,11 +130,11 @@ impl<'a> Core<'a> {
|
||||
self.registers[index]
|
||||
}
|
||||
|
||||
pub fn load<T>(&mut self, address: impl TryToUsize, size: impl TryToUsize) -> Result<Vec<&[T]>>
|
||||
pub fn load<T>(&mut self, address: impl AddressArg, size: impl TryToUsize) -> Result<Vec<&[T]>>
|
||||
where
|
||||
T: MemoryStorable,
|
||||
{
|
||||
let address = address.try_into().context("address can not be negative")?;
|
||||
let address = address.to_address_usize()?;
|
||||
let size = size.try_into().context("size can not be negative")?;
|
||||
self.memory.load(address, size)
|
||||
}
|
||||
@@ -141,8 +148,13 @@ impl<'a> Core<'a> {
|
||||
(memory, crossbars)
|
||||
}
|
||||
|
||||
pub fn memset(&mut self, address: impl TryToUsize, size: impl TryToUsize, val: u8) -> Result<()> {
|
||||
let address = address.try_into().context("address can not be negative")?;
|
||||
pub fn memset(
|
||||
&mut self,
|
||||
address: impl AddressArg,
|
||||
size: impl TryToUsize,
|
||||
val: u8,
|
||||
) -> Result<()> {
|
||||
let address = address.to_address_usize()?;
|
||||
let size = size.try_into().context("size can not be negative")?;
|
||||
self.memory.memset(address, size, val)
|
||||
}
|
||||
|
||||
@@ -1,10 +1,11 @@
|
||||
use paste::paste;
|
||||
use std::convert::TryFrom;
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default)]
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
|
||||
pub struct InstructionData {
|
||||
core_indx: i32,
|
||||
rd: i32,
|
||||
r1: i32,
|
||||
core_indx: u16,
|
||||
rd: u8,
|
||||
r1: u8,
|
||||
//r2 imm mbiw imm_core
|
||||
r2_or_imm: i32,
|
||||
//offset_select imm_relu ibiw
|
||||
@@ -16,18 +17,30 @@ pub struct InstructionData {
|
||||
}
|
||||
|
||||
impl InstructionData {
|
||||
pub fn core_indx(&self) -> i32 {
|
||||
pub fn core_indx_u16(&self) -> u16 {
|
||||
self.core_indx
|
||||
}
|
||||
|
||||
pub fn rd(&self) -> i32 {
|
||||
pub fn core_indx(&self) -> i32 {
|
||||
i32::from(self.core_indx)
|
||||
}
|
||||
|
||||
pub fn rd_u8(&self) -> u8 {
|
||||
self.rd
|
||||
}
|
||||
|
||||
pub fn r1(&self) -> i32 {
|
||||
pub fn rd(&self) -> i32 {
|
||||
i32::from(self.rd)
|
||||
}
|
||||
|
||||
pub fn r1_u8(&self) -> u8 {
|
||||
self.r1
|
||||
}
|
||||
|
||||
pub fn r1(&self) -> i32 {
|
||||
i32::from(self.r1)
|
||||
}
|
||||
|
||||
pub fn r2(&self) -> i32 {
|
||||
self.r2_or_imm
|
||||
}
|
||||
@@ -49,26 +62,26 @@ impl InstructionData {
|
||||
}
|
||||
|
||||
pub fn get_core_rd_r1(&self) -> (i32, i32, i32) {
|
||||
(self.core_indx, self.rd, self.r1)
|
||||
(self.core_indx(), self.rd(), self.r1())
|
||||
}
|
||||
|
||||
pub fn get_core_rd_r1_r2(&self) -> (i32, i32, i32, i32) {
|
||||
(self.core_indx, self.rd, self.r1, self.r2_or_imm)
|
||||
(self.core_indx(), self.rd(), self.r1(), self.r2_or_imm)
|
||||
}
|
||||
|
||||
pub fn get_core_rd_imm(&self) -> (i32, i32, i32) {
|
||||
(self.core_indx, self.rd, self.r2_or_imm)
|
||||
(self.core_indx(), self.rd(), self.r2_or_imm)
|
||||
}
|
||||
|
||||
pub fn get_core_rd_r1_imm(&self) -> (i32, i32, i32, i32) {
|
||||
(self.core_indx, self.rd, self.r1, self.r2_or_imm)
|
||||
(self.core_indx(), self.rd(), self.r1(), self.r2_or_imm)
|
||||
}
|
||||
|
||||
pub fn get_core_rd_r1_r2_immlen_offset(&self) -> (i32, i32, i32, i32, i32, i32, i32) {
|
||||
(
|
||||
self.core_indx,
|
||||
self.rd,
|
||||
self.r1,
|
||||
self.core_indx(),
|
||||
self.rd(),
|
||||
self.r1(),
|
||||
self.r2_or_imm,
|
||||
self.generic3,
|
||||
self.generic1,
|
||||
@@ -78,9 +91,9 @@ impl InstructionData {
|
||||
|
||||
pub fn get_core_rd_r1_mbiw_immrelu_immgroup(&self) -> (i32, i32, i32, i32, i32, i32) {
|
||||
(
|
||||
self.core_indx,
|
||||
self.rd,
|
||||
self.r1,
|
||||
self.core_indx(),
|
||||
self.rd(),
|
||||
self.r1(),
|
||||
self.r2_or_imm,
|
||||
self.generic1,
|
||||
self.generic2,
|
||||
@@ -100,7 +113,7 @@ impl InstructionData {
|
||||
}
|
||||
|
||||
pub(crate) fn get_core_immcore(&self) -> (i32, i32) {
|
||||
(self.core_indx, self.r2_or_imm)
|
||||
(self.core_indx(), self.r2_or_imm)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -216,6 +229,18 @@ impl InstructionDataBuilder {
|
||||
common_getter_setter![imm_group];
|
||||
common_getter_setter![imm_core];
|
||||
|
||||
pub fn set_core_indx_u16(&mut self, val: u16) -> &mut Self {
|
||||
self.set_core_indx(i32::from(val))
|
||||
}
|
||||
|
||||
pub fn set_rd_u8(&mut self, val: u8) -> &mut Self {
|
||||
self.set_rd(i32::from(val))
|
||||
}
|
||||
|
||||
pub fn set_r1_u8(&mut self, val: u8) -> &mut Self {
|
||||
self.set_r1(i32::from(val))
|
||||
}
|
||||
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
core_indx: Fixer::Edit(0),
|
||||
@@ -253,7 +278,12 @@ impl InstructionDataBuilder {
|
||||
}
|
||||
|
||||
fn check_sanity(&self) {
|
||||
assert!(!(self.get_r2() != 0 && self.get_imm() != 0 && self.get_mbiw() != 0 && self.get_imm_core() != 0));
|
||||
assert!(
|
||||
!(self.get_r2() != 0
|
||||
&& self.get_imm() != 0
|
||||
&& self.get_mbiw() != 0
|
||||
&& self.get_imm_core() != 0)
|
||||
);
|
||||
assert!(
|
||||
!(self.get_ibiw() != 0 && self.get_offset_select() != 0 && self.get_imm_relu() != 0)
|
||||
);
|
||||
@@ -265,9 +295,9 @@ impl InstructionDataBuilder {
|
||||
pub fn build(&mut self) -> InstructionData {
|
||||
self.check_sanity();
|
||||
let inst_data = InstructionData {
|
||||
core_indx: self.get_core_indx(),
|
||||
rd: self.get_rd(),
|
||||
r1: self.get_r1(),
|
||||
core_indx: u16::try_from(self.get_core_indx()).expect("core index does not fit in u16"),
|
||||
rd: u8::try_from(self.get_rd()).expect("rd does not fit in u8"),
|
||||
r1: u8::try_from(self.get_r1()).expect("r1 does not fit in u8"),
|
||||
r2_or_imm: self.get_r2() + self.get_imm() + self.get_mbiw() + self.get_imm_core(),
|
||||
generic1: self.get_offset_select() + self.get_ibiw() + self.get_imm_relu(),
|
||||
generic2: self.get_offset_value() + self.get_obiw() + self.get_imm_group(),
|
||||
@@ -281,6 +311,10 @@ impl InstructionDataBuilder {
|
||||
self.set_rd(rd).set_r1(r1).set_r2(r2)
|
||||
}
|
||||
|
||||
pub fn set_rdr1r2_u8(&mut self, rd: u8, r1: u8, r2: i32) -> &mut Self {
|
||||
self.set_rd_u8(rd).set_r1_u8(r1).set_r2(r2)
|
||||
}
|
||||
|
||||
pub fn set_offset_select_value(&mut self, offset_select: i32, offset_value: i32) -> &mut Self {
|
||||
self.set_offset_select(offset_select)
|
||||
.set_offset_value(offset_value)
|
||||
@@ -290,14 +324,26 @@ impl InstructionDataBuilder {
|
||||
self.set_rd(rd).set_r1(r1).set_imm(imm)
|
||||
}
|
||||
|
||||
pub fn set_rdr1imm_u8(&mut self, rd: u8, r1: u8, imm: i32) -> &mut Self {
|
||||
self.set_rd_u8(rd).set_r1_u8(r1).set_imm(imm)
|
||||
}
|
||||
|
||||
pub fn set_rdr1(&mut self, rd: i32, r1: i32) -> &mut Self {
|
||||
self.set_rd(rd).set_r1(r1)
|
||||
}
|
||||
|
||||
pub fn set_rdr1_u8(&mut self, rd: u8, r1: u8) -> &mut Self {
|
||||
self.set_rd_u8(rd).set_r1_u8(r1)
|
||||
}
|
||||
|
||||
pub fn set_rdimm(&mut self, rd: i32, imm: i32) -> &mut Self {
|
||||
self.set_rd(rd).set_imm(imm)
|
||||
}
|
||||
|
||||
pub fn set_rdimm_u8(&mut self, rd: u8, imm: i32) -> &mut Self {
|
||||
self.set_rd_u8(rd).set_imm(imm)
|
||||
}
|
||||
|
||||
pub fn set_ibiw_obiw(&mut self, ibiw: i32, obiw: i32) -> &mut Self {
|
||||
self.set_ibiw(ibiw).set_obiw(obiw)
|
||||
}
|
||||
|
||||
@@ -1,17 +1,22 @@
|
||||
use crate::{
|
||||
cpu::{CPU, crossbar}, instruction_set::{
|
||||
cpu::{CPU, crossbar},
|
||||
instruction_set::{
|
||||
Instruction, InstructionData, InstructionStatus, InstructionType, VectorBitWith,
|
||||
helper::add_all,
|
||||
}, memory_manager::{
|
||||
},
|
||||
memory_manager::{
|
||||
MemoryStorable,
|
||||
type_traits::{FromFloat, UpcastDestTraits, UpcastSlice},
|
||||
}, tracing::TRACER, utility::{add_offset_r1, add_offset_r2, add_offset_rd}
|
||||
},
|
||||
tracing::TRACER,
|
||||
utility::{AddressArg, add_offset_r1, add_offset_r2, add_offset_rd},
|
||||
};
|
||||
use aligned_vec::{AVec, ConstAlign};
|
||||
use anyhow::{Context, Result, ensure};
|
||||
use rayon::prelude::*;
|
||||
|
||||
use paste::paste;
|
||||
use std::{borrow::Cow, cell::OnceCell, collections::HashMap};
|
||||
use std::{borrow::Cow, cell::OnceCell, collections::HashMap, mem::size_of};
|
||||
use std::{collections::HashSet, sync::LazyLock};
|
||||
|
||||
macro_rules! add_name {
|
||||
@@ -30,7 +35,7 @@ macro_rules! add_name_simd {
|
||||
};
|
||||
}
|
||||
|
||||
static NAMES: LazyLock<HashMap<usize, &'static str>> = LazyLock::new(|| {
|
||||
pub static NAMES: LazyLock<HashMap<usize, &'static str>> = LazyLock::new(|| {
|
||||
let mut hash = HashMap::new();
|
||||
add_name!(hash, sldi);
|
||||
add_name!(hash, sld);
|
||||
@@ -53,7 +58,7 @@ static NAMES: LazyLock<HashMap<usize, &'static str>> = LazyLock::new(|| {
|
||||
add_name_simd!(hash, vtanh);
|
||||
add_name_simd!(hash, vsigm);
|
||||
add_name_simd!(hash, vsoftmax);
|
||||
add_name!(hash, vmv);
|
||||
add_name_simd!(hash, vmv);
|
||||
add_name!(hash, vrsu);
|
||||
add_name!(hash, vrsl);
|
||||
add_name!(hash, ld);
|
||||
@@ -76,8 +81,8 @@ pub fn functor_to_name(functor: usize) -> &'static str {
|
||||
///////////////////////////////////////////////////////////////
|
||||
/////////////////Scalar/register Instructions//////////////////
|
||||
///////////////////////////////////////////////////////////////
|
||||
pub fn sldi(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
{
|
||||
#[inline(never)]
|
||||
pub fn sldi(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_sldi(cores, data);
|
||||
let (core_indx, rd, imm) = data.get_core_rd_imm();
|
||||
let core = cores.core(core_indx);
|
||||
@@ -86,6 +91,7 @@ pub fn sldi(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn sld(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_sld(cores, data);
|
||||
let (core_indx, rd, r1) = data.get_core_rd_r1();
|
||||
@@ -100,6 +106,7 @@ pub fn sld(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn sadd(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_sadd(cores, data);
|
||||
let (core_indx, rd, r1, r2) = data.get_core_rd_r1_r2();
|
||||
@@ -110,6 +117,7 @@ pub fn sadd(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn ssub(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_ssub(cores, data);
|
||||
let (core_indx, rd, r1, r2) = data.get_core_rd_r1_r2();
|
||||
@@ -120,6 +128,7 @@ pub fn ssub(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn smul(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_smul(cores, data);
|
||||
let (core_indx, rd, r1, r2) = data.get_core_rd_r1_r2();
|
||||
@@ -130,6 +139,7 @@ pub fn smul(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn saddi(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_saddi(cores, data);
|
||||
let (core_indx, rd, r1, imm) = data.get_core_rd_r1_imm();
|
||||
@@ -139,6 +149,7 @@ pub fn saddi(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn smuli(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_smuli(cores, data);
|
||||
let (core_indx, rd, r1, imm) = data.get_core_rd_r1_imm();
|
||||
@@ -159,8 +170,6 @@ macro_rules! add_simd_to_map {
|
||||
tmp.insert((32_usize,64_usize), ([<$id _impl>]::<f32, f64> as InstructionType));
|
||||
tmp.insert((64_usize,32_usize), ([<$id _impl>]::<f64, f32> as InstructionType));
|
||||
tmp.insert((64_usize,64_usize), ([<$id _impl>]::<f64, f64> as InstructionType));
|
||||
//TODO WTF WHY
|
||||
tmp.insert((8_usize,8_usize), ([<$id _impl>]::<f32, f32> as InstructionType));
|
||||
$storage.insert($id as *const () as usize, tmp);
|
||||
|
||||
}
|
||||
@@ -180,6 +189,7 @@ static SIMD: LazyLock<HashMap<usize, HashMap<(usize, usize), InstructionType>>>
|
||||
add_simd_to_map!(storage, vtanh);
|
||||
add_simd_to_map!(storage, vsigm);
|
||||
add_simd_to_map!(storage, vsoftmax);
|
||||
add_simd_to_map!(storage, vmv);
|
||||
add_simd_to_map!(storage, mvmul);
|
||||
storage
|
||||
});
|
||||
@@ -213,14 +223,25 @@ pub fn is_setbw(functor: InstructionType) -> bool {
|
||||
functor as usize == setbw as *const () as usize
|
||||
}
|
||||
|
||||
fn vector_lengths<F>(imm_len: i32) -> Result<(usize, usize)> {
|
||||
let element_count: usize = imm_len.try_into().context("imm_len can not be negative")?;
|
||||
let byte_len = element_count
|
||||
.checked_mul(size_of::<F>())
|
||||
.context("vector byte length overflow")?;
|
||||
Ok((element_count, byte_len))
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn setbw(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!("You are calling a placeholder, this instruction is resolved in the construction phase");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn mvmul(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn mvm_impl_internal<F, M, T>(
|
||||
cores: &mut CPU,
|
||||
data: InstructionData,
|
||||
@@ -229,47 +250,64 @@ where
|
||||
[F]: UpcastSlice<T> + UpcastSlice<M>,
|
||||
[M]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
M: UpcastDestTraits<M> + MemoryStorable + FromFloat,
|
||||
// Add faer::ComplexField HERE, directly bounding M for this function only
|
||||
M: UpcastDestTraits<M> + MemoryStorable + FromFloat + faer_traits::ComplexField,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
TRACER.lock().unwrap().pre_mvm::<F,M,T>(cores, data);
|
||||
TRACER.lock().unwrap().pre_mvm::<F, M, T>(cores, data);
|
||||
|
||||
let (core_indx, rd, r1, mbiw, relu, group) = data.get_core_rd_r1_mbiw_immrelu_immgroup();
|
||||
let group: usize = group.try_into().context("group can not be negative")?;
|
||||
|
||||
let core = cores.core(core_indx);
|
||||
let r1_val = core.register(r1);
|
||||
let rd_val = core.register(rd);
|
||||
|
||||
let (memory, crossbars) = core.get_memory_crossbar();
|
||||
let crossbar = crossbars.get_mut(group).unwrap();
|
||||
let crossbar_stored_bytes = crossbar.stored_bytes();
|
||||
let crossbar_byte_width = crossbar.width();
|
||||
//Fix this
|
||||
let crossbar_elem_width = crossbar_byte_width / size_of::<M>();
|
||||
ensure!(
|
||||
crossbar_byte_width & size_of::<M>() == 0,
|
||||
"M not divisor of the crosbbar size"
|
||||
);
|
||||
let crossbar_height = crossbar.height();
|
||||
let crossbar_byte_size = crossbar_byte_width * crossbar_height;
|
||||
let crossbar_stored_bytes = crossbar.stored_bytes();
|
||||
let bytes_per_column = crossbar_height * size_of::<M>();
|
||||
ensure!(bytes_per_column != 0, "crossbar height can not be zero");
|
||||
ensure!(
|
||||
crossbar_stored_bytes % bytes_per_column == 0,
|
||||
"Stored crossbar bytes do not describe an integral number of columns"
|
||||
);
|
||||
let crossbar_elem_width = crossbar_stored_bytes / bytes_per_column;
|
||||
ensure!(
|
||||
crossbar_elem_width != 0,
|
||||
"Crossbar contains no stored columns"
|
||||
);
|
||||
|
||||
let loads = memory
|
||||
.reserve_load(r1_val, crossbar_height * size_of::<F>())?
|
||||
.execute_load::<F>()?;
|
||||
let load = loads[0];
|
||||
let vec: Cow<[M]> = load.up();
|
||||
let matrix = crossbar.load::<M>(crossbar_byte_size)?[0];
|
||||
let mut res = Vec::with_capacity(crossbar_elem_width);
|
||||
let mut partial :AVec<M, _> = AVec::<M, ConstAlign<64>>::with_capacity(64, vec.len());
|
||||
partial.resize(vec.len(), M::from_f32(0.0));
|
||||
let matrix = crossbar.load::<M>(crossbar_stored_bytes)?[0];
|
||||
|
||||
for x in 0..crossbar_elem_width {
|
||||
partial[0] = vec[0] * matrix[x];
|
||||
for y in 1..crossbar_height {
|
||||
partial[y] = vec[y] * matrix[y * crossbar_elem_width + x];
|
||||
}
|
||||
// --- FAER IMPLEMENTATION ---
|
||||
|
||||
// 1. Explicitly create a Matrix Reference (MatRef)
|
||||
let matrix_view = faer::mat::MatRef::from_row_major_slice(
|
||||
matrix.as_ref(),
|
||||
crossbar_height,
|
||||
crossbar_elem_width,
|
||||
);
|
||||
|
||||
// 2. Explicitly create a Column Vector Reference (ColRef)
|
||||
// Using `ColRef` here guarantees we don't accidentally get a RowRef (Fixes E0277)
|
||||
let vec_view = faer::col::ColRef::from_slice(vec.as_ref());
|
||||
|
||||
let res_col: faer::col::Col<M> = matrix_view.transpose() * vec_view;
|
||||
|
||||
// 4. Convert back to standard Rust Vec
|
||||
// try_as_slice() returns an Option<&[M]>.
|
||||
// We can safely unwrap() because a freshly allocated, owned Col is ALWAYS contiguous!
|
||||
let mut res: Vec<M> = (0..crossbar_elem_width).map(|i| res_col[i]).collect();
|
||||
|
||||
// --- END FAER ---
|
||||
|
||||
let mut acc = add_all(partial.as_slice());
|
||||
res.push(acc);
|
||||
}
|
||||
if relu != 0 {
|
||||
res.iter_mut().for_each(|x| {
|
||||
if *x < M::from_f32(0.0) {
|
||||
@@ -277,16 +315,20 @@ where
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
ensure!(
|
||||
res.len() == crossbar_elem_width,
|
||||
"mvm generate a vector bigger thant it's requested elements"
|
||||
"mvm generate a vector bigger thant it's requested elements"
|
||||
);
|
||||
|
||||
let res_up: Cow<[T]> = res.as_slice().up();
|
||||
core.execute_store(rd_val, res_up.as_ref());
|
||||
TRACER.lock().unwrap().post_mvm::<F,M,T>(cores, data);
|
||||
|
||||
TRACER.lock().unwrap().post_mvm::<F, M, T>(cores, data);
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn mvmul_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
where
|
||||
[F]: UpcastSlice<T> + UpcastSlice<f32> + UpcastSlice<f64>,
|
||||
@@ -307,17 +349,19 @@ where
|
||||
}
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vvadd(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn vvadd_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
TRACER.lock().unwrap().pre_vvadd::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().pre_vvadd::<F, T>(cores, data);
|
||||
let (core_indx, rd, r1, r2, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let core = cores.core(core_indx);
|
||||
@@ -327,11 +371,11 @@ where
|
||||
let r1_val = add_offset_r1(r1_val, offset_select, offset_value);
|
||||
let r2_val = add_offset_r2(r2_val, offset_select, offset_value);
|
||||
let rd_val = add_offset_rd(rd_val, offset_select, offset_value);
|
||||
let imm_len: usize = imm_len.try_into().context("imm_len can not be negative")?;
|
||||
let (element_count, byte_len) = vector_lengths::<F>(imm_len)?;
|
||||
|
||||
let loads = core
|
||||
.reserve_load(r1_val, imm_len)?
|
||||
.reserve_load(r2_val, imm_len)?
|
||||
.reserve_load(r1_val, byte_len)?
|
||||
.reserve_load(r2_val, byte_len)?
|
||||
.execute_load::<F>()?;
|
||||
let (load1, load2) = (loads[0], loads[1]);
|
||||
let res: Vec<F> = load1
|
||||
@@ -340,26 +384,28 @@ where
|
||||
.map(|(&a, &b)| a + b)
|
||||
.collect();
|
||||
ensure!(
|
||||
imm_len / size_of::<F>() == res.len(),
|
||||
element_count == res.len(),
|
||||
"vvadd generate a vector bigger thant it's requested elements"
|
||||
);
|
||||
let res_up: Cow<[T]> = res.as_slice().up();
|
||||
core.execute_store(rd_val, res_up.as_ref());
|
||||
TRACER.lock().unwrap().post_vvadd::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().post_vvadd::<F, T>(cores, data);
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vvsub(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn vvsub_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
TRACER.lock().unwrap().pre_vvsub::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().pre_vvsub::<F, T>(cores, data);
|
||||
let (core_indx, rd, r1, r2, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let core = cores.core(core_indx);
|
||||
@@ -369,11 +415,11 @@ where
|
||||
let r1_val = add_offset_r1(r1_val, offset_select, offset_value);
|
||||
let r2_val = add_offset_r2(r2_val, offset_select, offset_value);
|
||||
let rd_val = add_offset_rd(rd_val, offset_select, offset_value);
|
||||
let imm_len: usize = imm_len.try_into().context("imm_len can not be negative")?;
|
||||
let (element_count, byte_len) = vector_lengths::<F>(imm_len)?;
|
||||
|
||||
let loads = core
|
||||
.reserve_load(r1_val, imm_len)?
|
||||
.reserve_load(r2_val, imm_len)?
|
||||
.reserve_load(r1_val, byte_len)?
|
||||
.reserve_load(r2_val, byte_len)?
|
||||
.execute_load::<F>()?;
|
||||
let (load1, load2) = (loads[0], loads[1]);
|
||||
let res: Vec<F> = load1
|
||||
@@ -382,7 +428,7 @@ where
|
||||
.map(|(&a, &b)| a - b)
|
||||
.collect();
|
||||
ensure!(
|
||||
imm_len / size_of::<F>() == res.len(),
|
||||
element_count == res.len(),
|
||||
"vvadd generate a vector bigger thant it's requested elements"
|
||||
);
|
||||
let res_up: Cow<[T]> = res.as_slice().up();
|
||||
@@ -394,13 +440,14 @@ pub fn vvmul(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn vvmul_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
TRACER.lock().unwrap().pre_vvmul::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().pre_vvmul::<F, T>(cores, data);
|
||||
let (core_indx, rd, r1, r2, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let core = cores.core(core_indx);
|
||||
@@ -410,10 +457,10 @@ where
|
||||
let r1_val = add_offset_r1(r1_val, offset_select, offset_value);
|
||||
let r2_val = add_offset_r2(r2_val, offset_select, offset_value);
|
||||
let rd_val = add_offset_rd(rd_val, offset_select, offset_value);
|
||||
let imm_len: usize = imm_len.try_into().context("imm_len can not be negative")?;
|
||||
let (element_count, byte_len) = vector_lengths::<F>(imm_len)?;
|
||||
let loads = core
|
||||
.reserve_load(r1_val, imm_len)?
|
||||
.reserve_load(r2_val, imm_len)?
|
||||
.reserve_load(r1_val, byte_len)?
|
||||
.reserve_load(r2_val, byte_len)?
|
||||
.execute_load::<F>()?;
|
||||
let (load1, load2) = (loads[0], loads[1]);
|
||||
let res: Vec<F> = load1
|
||||
@@ -422,7 +469,7 @@ where
|
||||
.map(|(&a, &b)| a * b)
|
||||
.collect();
|
||||
ensure!(
|
||||
imm_len / size_of::<F>() == res.len(),
|
||||
element_count == res.len(),
|
||||
"vvadd generate a vector bigger thant it's requested elements"
|
||||
);
|
||||
let res_up: Cow<[T]> = res.as_slice().up();
|
||||
@@ -430,17 +477,19 @@ where
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vvdmul(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn vvdmul_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
TRACER.lock().unwrap().pre_vvdmul::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().pre_vvdmul::<F, T>(cores, data);
|
||||
let (core_indx, rd, r1, r2, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let core = cores.core(core_indx);
|
||||
@@ -449,9 +498,10 @@ where
|
||||
let rd_val = core.register(rd);
|
||||
let r1_val = add_offset_r1(r1_val, offset_select, offset_value);
|
||||
let r2_val = add_offset_r2(r2_val, offset_select, offset_value);
|
||||
let (_, byte_len) = vector_lengths::<F>(imm_len)?;
|
||||
let loads = core
|
||||
.reserve_load(r1_val, imm_len)?
|
||||
.reserve_load(r2_val, imm_len)?
|
||||
.reserve_load(r1_val, byte_len)?
|
||||
.reserve_load(r2_val, byte_len)?
|
||||
.execute_load::<F>()?;
|
||||
let (load1, load2) = (loads[0], loads[1]);
|
||||
let res: [F; 1] = [load1
|
||||
@@ -466,17 +516,19 @@ where
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vvmax(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn vvmax_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
TRACER.lock().unwrap().pre_vvmax::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().pre_vvmax::<F, T>(cores, data);
|
||||
let (core_indx, rd, r1, r2, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let core = cores.core(core_indx);
|
||||
@@ -486,10 +538,11 @@ where
|
||||
let r1_val = add_offset_r1(r1_val, offset_select, offset_value);
|
||||
let r2_val = add_offset_r2(r2_val, offset_select, offset_value);
|
||||
let rd_val = add_offset_rd(rd_val, offset_select, offset_value);
|
||||
let (_, byte_len) = vector_lengths::<F>(imm_len)?;
|
||||
|
||||
let loads = core
|
||||
.reserve_load(r1_val, imm_len)?
|
||||
.reserve_load(r2_val, imm_len)?
|
||||
.reserve_load(r1_val, byte_len)?
|
||||
.reserve_load(r2_val, byte_len)?
|
||||
.execute_load::<F>()?;
|
||||
let (load1, load2) = (loads[0], loads[1]);
|
||||
let res: Vec<F> = load1
|
||||
@@ -503,29 +556,33 @@ where
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vvsll(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!(
|
||||
"Shift left on floating point what does it means? who has generated this instruction???"
|
||||
);
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vvsra(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!(
|
||||
"Shift right on floating point what does it means? who has generated this instruction???"
|
||||
);
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vavg(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn vavg_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
TRACER.lock().unwrap().pre_vavg::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().pre_vavg::<F, T>(cores, data);
|
||||
let (core_indx, rd, r1, r2, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let core = cores.core(core_indx);
|
||||
@@ -533,9 +590,13 @@ where
|
||||
let r2_val = r2;
|
||||
ensure!(r2_val == 1, "Stride different than 1 not supported");
|
||||
let rd_val = core.register(rd);
|
||||
ensure!(offset_select == 1, "Offset select cannot be different from 1");
|
||||
ensure!(
|
||||
offset_select == 1,
|
||||
"Offset select cannot be different from 1"
|
||||
);
|
||||
let r1_val = add_offset_r1(r1_val, offset_select, offset_value);
|
||||
let loads = core.reserve_load(r1_val, imm_len)?.execute_load::<F>()?;
|
||||
let (_, byte_len) = vector_lengths::<F>(imm_len)?;
|
||||
let loads = core.reserve_load(r1_val, byte_len)?.execute_load::<F>()?;
|
||||
let load1 = loads[0];
|
||||
let len = load1.len();
|
||||
let res: [F; _] =
|
||||
@@ -545,17 +606,19 @@ where
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vrelu(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn vrelu_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
TRACER.lock().unwrap().pre_vrelu::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().pre_vrelu::<F, T>(cores, data);
|
||||
let (core_indx, rd, r1, r2, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let core = cores.core(core_indx);
|
||||
@@ -563,7 +626,8 @@ where
|
||||
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 loads = core.reserve_load(r1_val, imm_len)?.execute_load::<F>()?;
|
||||
let (_, byte_len) = vector_lengths::<F>(imm_len)?;
|
||||
let loads = core.reserve_load(r1_val, byte_len)?.execute_load::<F>()?;
|
||||
let load1 = loads[0];
|
||||
let res: Vec<F> = load1
|
||||
.iter()
|
||||
@@ -575,17 +639,19 @@ where
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vtanh(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn vtanh_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
TRACER.lock().unwrap().pre_vtanh::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().pre_vtanh::<F, T>(cores, data);
|
||||
let (core_indx, rd, r1, r2, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let core = cores.core(core_indx);
|
||||
@@ -594,7 +660,8 @@ where
|
||||
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 loads = core.reserve_load(r1_val, imm_len)?.execute_load::<F>()?;
|
||||
let (_, byte_len) = vector_lengths::<F>(imm_len)?;
|
||||
let loads = core.reserve_load(r1_val, byte_len)?.execute_load::<F>()?;
|
||||
let load1 = loads[0];
|
||||
let res: Vec<F> = load1.iter().map(|&a| a.tanh()).collect();
|
||||
|
||||
@@ -603,17 +670,19 @@ where
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vsigm(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn vsigm_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
TRACER.lock().unwrap().pre_vsigm::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().pre_vsigm::<F, T>(cores, data);
|
||||
let (core_indx, rd, r1, r2, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let core = cores.core(core_indx);
|
||||
@@ -621,7 +690,8 @@ where
|
||||
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 loads = core.reserve_load(r1_val, imm_len)?.execute_load::<F>()?;
|
||||
let (_, byte_len) = vector_lengths::<F>(imm_len)?;
|
||||
let loads = core.reserve_load(r1_val, byte_len)?.execute_load::<F>()?;
|
||||
let load1 = loads[0];
|
||||
let res: Vec<F> = load1.iter().map(|&a| a.sigm()).collect();
|
||||
let res_up: Cow<[T]> = res.as_slice().up();
|
||||
@@ -629,17 +699,22 @@ where
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vsoftmax(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
pub(super) fn vsoftmax_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
#[inline(never)]
|
||||
pub(super) fn vsoftmax_impl<F, T>(
|
||||
cores: &mut CPU,
|
||||
data: InstructionData,
|
||||
) -> Result<InstructionStatus>
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
TRACER.lock().unwrap().pre_vsoftmax::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().pre_vsoftmax::<F, T>(cores, data);
|
||||
let (core_indx, rd, r1, r2, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let core = cores.core(core_indx);
|
||||
@@ -647,7 +722,8 @@ where
|
||||
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 loads = core.reserve_load(r1_val, imm_len)?.execute_load::<F>()?;
|
||||
let (_, byte_len) = vector_lengths::<F>(imm_len)?;
|
||||
let loads = core.reserve_load(r1_val, byte_len)?.execute_load::<F>()?;
|
||||
let load1 = loads[0];
|
||||
ensure!(!load1.is_empty(), "vsoftmax does not support empty vectors");
|
||||
let max_val = load1
|
||||
@@ -656,27 +732,63 @@ where
|
||||
.reduce(|a, b| if a > b { a } else { b })
|
||||
.unwrap();
|
||||
let exp_values: Vec<F> = load1.iter().map(|&a| (a - max_val).exp()).collect();
|
||||
let sum = exp_values
|
||||
.iter()
|
||||
.copied()
|
||||
.reduce(|a, b| a + b)
|
||||
.unwrap();
|
||||
ensure!(sum > 0.0.into(), "vsoftmax normalization sum must be positive");
|
||||
let sum = exp_values.iter().copied().reduce(|a, b| a + b).unwrap();
|
||||
ensure!(
|
||||
sum > 0.0.into(),
|
||||
"vsoftmax normalization sum must be positive"
|
||||
);
|
||||
let res: Vec<F> = exp_values.iter().map(|&a| a / sum).collect();
|
||||
let res_up: Cow<[T]> = res.as_slice().up();
|
||||
core.execute_store(rd_val, res_up.as_ref());
|
||||
TRACER.lock().unwrap().post_vsoftmax::<F,T>(cores, data);
|
||||
TRACER.lock().unwrap().post_vsoftmax::<F, T>(cores, data);
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vmv(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
todo!()
|
||||
panic!("You are calling a placeholder, the real call is the generic version");
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub(super) fn vmv_impl<F, T>(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
where
|
||||
F: Copy + MemoryStorable,
|
||||
{
|
||||
let (core_indx, rd, r1, r2, imm_len, _, _) = data.get_core_rd_r1_r2_immlen_offset();
|
||||
let core = cores.core(core_indx);
|
||||
let source = core.register(r1).to_address_usize()?;
|
||||
let destination = core.register(rd);
|
||||
let stride: usize = core
|
||||
.register(r2)
|
||||
.try_into()
|
||||
.context("vmv stride can not be negative")?;
|
||||
let (element_count, _) = vector_lengths::<F>(imm_len)?;
|
||||
let stride_bytes = stride
|
||||
.checked_mul(size_of::<F>())
|
||||
.context("vmv byte stride overflow")?;
|
||||
let mut result = Vec::with_capacity(element_count);
|
||||
for index in 0..element_count {
|
||||
let offset = index
|
||||
.checked_mul(stride_bytes)
|
||||
.context("vmv source offset overflow")?;
|
||||
let address = source
|
||||
.checked_add(offset)
|
||||
.context("vmv source address overflow")?;
|
||||
result.push(
|
||||
core.reserve_load(address, size_of::<F>())?
|
||||
.execute_load::<F>()?[0][0],
|
||||
);
|
||||
}
|
||||
core.execute_store(destination, &result)?;
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vrsu(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
todo!()
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn vrsl(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
todo!()
|
||||
}
|
||||
@@ -684,6 +796,7 @@ pub fn vrsl(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
///////////////////////////////////////////////////////////////
|
||||
///Communication/synchronization Instructions/////////////////
|
||||
///////////////////////////////////////////////////////////////
|
||||
#[inline(never)]
|
||||
pub fn ld(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_ld(cores, data);
|
||||
let (core, rd, r1, _, imm_len, offset_select, offset_value) =
|
||||
@@ -700,6 +813,7 @@ pub fn ld(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn st(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_st(cores, data);
|
||||
let (core, rd, r1, _, imm_len, offset_select, offset_value) =
|
||||
@@ -716,6 +830,7 @@ pub fn st(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn lldi(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_lldi(cores, data);
|
||||
let (core, rd, imm) = data.get_core_rd_imm();
|
||||
@@ -732,6 +847,7 @@ pub fn lldi(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn lmv(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_lmv(cores, data);
|
||||
let (core, rd, r1, _, imm_len, offset_select, offset_value) =
|
||||
@@ -748,20 +864,32 @@ pub fn lmv(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus>
|
||||
Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn isa_send(functor: usize) -> bool {
|
||||
(send as *const () as usize) == functor
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn send(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_send(cores, data);
|
||||
Ok(InstructionStatus::Sending(data))
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn isa_recv(functor: usize) -> bool {
|
||||
(recv as *const () as usize) == functor
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn recv(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
TRACER.lock().unwrap().pre_recv(cores, data);
|
||||
Ok(InstructionStatus::Reciving(data))
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn wait(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
Ok(InstructionStatus::Waiting(data))
|
||||
}
|
||||
|
||||
#[inline(never)]
|
||||
pub fn sync(cores: &mut CPU, data: InstructionData) -> Result<InstructionStatus> {
|
||||
Ok(InstructionStatus::Sync(data))
|
||||
}
|
||||
|
||||
@@ -7,14 +7,14 @@ use crate::{
|
||||
};
|
||||
use anyhow::{Context, Result};
|
||||
use std::mem::swap;
|
||||
pub mod helper;
|
||||
pub mod instruction_data;
|
||||
pub mod isa;
|
||||
pub mod helper;
|
||||
|
||||
#[derive(Clone, Copy, Debug)]
|
||||
pub struct Instruction {
|
||||
pub data: InstructionData,
|
||||
functor: InstructionType,
|
||||
pub functor: InstructionType,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
@@ -41,7 +41,8 @@ impl Instruction {
|
||||
}
|
||||
|
||||
pub fn execute<'a, 'b>(&'b self, cpu: &mut CPU<'a>) -> InstructionStatus
|
||||
where 'a : 'b
|
||||
where
|
||||
'a: 'b,
|
||||
{
|
||||
(self.functor)(cpu, self.data)
|
||||
.with_context(|| format!("Instruction: {}", functor_to_name(self.functor as usize)))
|
||||
|
||||
@@ -567,7 +567,7 @@ fn json_to_send(
|
||||
let (offset_select, offset_value) = json_to_offset(json.get("offset").unwrap());
|
||||
inst_data_builder
|
||||
.set_rd(rd)
|
||||
.set_imm_core(core)
|
||||
.set_imm_core(core + 1)
|
||||
.set_imm_len(size)
|
||||
.set_offset_select(offset_select)
|
||||
.set_offset_value(offset_value);
|
||||
@@ -588,7 +588,7 @@ fn json_to_recv(
|
||||
let (offset_select, offset_value) = json_to_offset(json.get("offset").unwrap());
|
||||
inst_data_builder
|
||||
.set_rd(rd)
|
||||
.set_imm_core(core)
|
||||
.set_imm_core(core + 1)
|
||||
.set_imm_len(size)
|
||||
.set_offset_select(offset_select)
|
||||
.set_offset_value(offset_value);
|
||||
@@ -601,7 +601,11 @@ fn json_to_wait(
|
||||
inst_data_builder: &mut InstructionDataBuilder,
|
||||
json: &Value,
|
||||
) -> Result<()> {
|
||||
todo!("Not present in the compiler");
|
||||
inst_data_builder.set_offset_select_value(
|
||||
json_i64!(json, "event_register") as i32,
|
||||
json_i64!(json, "wait_value") as i32,
|
||||
);
|
||||
inst_builder.make_inst(wait, inst_data_builder.build());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -610,7 +614,10 @@ fn json_to_sync(
|
||||
inst_data_builder: &mut InstructionDataBuilder,
|
||||
json: &Value,
|
||||
) -> Result<()> {
|
||||
todo!("Not present in the compiler");
|
||||
inst_data_builder
|
||||
.set_imm_core(json_i64!(json, "core") as i32 + 1)
|
||||
.set_offset_select_value(json_i64!(json, "event_register") as i32, 0);
|
||||
inst_builder.make_inst(sync, inst_data_builder.build());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
+15
-28
@@ -1,45 +1,32 @@
|
||||
use core::panic;
|
||||
use std::collections::HashMap;
|
||||
|
||||
use serde_json::{Map, Value};
|
||||
use serde_json::Value;
|
||||
use std::{fs::File, io::BufReader};
|
||||
|
||||
use crate::{
|
||||
CoreInstructionsBuilder, Executable,
|
||||
cpu::{CPU, crossbar::{self, Crossbar}},
|
||||
instruction_set::{
|
||||
InstructionsBuilder,
|
||||
instruction_data::{self, InstructionData, InstructionDataBuilder},
|
||||
},
|
||||
json_to_instruction::{self, json_isa},
|
||||
memory_manager::type_traits::TryToUsize,
|
||||
cpu::{CPU, crossbar::Crossbar},
|
||||
instruction_set::{InstructionsBuilder, instruction_data::InstructionDataBuilder},
|
||||
json_to_instruction::json_isa,
|
||||
};
|
||||
|
||||
|
||||
pub fn json_to_executor<'a>(
|
||||
pub fn json_to_executor<'a, 'b>(
|
||||
config: Value,
|
||||
mut cores: impl Iterator<Item = &'a Value>,
|
||||
crossbars : Vec<Vec<&'a Crossbar>>
|
||||
cores: &'b mut Vec<BufReader<File>>,
|
||||
crossbars: Vec<Vec<&'a Crossbar>>,
|
||||
) -> Executable<'a> {
|
||||
let cell_precision = config.get("cell_precision").unwrap().as_i64().unwrap() as i32;
|
||||
let core_cnt = config.get("core_cnt").unwrap().as_i64().unwrap() as i32 - 1;
|
||||
let xbar_count = config.get("xbar_array_count").unwrap().as_i64().unwrap() as i32;
|
||||
let xbar_size = config.get("xbar_size").unwrap().as_array().unwrap();
|
||||
let rows_crossbar = xbar_size[0].as_i64().unwrap() as i32;
|
||||
let column_corssbar = xbar_size[1].as_i64().unwrap() as i32;
|
||||
let core_cnt = config.get("core_cnt").unwrap().as_i64().unwrap() as i32;
|
||||
|
||||
let mut cpu = CPU::new(core_cnt, crossbars);
|
||||
let cpu = CPU::new(core_cnt, crossbars);
|
||||
let mut core_insts_builder = CoreInstructionsBuilder::new(core_cnt as usize);
|
||||
cores.next();
|
||||
for core_indx in 1..=core_cnt {
|
||||
for (external_core_indx, json_core_reader) in cores.iter_mut().enumerate() {
|
||||
let core_indx = external_core_indx as i32 + 1;
|
||||
let mut insts_builder = InstructionsBuilder::new();
|
||||
let mut inst_data_builder = InstructionDataBuilder::new();
|
||||
inst_data_builder.set_core_indx(core_indx).fix_core_indx();
|
||||
let json_core = cores
|
||||
.next()
|
||||
.unwrap_or_else(|| panic!("cores files less than {}", core_indx ));
|
||||
let json_core: Value = serde_json::from_reader(json_core_reader)
|
||||
.unwrap_or_else(|err| panic!("failed to parse core{}: {}", external_core_indx, err));
|
||||
let json_core_insts = json_core
|
||||
.as_array()
|
||||
.unwrap_or_else(|| panic!("core{} has not a list of instruction", core_indx));
|
||||
.unwrap_or_else(|| panic!("core{} has not a list of instruction", external_core_indx));
|
||||
for json_inst in json_core_insts {
|
||||
json_isa::json_to_instruction(&mut insts_builder, &mut inst_data_builder, json_inst);
|
||||
}
|
||||
|
||||
@@ -1,2 +1,2 @@
|
||||
mod json_isa;
|
||||
pub(crate) mod json_isa;
|
||||
pub mod json_to_executor;
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
use std::cmp::min;
|
||||
use std::fmt::Debug;
|
||||
|
||||
use anyhow::{Context, Result, bail, ensure};
|
||||
@@ -77,7 +78,11 @@ impl CoreMemory {
|
||||
}
|
||||
}
|
||||
|
||||
pub fn reserve_load(&mut self, address: impl TryToUsize, size: impl TryToUsize) -> Result<&mut Self>
|
||||
pub fn reserve_load(
|
||||
&mut self,
|
||||
address: impl TryToUsize,
|
||||
size: impl TryToUsize,
|
||||
) -> Result<&mut Self>
|
||||
where {
|
||||
let address = address.try_into().context("address can not be negative")?;
|
||||
let size = size.try_into().context("size can not be negative")?;
|
||||
@@ -86,7 +91,8 @@ where {
|
||||
size,
|
||||
};
|
||||
if self.memory.len() < address + size {
|
||||
self.memory.resize((address + size) * 2, 0);
|
||||
self.memory
|
||||
.resize(min((address + size) * 2, u32::MAX as usize), 0);
|
||||
}
|
||||
self.load_requests.push(load_request);
|
||||
Ok(self)
|
||||
@@ -104,14 +110,24 @@ where {
|
||||
for (load_index, load_request) in load_requests.drain(..).enumerate() {
|
||||
let LoadRequest { index, size } = load_request;
|
||||
let memory_slice = &memory[index..index + size];
|
||||
let memory_slice = unsafe { slice_from_u8(memory_slice) }
|
||||
.with_context(|| format!("Load number: {} Accessing from {} to {}", load_index, index, index + size))?;
|
||||
let memory_slice = unsafe { slice_from_u8(memory_slice) }.with_context(|| {
|
||||
format!(
|
||||
"Load number: {} Accessing from {} to {}",
|
||||
load_index,
|
||||
index,
|
||||
index + size
|
||||
)
|
||||
})?;
|
||||
res.push(memory_slice);
|
||||
}
|
||||
Ok(res)
|
||||
}
|
||||
|
||||
pub fn load_const<T>(&self, address: impl TryToUsize, size: impl TryToUsize) -> Result<Vec<&[T]>>
|
||||
pub fn load_const<T>(
|
||||
&self,
|
||||
address: impl TryToUsize,
|
||||
size: impl TryToUsize,
|
||||
) -> Result<Vec<&[T]>>
|
||||
where
|
||||
T: MemoryStorable,
|
||||
{
|
||||
@@ -124,7 +140,7 @@ where {
|
||||
let mut res = Vec::new();
|
||||
let memory_slice = &memory[address..address + size];
|
||||
let memory_slice = unsafe { slice_from_u8(memory_slice) }
|
||||
.with_context(|| format!("Accessing from {} to {}", address, address + size))?;
|
||||
.with_context(|| format!("Accessing from {} to {}", address, address + size))?;
|
||||
res.push(memory_slice);
|
||||
Ok(res)
|
||||
}
|
||||
@@ -154,7 +170,12 @@ where {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn memset(&mut self, address: impl TryToUsize, size: impl TryToUsize, val: u8) -> Result<()> {
|
||||
pub fn memset(
|
||||
&mut self,
|
||||
address: impl TryToUsize,
|
||||
size: impl TryToUsize,
|
||||
val: u8,
|
||||
) -> Result<()> {
|
||||
let address = address.try_into().expect("address can not be negative");
|
||||
let size = size.try_into().expect("size can not be negative");
|
||||
let Self { memory, .. } = self;
|
||||
@@ -173,11 +194,11 @@ where {
|
||||
}
|
||||
}
|
||||
|
||||
pub fn get_len(&self) ->usize {
|
||||
pub fn get_len(&self) -> usize {
|
||||
self.memory.len()
|
||||
}
|
||||
|
||||
pub(crate) fn clear(&mut self) {
|
||||
pub(crate) fn clear(&mut self) {
|
||||
self.memory.clear();
|
||||
}
|
||||
}
|
||||
@@ -256,7 +277,11 @@ mod test {
|
||||
let mut data = [0_f32; 2];
|
||||
data[0] = loads[0][0] + loads[1][0];
|
||||
core_memory.execute_store(4, &data[0..1]).unwrap();
|
||||
let loads: &[f32] = core_memory.reserve_load(0, 16).unwrap().execute_load().unwrap()[0];
|
||||
let loads: &[f32] = core_memory
|
||||
.reserve_load(0, 16)
|
||||
.unwrap()
|
||||
.execute_load()
|
||||
.unwrap()[0];
|
||||
println!("{:?}", loads);
|
||||
assert!(loads[0] == 5_f32 && loads[1] == 12_f32 && loads[2] == 7_f32 && loads[3] == 15_f32)
|
||||
}
|
||||
|
||||
@@ -1,5 +1,3 @@
|
||||
|
||||
|
||||
use std::{
|
||||
borrow::Cow,
|
||||
fmt::Debug,
|
||||
@@ -9,32 +7,32 @@ use std::{
|
||||
use anyhow::Context;
|
||||
|
||||
pub trait FromFloat {
|
||||
fn from_f32(val :f32) -> Self;
|
||||
fn from_f64(val :f64) -> Self;
|
||||
fn from_f32(val: f32) -> Self;
|
||||
fn from_f64(val: f64) -> Self;
|
||||
}
|
||||
|
||||
impl FromFloat for f32 {
|
||||
fn from_f32(val :f32) -> Self {
|
||||
fn from_f32(val: f32) -> Self {
|
||||
val
|
||||
}
|
||||
|
||||
fn from_f64(val :f64) -> Self {
|
||||
fn from_f64(val: f64) -> Self {
|
||||
val as f32
|
||||
}
|
||||
}
|
||||
|
||||
impl FromFloat for f64 {
|
||||
fn from_f32(val :f32) -> Self {
|
||||
fn from_f32(val: f32) -> Self {
|
||||
val as f64
|
||||
}
|
||||
|
||||
fn from_f64(val :f64) -> Self {
|
||||
val
|
||||
fn from_f64(val: f64) -> Self {
|
||||
val
|
||||
}
|
||||
}
|
||||
|
||||
pub trait HasTanh {
|
||||
fn tanh(self) -> Self ;
|
||||
fn tanh(self) -> Self;
|
||||
}
|
||||
|
||||
impl HasTanh for f32 {
|
||||
@@ -50,7 +48,7 @@ impl HasTanh for f64 {
|
||||
}
|
||||
|
||||
pub trait HasSigm {
|
||||
fn sigm(self) -> Self ;
|
||||
fn sigm(self) -> Self;
|
||||
}
|
||||
|
||||
impl HasSigm for f32 {
|
||||
@@ -91,34 +89,36 @@ impl HasExp for f64 {
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
pub trait TryToUsize: TryInto<usize, Error = Self::TryError>
|
||||
where std::result::Result<usize, Self::TryError> : Context<usize, Self::TryError>
|
||||
pub trait TryToUsize: TryInto<usize, Error = Self::TryError>
|
||||
where
|
||||
std::result::Result<usize, Self::TryError>: Context<usize, Self::TryError>,
|
||||
{
|
||||
type TryError: Debug + Send + Sync + 'static + std::error::Error;
|
||||
}
|
||||
|
||||
impl<T, E> TryToUsize for T
|
||||
where
|
||||
impl<T, E> TryToUsize for T
|
||||
where
|
||||
T: TryInto<usize, Error = E>,
|
||||
E: Debug + Send + Sync + 'static + std::error::Error,
|
||||
std::result::Result<usize, E> : Context<usize, E>
|
||||
std::result::Result<usize, E>: Context<usize, E>,
|
||||
{
|
||||
type TryError = E;
|
||||
}
|
||||
|
||||
|
||||
pub trait FromUsize {
|
||||
fn from_usize(v: usize) -> Self;
|
||||
}
|
||||
|
||||
impl FromUsize for f32 {
|
||||
fn from_usize(v: usize) -> Self { v as f32 }
|
||||
fn from_usize(v: usize) -> Self {
|
||||
v as f32
|
||||
}
|
||||
}
|
||||
|
||||
impl FromUsize for f64 {
|
||||
fn from_usize(v: usize) -> Self { v as f64 }
|
||||
fn from_usize(v: usize) -> Self {
|
||||
v as f64
|
||||
}
|
||||
}
|
||||
|
||||
pub trait UpcastDestTraits<T>:
|
||||
|
||||
@@ -1,14 +1,22 @@
|
||||
#![allow(unused)]
|
||||
|
||||
use std::time::{Duration, SystemTime};
|
||||
use anyhow::{Result, bail};
|
||||
use std::{
|
||||
collections::{HashMap, HashSet},
|
||||
time::{Duration, SystemTime},
|
||||
};
|
||||
|
||||
use crate::{
|
||||
cpu::CPU,
|
||||
instruction_set::{Instruction, InstructionStatus, Instructions, isa::functor_to_name},
|
||||
instruction_set::{
|
||||
Instruction, InstructionStatus, Instructions,
|
||||
isa::{NAMES, functor_to_name, isa_recv, isa_send},
|
||||
},
|
||||
memory_manager::type_traits::TryToUsize,
|
||||
send_recv::{SendRecv, handle_send_recv},
|
||||
tracing::TRACER,
|
||||
};
|
||||
pub mod binary_to_instruction;
|
||||
pub mod cpu;
|
||||
pub mod instruction_set;
|
||||
pub mod json_to_instruction;
|
||||
@@ -80,6 +88,13 @@ pub struct Executable<'a> {
|
||||
send_recv: SendRecv,
|
||||
}
|
||||
|
||||
struct DeadlockInfo {
|
||||
cycle: String,
|
||||
states: String,
|
||||
}
|
||||
|
||||
type SyncEvents = Vec<[i32; 32]>;
|
||||
|
||||
fn print_status(core_instructions: &[CoreInstructions]) {
|
||||
let mut tot_instructions = 0;
|
||||
let mut progress = 0;
|
||||
@@ -111,7 +126,7 @@ impl<'a> Executable<'a> {
|
||||
}
|
||||
}
|
||||
|
||||
pub fn execute<'b>(&'b mut self)
|
||||
pub fn execute<'b>(&'b mut self) -> Result<()>
|
||||
where
|
||||
'a: 'b,
|
||||
{
|
||||
@@ -122,6 +137,7 @@ impl<'a> Executable<'a> {
|
||||
} = self;
|
||||
let mut cpu_progressed = 0;
|
||||
let max_core = cpu.num_core();
|
||||
let mut sync_events: SyncEvents = vec![[0; 32]; max_core];
|
||||
let mut cpu_index = 0;
|
||||
let mut now = SystemTime::now();
|
||||
|
||||
@@ -144,12 +160,21 @@ impl<'a> Executable<'a> {
|
||||
cpu_progressed = 0;
|
||||
*program_counter += 1;
|
||||
}
|
||||
if (now.elapsed().unwrap() > Duration::from_secs(1)) {
|
||||
print_status(&cores_instructions);
|
||||
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();
|
||||
}
|
||||
}
|
||||
handle_wait_sync(cpu, cores_instructions, core_result);
|
||||
if handle_wait_sync(cores_instructions, &mut sync_events, core_result) {
|
||||
cpu_progressed = 0;
|
||||
}
|
||||
match handle_send_recv(cpu, cores_instructions, send_recv, core_result) {
|
||||
(true, other_cpu_index) => {
|
||||
cpu_progressed = 0;
|
||||
@@ -169,6 +194,24 @@ impl<'a> Executable<'a> {
|
||||
}
|
||||
}
|
||||
print_status(cores_instructions);
|
||||
|
||||
if let Some(deadlock) = detect_deadlock(cores_instructions) {
|
||||
bail!(
|
||||
"Deadlock cycle detected: {} [{}]",
|
||||
deadlock.cycle,
|
||||
deadlock.states
|
||||
);
|
||||
}
|
||||
if cores_instructions
|
||||
.iter()
|
||||
.any(|core_inst| core_inst.program_counter < core_inst.instructions.len())
|
||||
{
|
||||
bail!("Execution stalled with unfinished instructions");
|
||||
}
|
||||
|
||||
#[cfg(feature = "profile_time")]
|
||||
TRACER.lock().unwrap().report();
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn cpu(&self) -> &CPU<'a> {
|
||||
@@ -190,12 +233,152 @@ impl<'a> Executable<'a> {
|
||||
}
|
||||
}
|
||||
|
||||
fn handle_wait_sync<'a, 'b, 'c>(
|
||||
cpu: &'b mut CPU<'a>,
|
||||
core_instructions: &'c mut [CoreInstructions],
|
||||
core_result: InstructionStatus,
|
||||
) where
|
||||
'a: 'b,
|
||||
'a: 'c,
|
||||
{
|
||||
fn detect_deadlock(cores_instructions: &[CoreInstructions]) -> Option<DeadlockInfo> {
|
||||
#[derive(Debug, PartialEq, Eq)]
|
||||
enum CoreState {
|
||||
SendingTo(i32, i32),
|
||||
ReceivingFrom(i32, i32),
|
||||
Working,
|
||||
Halted,
|
||||
}
|
||||
|
||||
let mut states = HashMap::new();
|
||||
|
||||
for core_inst in cores_instructions.iter() {
|
||||
if core_inst.program_counter >= core_inst.instructions.len() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let Instruction { data, functor } = core_inst.instructions[core_inst.program_counter];
|
||||
let functor_address = functor as usize;
|
||||
|
||||
let (this_core, target_core) = data.get_core_immcore();
|
||||
|
||||
if isa_recv(functor_address) {
|
||||
states.insert(
|
||||
this_core,
|
||||
CoreState::ReceivingFrom(target_core, data.imm_len()),
|
||||
);
|
||||
} else if isa_send(functor_address) {
|
||||
states.insert(this_core, CoreState::SendingTo(target_core, data.imm_len()));
|
||||
} else {
|
||||
states.insert(this_core, CoreState::Working);
|
||||
}
|
||||
}
|
||||
|
||||
let mut wait_for = HashMap::new();
|
||||
|
||||
for (&core_id, state) in states.iter() {
|
||||
match state {
|
||||
CoreState::SendingTo(target_core, size) => {
|
||||
let target_state = states.get(target_core).unwrap_or(&CoreState::Halted);
|
||||
if target_state != &CoreState::ReceivingFrom(core_id, *size) {
|
||||
wait_for.insert(core_id, *target_core);
|
||||
}
|
||||
}
|
||||
CoreState::ReceivingFrom(target_core, size) => {
|
||||
let target_state = states.get(target_core).unwrap_or(&CoreState::Halted);
|
||||
if target_state != &CoreState::SendingTo(core_id, *size) {
|
||||
wait_for.insert(core_id, *target_core);
|
||||
}
|
||||
}
|
||||
CoreState::Working | CoreState::Halted => {}
|
||||
}
|
||||
}
|
||||
|
||||
let mut visited = HashSet::new();
|
||||
|
||||
for &start_core in wait_for.keys() {
|
||||
if visited.contains(&start_core) {
|
||||
continue;
|
||||
}
|
||||
|
||||
let mut path = Vec::new();
|
||||
let mut current_core = start_core;
|
||||
let mut in_path = HashSet::new();
|
||||
|
||||
while let Some(&waiting_for) = wait_for.get(¤t_core) {
|
||||
path.push(current_core);
|
||||
in_path.insert(current_core);
|
||||
visited.insert(current_core);
|
||||
|
||||
// Found a closed loop!
|
||||
if in_path.contains(&waiting_for) {
|
||||
let cycle_start = path.iter().position(|&c| c == waiting_for).unwrap();
|
||||
let cycle = &path[cycle_start..];
|
||||
let format_core = |core: &i32| (core - 1).to_string();
|
||||
|
||||
let cycle_str = cycle
|
||||
.iter()
|
||||
.map(format_core)
|
||||
.collect::<Vec<_>>()
|
||||
.join(" -> ");
|
||||
|
||||
let cycle = cycle
|
||||
.iter()
|
||||
.copied()
|
||||
.chain(std::iter::once(waiting_for))
|
||||
.collect::<Vec<_>>();
|
||||
let cycle_msg = format!("{} -> {}", cycle_str, waiting_for - 1);
|
||||
let states_msg = cycle
|
||||
.iter()
|
||||
.filter_map(|core| {
|
||||
states.get(core).map(|state| match state {
|
||||
CoreState::SendingTo(target, size) => {
|
||||
format!("core {} send {}B -> {}", core - 1, size, target - 1)
|
||||
}
|
||||
CoreState::ReceivingFrom(source, size) => {
|
||||
format!("core {} recv {}B <- {}", core - 1, size, source - 1)
|
||||
}
|
||||
CoreState::Working => format!("core {} working", core - 1),
|
||||
CoreState::Halted => format!("core {} halted", core - 1),
|
||||
})
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
.join(", ");
|
||||
|
||||
return Some(DeadlockInfo {
|
||||
cycle: cycle_msg,
|
||||
states: states_msg,
|
||||
});
|
||||
}
|
||||
|
||||
// Hit a known branch that didn't result in a cycle
|
||||
if visited.contains(&waiting_for) {
|
||||
break;
|
||||
}
|
||||
|
||||
current_core = waiting_for;
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
fn handle_wait_sync(
|
||||
core_instructions: &mut [CoreInstructions],
|
||||
events: &mut SyncEvents,
|
||||
core_result: InstructionStatus,
|
||||
) -> bool {
|
||||
match core_result {
|
||||
InstructionStatus::Sync(data) => {
|
||||
let (source, target) = data.get_core_immcore();
|
||||
let register = data.offset_select() as usize;
|
||||
events[target as usize][register] += 1;
|
||||
core_instructions[source as usize].program_counter += 1;
|
||||
true
|
||||
}
|
||||
InstructionStatus::Waiting(data) => {
|
||||
let core = data.core_indx() as usize;
|
||||
let register = data.offset_select() as usize;
|
||||
let value = data.offset_value();
|
||||
if events[core][register] >= value {
|
||||
events[core][register] -= value;
|
||||
core_instructions[core].program_counter += 1;
|
||||
true
|
||||
} else {
|
||||
false
|
||||
}
|
||||
}
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -41,16 +41,17 @@ impl SendRecv {
|
||||
}
|
||||
}
|
||||
|
||||
pub fn handle_send_recv<'a, 'b >(
|
||||
pub fn handle_send_recv<'a, 'b>(
|
||||
cpu: &'b mut CPU<'a>,
|
||||
core_instructions: & mut [CoreInstructions],
|
||||
send_recv: & mut SendRecv,
|
||||
core_instructions: &mut [CoreInstructions],
|
||||
send_recv: &mut SendRecv,
|
||||
core_result: InstructionStatus,
|
||||
) -> (bool, usize)
|
||||
where 'a : 'b
|
||||
where
|
||||
'a: 'b,
|
||||
{
|
||||
let transfer_memory = |cpu: &'b mut CPU<'a>,
|
||||
core_instructions: & mut [CoreInstructions],
|
||||
core_instructions: &mut [CoreInstructions],
|
||||
sender: Option<SendRecvInfo>,
|
||||
receiver: Option<SendRecvInfo>| {
|
||||
if let Some(sender) = sender
|
||||
@@ -58,6 +59,20 @@ where 'a : 'b
|
||||
&& sender.internal_core == receiver.external_core
|
||||
&& receiver.internal_core == sender.external_core
|
||||
{
|
||||
{
|
||||
let sender = &mut core_instructions[sender.internal_core];
|
||||
let pc = sender.program_counter;
|
||||
let inst = sender.instructions.get(pc).unwrap();
|
||||
let data = inst.data;
|
||||
TRACER.lock().unwrap().pre_send(cpu, data);
|
||||
}
|
||||
{
|
||||
let recv = &mut core_instructions[receiver.internal_core];
|
||||
let pc = recv.program_counter;
|
||||
let inst = recv.instructions.get(pc).unwrap();
|
||||
let data = inst.data;
|
||||
TRACER.lock().unwrap().pre_recv(cpu, data);
|
||||
}
|
||||
let [sender_core, reciver_core] =
|
||||
cpu.get_multiple_cores([sender.internal_core, receiver.internal_core]);
|
||||
let memory = sender_core
|
||||
@@ -119,7 +134,7 @@ where 'a : 'b
|
||||
send_recv.sending[sender] = None;
|
||||
send_recv.receiving[receiver] = None;
|
||||
}
|
||||
(transfered, receiver)
|
||||
(transfered, if transfered { receiver } else { 0 })
|
||||
}
|
||||
InstructionStatus::Reciving(instruction_data) => {
|
||||
let (core_idx, imm_core) = instruction_data.get_core_immcore();
|
||||
@@ -148,7 +163,7 @@ where 'a : 'b
|
||||
send_recv.sending[sender] = None;
|
||||
send_recv.receiving[receiver] = None;
|
||||
}
|
||||
(transfered, sender)
|
||||
(transfered, if transfered { sender } else { 0 })
|
||||
}
|
||||
_ => (false, 0),
|
||||
}
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
|
||||
use std::fs::File;
|
||||
|
||||
use crate::{
|
||||
@@ -13,64 +12,47 @@ use crate::{
|
||||
};
|
||||
use std::io::Write;
|
||||
|
||||
#[cfg(not(feature = "tracing"))]
|
||||
#[cfg(not(any(feature = "tracing", feature = "profile_time")))]
|
||||
impl Trace {
|
||||
///////////////////////////////////////////////////////////////
|
||||
/////////////////Scalar/register Instructions//////////////////
|
||||
///////////////////////////////////////////////////////////////
|
||||
|
||||
pub fn pre_sldi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
pub fn pre_sldi(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
}
|
||||
pub fn post_sldi(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_sldi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_sld(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_sld(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_sld(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_sld(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_sadd(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_sadd(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_sadd(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_sadd(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_ssub(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_ssub(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_ssub(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_ssub(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_smul(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_smul(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_smul(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_smul(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_saddi(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_saddi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_saddi(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_saddi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_smuli(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_smuli(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
|
||||
pub fn post_smuli(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_smuli(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
/////////////////////////////////////////////////////////////////
|
||||
///////////////////Matrix/vector Instructions////////////////////
|
||||
/////////////////////////////////////////////////////////////////
|
||||
|
||||
pub fn pre_setbw(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_setbw(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_setbw(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_setbw(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_mvm<F, M, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
@@ -80,7 +62,7 @@ impl Trace {
|
||||
M: UpcastDestTraits<M> + MemoryStorable + FromFloat,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.mvm_impl::<F,M,T>(cores, data, "Pre");
|
||||
self.mvm_impl::<F, M, T>(cores, data, "Pre");
|
||||
}
|
||||
|
||||
pub fn post_mvm<F, M, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
@@ -91,11 +73,15 @@ impl Trace {
|
||||
M: UpcastDestTraits<M> + MemoryStorable + FromFloat,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.mvm_impl::<F,M,T>(cores, data, "Post");
|
||||
self.mvm_impl::<F, M, T>(cores, data, "Post");
|
||||
}
|
||||
|
||||
pub fn mvm_impl<F, M, T>(&mut self, cores: &mut CPU, data: InstructionData, prefix : &'static str)
|
||||
where
|
||||
pub fn mvm_impl<F, M, T>(
|
||||
&mut self,
|
||||
cores: &mut CPU,
|
||||
data: InstructionData,
|
||||
prefix: &'static str,
|
||||
) where
|
||||
[F]: UpcastSlice<T> + UpcastSlice<M>,
|
||||
[M]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
@@ -267,39 +253,27 @@ impl Trace {
|
||||
/////////////////////////////////////////////////////////////////
|
||||
/////Communication/synchronization Instructions/////////////////
|
||||
/////////////////////////////////////////////////////////////////
|
||||
pub fn pre_ld(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_ld(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_ld(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_ld(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_st(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_st(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_st(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_st(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_lldi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_lldi(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_lldi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_lldi(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_lmv(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_lmv(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_lmv(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_lmv(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_send(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_send(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_send(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_send(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn pre_recv(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn pre_recv(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
|
||||
pub fn post_recv(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
}
|
||||
pub fn post_recv(&mut self, cores: &mut CPU, data: InstructionData) {}
|
||||
}
|
||||
|
||||
@@ -1,52 +1,30 @@
|
||||
mod tracing_isa;
|
||||
mod disable;
|
||||
mod pretty_print;
|
||||
use std::{fs::File, path::{ PathBuf}};
|
||||
use std::sync::{LazyLock, Mutex};
|
||||
#[cfg(feature = "profile_time")]
|
||||
mod profile;
|
||||
|
||||
#[cfg(feature = "profile_time")]
|
||||
use profile::Trace;
|
||||
|
||||
#[cfg(feature = "tracing")]
|
||||
mod trace;
|
||||
#[cfg(feature = "tracing")]
|
||||
use trace::Trace;
|
||||
|
||||
use crate::Executable;
|
||||
#[cfg(not(any(feature = "tracing", feature = "profile_time")))]
|
||||
use std::path::PathBuf;
|
||||
use std::sync::{LazyLock, Mutex};
|
||||
|
||||
#[cfg(feature = "tracing")]
|
||||
pub struct Trace {
|
||||
out_files : Vec<File>
|
||||
}
|
||||
#[cfg(not(any(feature = "tracing", feature = "profile_time")))]
|
||||
pub struct Trace {}
|
||||
|
||||
|
||||
#[cfg(feature = "tracing")]
|
||||
#[cfg(not(any(feature = "tracing", feature = "profile_time")))]
|
||||
impl Trace {
|
||||
fn new() -> Self {
|
||||
Self { out_files : Vec::new()}
|
||||
Self {}
|
||||
}
|
||||
|
||||
|
||||
pub fn init(&mut self, num_core : usize , mut path : PathBuf) {
|
||||
path.pop();
|
||||
for i in 0..num_core {
|
||||
path.push(format!("TraceCore{}", i));
|
||||
let file = File::create(&path).expect("Can not create file");
|
||||
self.out_files.push(file);
|
||||
path.pop();
|
||||
}
|
||||
}
|
||||
pub fn init(&mut self, num_core: usize, path: PathBuf) {}
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "tracing"))]
|
||||
pub struct Trace {
|
||||
}
|
||||
|
||||
|
||||
#[cfg(not(feature = "tracing"))]
|
||||
impl Trace {
|
||||
fn new() -> Self {
|
||||
Self { }
|
||||
}
|
||||
|
||||
|
||||
pub fn init(&mut self, num_core : usize, path : PathBuf ) {
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
pub static TRACER: LazyLock<Mutex<Trace>> = LazyLock::new(|| { Trace::new().into()});
|
||||
|
||||
pub static TRACER: LazyLock<Mutex<Trace>> = LazyLock::new(|| Trace::new().into());
|
||||
|
||||
@@ -0,0 +1,73 @@
|
||||
use std::{collections::HashMap, path::PathBuf, time::Instant};
|
||||
|
||||
use crate::tracing::profile::profile_analysis::{
|
||||
analyze_timings, generate_interactive_report, print_textual_report,
|
||||
};
|
||||
|
||||
pub mod profile_analysis;
|
||||
pub mod profile_isa;
|
||||
|
||||
pub struct Trace {
|
||||
instruction_times: HashMap<String, Vec<(u128, u128)>>,
|
||||
core_start_time: HashMap<usize, Option<Instant>>,
|
||||
start_time: Instant,
|
||||
}
|
||||
|
||||
impl Trace {
|
||||
pub fn new() -> Self {
|
||||
let mut instruction_times = HashMap::new();
|
||||
instruction_times.insert("sldi".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("sld".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("sadd".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("ssub".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("smul".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("saddi".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("smuli".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("setbw".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("mvmul".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vvadd".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vvsub".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vvmul".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vvdmul".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vvmax".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vvsll".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vvsra".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vavg".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vrelu".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vtanh".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vsigm".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vsoftmax".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vmv".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vrsu".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("vrsl".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("ld".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("st".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("lldi".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("lmv".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("send".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("recv".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("wait".to_string(), Vec::with_capacity(20000));
|
||||
instruction_times.insert("sync".to_string(), Vec::with_capacity(20000));
|
||||
Self {
|
||||
instruction_times,
|
||||
core_start_time: HashMap::new(),
|
||||
start_time: Instant::now(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn init(&mut self, num_core: usize, path: PathBuf) {
|
||||
for i in 0..num_core {
|
||||
self.core_start_time.insert(i, None);
|
||||
}
|
||||
}
|
||||
|
||||
pub fn report(&self) {
|
||||
let res = analyze_timings(&self.instruction_times);
|
||||
print_textual_report(&res);
|
||||
generate_interactive_report(
|
||||
&self.instruction_times,
|
||||
&["mvmul", "recv"],
|
||||
"/tmp/report.html",
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,192 @@
|
||||
use comfy_table::{Cell, Table, modifiers::UTF8_ROUND_CORNERS, presets::UTF8_FULL};
|
||||
use statrs::statistics::{Data, Distribution, Max, Min, OrderStatistics};
|
||||
use std::collections::HashMap;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct InstructionStats {
|
||||
pub name: String,
|
||||
pub count: usize,
|
||||
pub total_time: u128,
|
||||
pub min: f64,
|
||||
pub max: f64,
|
||||
pub mean: f64,
|
||||
pub median: f64,
|
||||
pub std_dev: f64,
|
||||
pub cv: f64,
|
||||
pub p95: f64,
|
||||
pub p99: f64,
|
||||
pub skewness: f64,
|
||||
pub kurtosis: f64,
|
||||
}
|
||||
|
||||
fn format_time(ns: f64) -> String {
|
||||
if ns.is_nan() {
|
||||
return "NaN".to_string();
|
||||
}
|
||||
|
||||
if ns >= 1_000_000_000.0 {
|
||||
format!("{:.2} s", ns / 1_000_000_000.0)
|
||||
} else if ns >= 1_000_000.0 {
|
||||
format!("{:.2} ms", ns / 1_000_000.0)
|
||||
} else if ns >= 1_000.0 {
|
||||
format!("{:.2} µs", ns / 1_000.0)
|
||||
} else {
|
||||
format!("{:.2} ns", ns)
|
||||
}
|
||||
}
|
||||
|
||||
fn calculate_skewness_kurtosis(times: &[f64], mean: f64, std_dev: f64) -> (f64, f64) {
|
||||
let n = times.len() as f64;
|
||||
|
||||
if n < 4.0 || std_dev == 0.0 {
|
||||
return (f64::NAN, f64::NAN);
|
||||
}
|
||||
|
||||
let mut sum_m3 = 0.0;
|
||||
let mut sum_m4 = 0.0;
|
||||
|
||||
for &x in times {
|
||||
let deviation = x - mean;
|
||||
sum_m3 += deviation.powi(3);
|
||||
sum_m4 += deviation.powi(4);
|
||||
}
|
||||
|
||||
let m3 = sum_m3 / n;
|
||||
let m4 = sum_m4 / n;
|
||||
|
||||
let skewness = m3 / std_dev.powi(3);
|
||||
let kurtosis = (m4 / std_dev.powi(4)) - 3.0;
|
||||
|
||||
(skewness, kurtosis)
|
||||
}
|
||||
|
||||
pub fn analyze_timings(timings: &HashMap<String, Vec<(u128, u128)>>) -> Vec<InstructionStats> {
|
||||
let mut results = Vec::new();
|
||||
|
||||
for (instruction, times) in timings {
|
||||
let count = times.len();
|
||||
if count == 0 {
|
||||
continue;
|
||||
}
|
||||
|
||||
// Extract ONLY the duration (the second element of the tuple) for stats
|
||||
let durations: Vec<u128> = times.iter().map(|&(_, duration)| duration).collect();
|
||||
let total_time: u128 = durations.iter().sum();
|
||||
|
||||
let f64_times: Vec<f64> = durations.iter().map(|&t| t as f64).collect();
|
||||
let mut data = Data::new(f64_times.clone());
|
||||
|
||||
let mean = data.mean().unwrap_or(0.0);
|
||||
let std_dev = data.std_dev().unwrap_or(0.0);
|
||||
let cv = if mean > 0.0 { std_dev / mean } else { 0.0 };
|
||||
|
||||
let (skewness, kurtosis) = calculate_skewness_kurtosis(&f64_times, mean, std_dev);
|
||||
|
||||
results.push(InstructionStats {
|
||||
name: instruction.clone(),
|
||||
count,
|
||||
total_time,
|
||||
min: data.min(),
|
||||
max: data.max(),
|
||||
mean,
|
||||
median: data.median(),
|
||||
std_dev,
|
||||
cv,
|
||||
p95: data.percentile(95),
|
||||
p99: data.percentile(99),
|
||||
skewness,
|
||||
kurtosis,
|
||||
});
|
||||
}
|
||||
|
||||
results.sort_by(|a, b| b.mean.partial_cmp(&a.mean).unwrap());
|
||||
results
|
||||
}
|
||||
|
||||
pub fn print_textual_report(stats: &[InstructionStats]) {
|
||||
let mut table = Table::new();
|
||||
table
|
||||
.load_preset(UTF8_FULL)
|
||||
.apply_modifier(UTF8_ROUND_CORNERS)
|
||||
.set_header(vec![
|
||||
"Instruction",
|
||||
"Count",
|
||||
"Total Time",
|
||||
"Mean",
|
||||
"Median",
|
||||
"Min",
|
||||
"Max",
|
||||
"P95",
|
||||
"P99",
|
||||
"StdDev",
|
||||
"CV",
|
||||
"Skewness",
|
||||
"Kurtosis",
|
||||
]);
|
||||
|
||||
for stat in stats {
|
||||
table.add_row(vec![
|
||||
Cell::new(&stat.name),
|
||||
Cell::new(stat.count.to_string()),
|
||||
Cell::new(format_time(stat.total_time as f64)), // Cast u128 to f64 for formatting
|
||||
Cell::new(format_time(stat.mean)),
|
||||
Cell::new(format_time(stat.median)),
|
||||
Cell::new(format_time(stat.min)),
|
||||
Cell::new(format_time(stat.max)),
|
||||
Cell::new(format_time(stat.p95)),
|
||||
Cell::new(format_time(stat.p99)),
|
||||
Cell::new(format_time(stat.std_dev)),
|
||||
Cell::new(format!("{:.3}", stat.cv)),
|
||||
Cell::new(format!("{:.2}", stat.skewness)),
|
||||
Cell::new(format!("{:.2}", stat.kurtosis)),
|
||||
]);
|
||||
}
|
||||
|
||||
println!("{table}");
|
||||
}
|
||||
|
||||
pub fn generate_interactive_report(
|
||||
timings: &HashMap<String, Vec<(u128, u128)>>,
|
||||
instructions_to_plot: &[&str], // <-- NEW: Only plot these
|
||||
file_path: &str,
|
||||
) {
|
||||
use plotly::common::{Line, Marker, Mode};
|
||||
use plotly::layout::{Axis, Layout};
|
||||
use plotly::{Plot, Scatter};
|
||||
use std::collections::HashMap;
|
||||
let mut plot = Plot::new();
|
||||
|
||||
for &instruction_name in instructions_to_plot {
|
||||
// Only proceed if the instruction exists in our timings map
|
||||
if let Some(times) = timings.get(instruction_name) {
|
||||
let x_axis: Vec<f64> = times.iter().map(|&(ts, _)| ts as f64).collect();
|
||||
let y_axis: Vec<f64> = times.iter().map(|&(_, dur)| dur as f64).collect();
|
||||
|
||||
let text_array: Vec<String> = times
|
||||
.iter()
|
||||
.map(|&(_, dur)| format_time(dur as f64))
|
||||
.collect();
|
||||
|
||||
let trace = Scatter::new(x_axis, y_axis)
|
||||
.name(instruction_name)
|
||||
.mode(Mode::LinesMarkers)
|
||||
.marker(Marker::new().size(4).opacity(0.6))
|
||||
.line(Line::new().width(1.0))
|
||||
.text_array(text_array)
|
||||
.hover_info(plotly::common::HoverInfo::All);
|
||||
|
||||
plot.add_trace(trace);
|
||||
}
|
||||
}
|
||||
|
||||
let layout = Layout::new()
|
||||
.title(plotly::common::Title::new(
|
||||
"Simulator Timeline: Top Offenders",
|
||||
))
|
||||
.x_axis(Axis::new().title(plotly::common::Title::new("Absolute Time (ns)")))
|
||||
.y_axis(Axis::new().title(plotly::common::Title::new("Execution Duration")));
|
||||
|
||||
plot.set_layout(layout);
|
||||
plot.write_html(file_path);
|
||||
println!("🌐 Interactive timeline saved to {}", file_path);
|
||||
}
|
||||
@@ -0,0 +1,365 @@
|
||||
use crate::{
|
||||
cpu::CPU,
|
||||
instruction_set::instruction_data::InstructionData,
|
||||
memory_manager::{
|
||||
MemoryStorable,
|
||||
type_traits::{FromFloat, UpcastDestTraits, UpcastSlice},
|
||||
},
|
||||
tracing::Trace,
|
||||
utility::{add_offset_r1, add_offset_rd},
|
||||
};
|
||||
use std::io::Write;
|
||||
use std::time::Instant;
|
||||
|
||||
#[cfg(feature = "profile_time")]
|
||||
impl Trace {
|
||||
///////////////////////////////////////////////////////////////
|
||||
/////////////////Scalar/register Instructions//////////////////
|
||||
///////////////////////////////////////////////////////////////
|
||||
|
||||
fn pre_impl(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
let (core_indx, rd, imm) = data.get_core_rd_imm();
|
||||
let core_indx = core_indx as usize;
|
||||
if self.core_start_time.get(&core_indx).unwrap().is_none() {
|
||||
self.core_start_time.insert(core_indx, Some(Instant::now()));
|
||||
}
|
||||
}
|
||||
|
||||
fn post_impl(&mut self, cores: &mut CPU, data: InstructionData, name: &'static str) {
|
||||
let (core_indx, rd, imm) = data.get_core_rd_imm();
|
||||
let core_indx = core_indx as usize;
|
||||
let Self {
|
||||
instruction_times,
|
||||
core_start_time,
|
||||
start_time,
|
||||
} = self;
|
||||
let now = Instant::now();
|
||||
instruction_times.get_mut(name).unwrap().push((
|
||||
now.duration_since(*start_time).as_nanos(),
|
||||
now.duration_since(core_start_time[&core_indx].unwrap())
|
||||
.as_nanos(),
|
||||
));
|
||||
self.core_start_time.insert(core_indx, None);
|
||||
}
|
||||
|
||||
pub fn pre_sldi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_sldi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "sldi");
|
||||
}
|
||||
|
||||
pub fn pre_sld(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_sld(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "sld");
|
||||
}
|
||||
|
||||
pub fn pre_sadd(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_sadd(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "sadd");
|
||||
}
|
||||
|
||||
pub fn pre_ssub(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_ssub(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "ssub");
|
||||
}
|
||||
|
||||
pub fn pre_smul(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_smul(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "smul");
|
||||
}
|
||||
|
||||
pub fn pre_saddi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_saddi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "saddi");
|
||||
}
|
||||
|
||||
pub fn pre_smuli(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_smuli(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "smuli");
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////
|
||||
///////////////////Matrix/vector Instructions////////////////////
|
||||
/////////////////////////////////////////////////////////////////
|
||||
|
||||
pub fn pre_setbw(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_setbw(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "setbw");
|
||||
}
|
||||
|
||||
pub fn pre_mvm<F, M, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T> + UpcastSlice<M>,
|
||||
[M]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
M: UpcastDestTraits<M> + MemoryStorable + FromFloat,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_mvm<F, M, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T> + UpcastSlice<M>,
|
||||
[M]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
M: UpcastDestTraits<M> + MemoryStorable + FromFloat,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.post_impl(cores, data, "mvmul");
|
||||
}
|
||||
|
||||
pub fn pre_vvadd<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_vvadd<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.post_impl(cores, data, "vvadd");
|
||||
}
|
||||
|
||||
pub fn pre_vvsub<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_vvsub<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.post_impl(cores, data, "vvsub");
|
||||
}
|
||||
|
||||
pub fn pre_vvmul<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_vvmul<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.post_impl(cores, data, "vvmul");
|
||||
}
|
||||
|
||||
pub fn pre_vvdmul<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_vvdmul<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.post_impl(cores, data, "vvdmul");
|
||||
}
|
||||
|
||||
pub fn pre_vvmax<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_vvmax<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.post_impl(cores, data, "vvmax");
|
||||
}
|
||||
|
||||
pub fn pre_vavg<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_vavg<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
self.post_impl(cores, data, "vavg");
|
||||
}
|
||||
|
||||
pub fn pre_vrelu<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_vrelu<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
self.post_impl(cores, data, "vrelu");
|
||||
}
|
||||
|
||||
pub fn pre_vtanh<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_vtanh<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
self.post_impl(cores, data, "vtanh");
|
||||
}
|
||||
|
||||
pub fn pre_vsigm<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_vsigm<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
self.post_impl(cores, data, "vsigm");
|
||||
}
|
||||
|
||||
pub fn pre_vsoftmax<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_vsoftmax<F, T>(&mut self, cores: &mut CPU, data: InstructionData)
|
||||
where
|
||||
[F]: UpcastSlice<T>,
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable + From<f32>,
|
||||
{
|
||||
self.post_impl(cores, data, "vsoftmax");
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////
|
||||
/////Communication/synchronization Instructions/////////////////
|
||||
/////////////////////////////////////////////////////////////////
|
||||
|
||||
pub fn pre_ld(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
pub fn post_ld(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "ld");
|
||||
}
|
||||
|
||||
pub fn pre_st(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_st(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "st");
|
||||
}
|
||||
|
||||
pub fn pre_lldi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_lldi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "lldi");
|
||||
}
|
||||
|
||||
pub fn pre_lmv(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_lmv(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "lmv");
|
||||
}
|
||||
|
||||
pub fn pre_send(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_send(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "send");
|
||||
}
|
||||
|
||||
pub fn pre_recv(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.pre_impl(cores, data);
|
||||
}
|
||||
|
||||
pub fn post_recv(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
self.post_impl(cores, data, "recv");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,26 @@
|
||||
use std::{fs::File, path::PathBuf};
|
||||
|
||||
pub mod pretty_print;
|
||||
pub mod tracing_isa;
|
||||
|
||||
pub struct Trace {
|
||||
out_files: Vec<File>,
|
||||
}
|
||||
|
||||
impl Trace {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
out_files: Vec::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn init(&mut self, num_core: usize, mut path: PathBuf) {
|
||||
path.pop();
|
||||
for i in 0..num_core {
|
||||
path.push(format!("TraceCore{}", i));
|
||||
let file = File::create(&path).expect("Can not create file");
|
||||
self.out_files.push(file);
|
||||
path.pop();
|
||||
}
|
||||
}
|
||||
}
|
||||
+34
-36
@@ -1,5 +1,5 @@
|
||||
use crate::tracing::pretty_print;
|
||||
use std::fs::File;
|
||||
use crate::{tracing::trace::pretty_print, utility::add_offset_r2};
|
||||
use std::{fs::File, mem::size_of};
|
||||
|
||||
use crate::{
|
||||
cpu::CPU,
|
||||
@@ -13,7 +13,6 @@ use crate::{
|
||||
};
|
||||
use std::io::Write;
|
||||
|
||||
#[cfg(feature = "tracing")]
|
||||
impl Trace {
|
||||
///////////////////////////////////////////////////////////////
|
||||
/////////////////Scalar/register Instructions//////////////////
|
||||
@@ -284,8 +283,6 @@ impl Trace {
|
||||
M: UpcastDestTraits<M> + MemoryStorable + FromFloat,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
use crate::tracing::pretty_print;
|
||||
|
||||
let (core_indx, rd, r1, mbiw, relu, group) = data.get_core_rd_r1_mbiw_immrelu_immgroup();
|
||||
let file: &mut File = self
|
||||
.out_files
|
||||
@@ -339,7 +336,7 @@ impl Trace {
|
||||
);
|
||||
pretty_print::print_slice::<_, F>(file, loads[0], 30);
|
||||
writeln!(file, "\tCrossbar[{}:{}](B): ", 0, crossbar_byte_size);
|
||||
pretty_print::print_slice::<_,F>(file, matrix, 30);
|
||||
pretty_print::print_slice::<_, F>(file, matrix, 30);
|
||||
writeln!(
|
||||
file,
|
||||
"\tLocal[{}:{}](out): ",
|
||||
@@ -358,8 +355,6 @@ impl Trace {
|
||||
T: UpcastDestTraits<T> + MemoryStorable,
|
||||
F: UpcastDestTraits<F> + MemoryStorable,
|
||||
{
|
||||
use crate::{tracing::pretty_print, utility::add_offset_r2};
|
||||
|
||||
let (core_indx, rd, r1, r2, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let file: &mut File = self
|
||||
@@ -385,7 +380,15 @@ impl Trace {
|
||||
writeln!(file, "\trs1({}): {}", r1, r1_val);
|
||||
writeln!(file, "\trs2({}): {}", r2, r2_val);
|
||||
writeln!(file, "{} Immediate:", prefix);
|
||||
writeln!(file, "\tLoad Length: {} bytes", imm_len);
|
||||
let element_count: usize = imm_len.try_into().expect("imm_len can not be negative");
|
||||
let byte_len = element_count
|
||||
.checked_mul(size_of::<F>())
|
||||
.expect("vector byte length overflow");
|
||||
writeln!(
|
||||
file,
|
||||
"\tLoad Length: {} elements ({} bytes)",
|
||||
element_count, byte_len
|
||||
);
|
||||
writeln!(
|
||||
file,
|
||||
"\toffset_select: {} offset_value: {}",
|
||||
@@ -394,23 +397,22 @@ impl Trace {
|
||||
let r1_final = add_offset_r1(r1_val, offset_select, offset_value);
|
||||
let r2_final = add_offset_r2(r2_val, offset_select, offset_value);
|
||||
let rd_final = add_offset_rd(rd_val, offset_select, offset_value);
|
||||
let imm_len: usize = imm_len.try_into().expect("imm_len can not be negative");
|
||||
let loads = core
|
||||
.reserve_load(r1_final, imm_len)
|
||||
.reserve_load(r1_final, byte_len)
|
||||
.unwrap()
|
||||
.reserve_load(r2_final, imm_len)
|
||||
.reserve_load(r2_final, byte_len)
|
||||
.unwrap()
|
||||
.reserve_load(rd_final, imm_len)
|
||||
.reserve_load(rd_final, byte_len)
|
||||
.unwrap()
|
||||
.execute_load::<F>()
|
||||
.unwrap();
|
||||
writeln!(file, "{} Memory:", prefix);
|
||||
write!(file, "\tLocal[{}:{}](A): ", r1_final, r1_final + imm_len);
|
||||
pretty_print::print_slice::<_,f32>(file, loads[0], 30);
|
||||
write!(file, "\tLocal[{}:{}](B): ", r2_final , r2_final + imm_len);
|
||||
pretty_print::print_slice::<_,f32>(file, loads[1], 30);
|
||||
write!(file, "\tLocal[{}:{}](out): ", rd_final, rd_final+ imm_len);
|
||||
pretty_print::print_slice::<_,f32>(file, loads[2], 30);
|
||||
write!(file, "\tLocal[{}:{}](A): ", r1_final, r1_final + byte_len);
|
||||
pretty_print::print_slice::<_, f32>(file, loads[0], 30);
|
||||
write!(file, "\tLocal[{}:{}](B): ", r2_final, r2_final + byte_len);
|
||||
pretty_print::print_slice::<_, f32>(file, loads[1], 30);
|
||||
write!(file, "\tLocal[{}:{}](out): ", rd_final, rd_final + byte_len);
|
||||
pretty_print::print_slice::<_, f32>(file, loads[2], 30);
|
||||
if prefix == "Post" {
|
||||
writeln!(file, "\n###############################################\n");
|
||||
}
|
||||
@@ -990,8 +992,6 @@ impl Trace {
|
||||
/////////////////////////////////////////////////////////////////
|
||||
|
||||
pub fn ld_impl(&mut self, cores: &mut CPU, data: InstructionData, prefix: &'static str) {
|
||||
use crate::tracing::pretty_print;
|
||||
|
||||
let (core, rd, r1, _, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let file: &mut File = self
|
||||
@@ -1026,9 +1026,9 @@ impl Trace {
|
||||
let core_memory = core.load::<u8>(rd_val, imm_len).unwrap();
|
||||
writeln!(file, "{} Memory:", prefix);
|
||||
writeln!(file, "\tHost[{}:{}]: ", r1_val, r1_val + imm_len as usize,);
|
||||
pretty_print::print_slice::<_,f32>(file, global_memory[0], 30);
|
||||
pretty_print::print_slice::<_, f32>(file, global_memory[0], 30);
|
||||
writeln!(file, "\tLocal[{}:{}]: ", rd_val, rd_val + imm_len as usize,);
|
||||
pretty_print::print_slice::<_,f32>(file, core_memory[0], 30);
|
||||
pretty_print::print_slice::<_, f32>(file, core_memory[0], 30);
|
||||
|
||||
if prefix == "Post" {
|
||||
writeln!(file, "\n###############################################\n");
|
||||
@@ -1044,8 +1044,6 @@ impl Trace {
|
||||
}
|
||||
|
||||
pub fn st_impl(&mut self, cores: &mut CPU, data: InstructionData, prefix: &'static str) {
|
||||
use crate::tracing::pretty_print;
|
||||
|
||||
let (core, rd, r1, _, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let file: &mut File = self
|
||||
@@ -1080,9 +1078,9 @@ impl Trace {
|
||||
let global_memory = host.load::<u8>(rd_val, imm_len).unwrap();
|
||||
writeln!(file, "{} Memory:", prefix);
|
||||
writeln!(file, "\tLocal[{}:{}]: ", r1_val, r1_val + imm_len as usize,);
|
||||
pretty_print::print_slice::<_,f32>(file, core_memory[0], 30);
|
||||
pretty_print::print_slice::<_, f32>(file, core_memory[0], 30);
|
||||
writeln!(file, "\tHost[{}:{}]: ", rd_val, rd_val + imm_len as usize,);
|
||||
pretty_print::print_slice::<_,f32>(file, global_memory[0], 30);
|
||||
pretty_print::print_slice::<_, f32>(file, global_memory[0], 30);
|
||||
|
||||
if prefix == "Post" {
|
||||
writeln!(file, "\n###############################################\n");
|
||||
@@ -1097,7 +1095,6 @@ impl Trace {
|
||||
self.st_impl(cores, data, "Pre");
|
||||
}
|
||||
|
||||
|
||||
pub fn pre_lldi(&mut self, cores: &mut CPU, data: InstructionData) {
|
||||
let (core, rd, imm) = data.get_core_rd_imm();
|
||||
let file: &mut File = self
|
||||
@@ -1137,9 +1134,7 @@ impl Trace {
|
||||
// Ok(InstructionStatus::Completed)
|
||||
}
|
||||
|
||||
fn lmv_impl (&mut self, cores: &mut CPU, data: InstructionData, prefix: &'static str) {
|
||||
use crate::tracing::pretty_print;
|
||||
|
||||
fn lmv_impl(&mut self, cores: &mut CPU, data: InstructionData, prefix: &'static str) {
|
||||
let (core, rd, r1, _, imm_len, offset_select, offset_value) =
|
||||
data.get_core_rd_r1_r2_immlen_offset();
|
||||
let file: &mut File = self
|
||||
@@ -1171,14 +1166,17 @@ impl Trace {
|
||||
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 core_memory = core
|
||||
.reserve_load(r1_val, imm_len).unwrap()
|
||||
.reserve_load(rd_val, imm_len).unwrap()
|
||||
.execute_load::<u8>().unwrap();
|
||||
.reserve_load(r1_val, imm_len)
|
||||
.unwrap()
|
||||
.reserve_load(rd_val, imm_len)
|
||||
.unwrap()
|
||||
.execute_load::<u8>()
|
||||
.unwrap();
|
||||
writeln!(file, "{} Memory:", prefix);
|
||||
writeln!(file, "\tLocal[{}:{}]: ", r1_val, r1_val + imm_len as usize,);
|
||||
pretty_print::print_slice::<_,f32>(file, core_memory[0], 30);
|
||||
pretty_print::print_slice::<_, f32>(file, core_memory[0], 30);
|
||||
writeln!(file, "\tLocal[{}:{}]: ", rd_val, rd_val + imm_len as usize,);
|
||||
pretty_print::print_slice::<_,f32>(file, core_memory[1], 30);
|
||||
pretty_print::print_slice::<_, f32>(file, core_memory[1], 30);
|
||||
|
||||
if prefix == "Post" {
|
||||
writeln!(file, "\n###############################################\n");
|
||||
@@ -1,39 +1,90 @@
|
||||
use anyhow::{Context, Result};
|
||||
use std::{fmt::Debug, mem::transmute};
|
||||
|
||||
use crate::memory_manager::type_traits::TryToUsize;
|
||||
|
||||
|
||||
fn add_offset_impl(address: usize, offset_select : i32, offset_value : i32, id:i32) -> usize{
|
||||
assert!(offset_select == 1 || offset_select == 2 || offset_select == 4 || offset_value == 0, "offset_select not a bit field");
|
||||
let offset_value = (offset_select & id) * offset_value;
|
||||
if offset_value > 0 {
|
||||
address + offset_value as usize
|
||||
} else {
|
||||
address - offset_value as usize
|
||||
}
|
||||
pub trait AddressArg {
|
||||
fn to_address_usize(self) -> Result<usize>;
|
||||
}
|
||||
|
||||
|
||||
pub fn add_offset_rd(address: impl TryToUsize, offset_select : i32, offset_value : i32) -> usize
|
||||
{
|
||||
let address = address.try_into().expect("address can not be negative");
|
||||
add_offset_impl(address, offset_select, offset_value, 4)
|
||||
impl AddressArg for usize {
|
||||
fn to_address_usize(self) -> Result<usize> {
|
||||
Ok(self)
|
||||
}
|
||||
}
|
||||
|
||||
pub fn add_offset_r1(address: impl TryToUsize, offset_select : i32, offset_value : i32) -> usize
|
||||
{
|
||||
let address = address.try_into().expect("address can not be negative");
|
||||
impl AddressArg for u32 {
|
||||
fn to_address_usize(self) -> Result<usize> {
|
||||
Ok(self as usize)
|
||||
}
|
||||
}
|
||||
|
||||
impl AddressArg for u64 {
|
||||
fn to_address_usize(self) -> Result<usize> {
|
||||
usize::try_from(self).context("address does not fit in usize")
|
||||
}
|
||||
}
|
||||
|
||||
impl AddressArg for i32 {
|
||||
fn to_address_usize(self) -> Result<usize> {
|
||||
Ok(self as u32 as usize)
|
||||
}
|
||||
}
|
||||
|
||||
impl AddressArg for i64 {
|
||||
fn to_address_usize(self) -> Result<usize> {
|
||||
usize::try_from(self).context("address can not be negative")
|
||||
}
|
||||
}
|
||||
|
||||
fn address_to_usize(address: i32) -> usize {
|
||||
address as u32 as usize
|
||||
}
|
||||
|
||||
fn add_offset_impl(address: usize, offset_select: i32, offset_value: i32, id: i32) -> usize {
|
||||
assert!(
|
||||
(0..=7).contains(&offset_select),
|
||||
"offset_select is not a 3-bit field"
|
||||
);
|
||||
if offset_select & id == 0 {
|
||||
address
|
||||
} else {
|
||||
address
|
||||
.checked_add_signed(offset_value as isize)
|
||||
.expect("address offset overflow")
|
||||
}
|
||||
}
|
||||
|
||||
pub fn add_offset_rd(address: i32, offset_select: i32, offset_value: i32) -> usize {
|
||||
let address = address_to_usize(address);
|
||||
add_offset_impl(address, offset_select, offset_value, 1)
|
||||
}
|
||||
|
||||
pub fn add_offset_r2(address: impl TryToUsize, offset_select : i32, offset_value : i32) -> usize
|
||||
{
|
||||
let address = address.try_into().expect("address can not be negative");
|
||||
pub fn add_offset_r1(address: i32, offset_select: i32, offset_value: i32) -> usize {
|
||||
let address = address_to_usize(address);
|
||||
add_offset_impl(address, offset_select, offset_value, 2)
|
||||
}
|
||||
|
||||
|
||||
pub fn pack_float_in_i32(val : impl TryInto<f32>) -> i32 {
|
||||
let val = val.try_into().unwrap_or_else( |x| panic!("Cannot parse into f32"));
|
||||
f32::to_bits(val).cast_signed()
|
||||
pub fn add_offset_r2(address: i32, offset_select: i32, offset_value: i32) -> usize {
|
||||
let address = address_to_usize(address);
|
||||
add_offset_impl(address, offset_select, offset_value, 4)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{add_offset_r1, add_offset_r2, add_offset_rd};
|
||||
|
||||
#[test]
|
||||
fn offset_select_uses_rd_rs1_rs2_bit_order() {
|
||||
assert_eq!(add_offset_rd(100, 1, 7), 107);
|
||||
assert_eq!(add_offset_r1(100, 2, 7), 107);
|
||||
assert_eq!(add_offset_r2(100, 4, 7), 107);
|
||||
assert_eq!(add_offset_rd(100, 6, 7), 100);
|
||||
assert_eq!(add_offset_r1(100, 7, -7), 93);
|
||||
}
|
||||
}
|
||||
|
||||
pub fn pack_float_in_i32(val: impl TryInto<f32>) -> i32 {
|
||||
let val = val
|
||||
.try_into()
|
||||
.unwrap_or_else(|x| panic!("Cannot parse into f32"));
|
||||
f32::to_bits(val).cast_signed()
|
||||
}
|
||||
|
||||
@@ -1,8 +1,13 @@
|
||||
use std::path::Path;
|
||||
|
||||
use pimcore::{Executable, cpu::CPU, instruction_set::{InstructionsBuilder, instruction_data::InstructionDataBuilder, isa::*}};
|
||||
use pimcore::{
|
||||
Executable,
|
||||
cpu::crossbar::Crossbar,
|
||||
instruction_set::{InstructionsBuilder, instruction_data::InstructionDataBuilder, isa::*},
|
||||
memory_manager::CoreMemory,
|
||||
};
|
||||
|
||||
fn simple_read(path: &Path) -> Vec<f32> {
|
||||
fn simple_read(path: &Path) -> Vec<f32> {
|
||||
if !path.exists() {
|
||||
panic!("{:?} not exists", path)
|
||||
}
|
||||
@@ -14,17 +19,13 @@ fn simple_read(path: &Path) -> Vec<f32> {
|
||||
}
|
||||
|
||||
/// mvmul Test
|
||||
fn mvmul_f32(err: &str)
|
||||
where
|
||||
{
|
||||
let mut cpu = CPU::new(0);
|
||||
cpu.reserve_crossbar(1, 1024 * size_of::<f32>(), 1024);
|
||||
let (memory, crossbars) = cpu.host().get_memory_crossbar();
|
||||
let matrix = simple_read(Path::new("B.txt")) ;
|
||||
|
||||
|
||||
crossbars.get_mut(0).unwrap().execute_store( &matrix).unwrap();
|
||||
let vector = simple_read(Path::new("A.txt"));
|
||||
fn mvmul_f32(err: &str) {
|
||||
let matrix = simple_read(Path::new("tests/B.txt"));
|
||||
let mut crossbar = Crossbar::new(1024 * size_of::<f32>(), 1024, CoreMemory::new());
|
||||
crossbar.execute_store(&matrix).unwrap();
|
||||
let mut cpu = pimcore::cpu::CPU::new(0, vec![vec![&crossbar]]);
|
||||
let (memory, _) = cpu.host().get_memory_crossbar();
|
||||
let vector = simple_read(Path::new("tests/A.txt"));
|
||||
memory.execute_store(0, &vector).unwrap();
|
||||
|
||||
let mut inst_builder = InstructionsBuilder::new();
|
||||
@@ -33,7 +34,9 @@ where
|
||||
inst_builder.make_inst(sldi, idata_build.set_rdimm(1, 0).build());
|
||||
inst_builder.make_inst(
|
||||
sldi,
|
||||
idata_build.set_rdimm(3, 1024 * size_of::<f32>() as i32).build(),
|
||||
idata_build
|
||||
.set_rdimm(3, 1024 * size_of::<f32>() as i32)
|
||||
.build(),
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
setbw,
|
||||
@@ -45,7 +48,7 @@ where
|
||||
mvmul,
|
||||
idata_build
|
||||
.set_rdr1(3, 1)
|
||||
.set_mbiw_immrelu_immgroup(8*size_of::<f32>() as i32, 0, 0)
|
||||
.set_mbiw_immrelu_immgroup(8 * size_of::<f32>() as i32, 0, 0)
|
||||
.build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
@@ -56,8 +59,11 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<f32>(1024 * size_of::<f32>(), 1024*size_of::<f32>()).unwrap()[0].iter().zip(
|
||||
simple_read(Path::new("X.txt")) ).all(|(&a,b) : (&f32, f32)| {a-b < 0.001}),
|
||||
.load::<f32>(1024 * size_of::<f32>(), 1024 * size_of::<f32>())
|
||||
.unwrap()[0]
|
||||
.iter()
|
||||
.zip(simple_read(Path::new("tests/X.txt")))
|
||||
.all(|(&a, b): (&f32, f32)| { a - b < 0.001 }),
|
||||
"Wrong result for {}",
|
||||
err
|
||||
);
|
||||
@@ -66,8 +72,4 @@ where
|
||||
#[test]
|
||||
fn mvmul_big_test() {
|
||||
mvmul_f32("mvmul_f32");
|
||||
|
||||
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
use pimcore::cpu::CPU;
|
||||
|
||||
pub fn empty_cpu(num_cores: usize) -> CPU<'static> {
|
||||
CPU::new(num_cores, vec![Vec::new(); num_cores + 1])
|
||||
}
|
||||
@@ -1,51 +1,106 @@
|
||||
use std::{fs, io::BufReader, path::Path};
|
||||
use std::{
|
||||
fs::{self, File},
|
||||
io::BufReader,
|
||||
path::{Path, PathBuf},
|
||||
};
|
||||
|
||||
use anyhow::{Context, Result};
|
||||
use pimcore::json_to_instruction::json_to_executor;
|
||||
use pimcore::{
|
||||
cpu::crossbar::Crossbar, json_to_instruction::json_to_executor, memory_manager::CoreMemory,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
fn collect_json_from_subfolders<P: AsRef<Path>>(root: P) -> Result<Vec<(Value, Vec<Value>)>> {
|
||||
fn collect_examples<P: AsRef<Path>>(root: P) -> Result<Vec<PathBuf>> {
|
||||
let mut result = Vec::new();
|
||||
for entry in fs::read_dir(root)? {
|
||||
let entry = entry.context("Root not found")?;
|
||||
let path = entry.path();
|
||||
if path.is_dir() {
|
||||
let mut cores = Vec::new();
|
||||
let mut config: Option<Value> = None;
|
||||
for sub_entry in fs::read_dir(&path)
|
||||
.with_context(|| format!("File {} not readable", path.display()))?
|
||||
{
|
||||
let sub_entry =
|
||||
sub_entry.with_context(|| format!("File {} not readable", path.display()))?;
|
||||
let sub_path = sub_entry.path();
|
||||
if sub_path.is_file()
|
||||
&& sub_path.extension().and_then(|s| s.to_str()) == Some("json")
|
||||
{
|
||||
let file = fs::File::open(&sub_path)
|
||||
.with_context(|| format!("Subpath {} not opened", sub_path.display()))?;
|
||||
let reader = BufReader::new(file);
|
||||
let val: Value = serde_json::from_reader(reader).with_context(|| format!(
|
||||
"Serde reader fail for subpath {}",
|
||||
sub_path.display()
|
||||
))?;
|
||||
if sub_path.file_name().unwrap() == "config.json" {
|
||||
config = Some(val);
|
||||
} else {
|
||||
cores.push(val);
|
||||
}
|
||||
}
|
||||
}
|
||||
result.push((config.unwrap(), cores));
|
||||
result.push(path);
|
||||
}
|
||||
}
|
||||
Ok(result)
|
||||
}
|
||||
|
||||
fn core_sort_key(path: &Path) -> i32 {
|
||||
let stem = path.file_stem().unwrap().to_str().unwrap();
|
||||
stem[5..].parse::<i32>().unwrap()
|
||||
}
|
||||
|
||||
fn crossbar_sort_key(path: &Path) -> i32 {
|
||||
let stem = path.file_stem().unwrap().to_str().unwrap();
|
||||
stem[9..].parse::<i32>().unwrap()
|
||||
}
|
||||
|
||||
fn load_crossbars(folder: &Path, config: &Value) -> Result<Vec<Vec<Crossbar>>> {
|
||||
let xbar_size = config.get("xbar_size").unwrap().as_array().unwrap();
|
||||
let rows = xbar_size[0].as_i64().unwrap() as usize;
|
||||
let cols = xbar_size[1].as_i64().unwrap() as usize;
|
||||
let core_cnt = config.get("core_cnt").unwrap().as_i64().unwrap() as usize;
|
||||
let mut owned_crossbars = Vec::with_capacity(core_cnt + 1);
|
||||
owned_crossbars.push(Vec::new());
|
||||
|
||||
for core_idx in 0..core_cnt {
|
||||
let core_folder = folder.join(format!("core_{core_idx}"));
|
||||
let mut core_crossbars = Vec::new();
|
||||
if core_folder.is_dir() {
|
||||
let mut paths: Vec<_> = fs::read_dir(&core_folder)?
|
||||
.map(|entry| entry.map(|entry| entry.path()))
|
||||
.collect::<std::io::Result<Vec<_>>>()?;
|
||||
paths.sort_by_cached_key(|path| crossbar_sort_key(path));
|
||||
for path in paths {
|
||||
if path.extension().and_then(|ext| ext.to_str()) != Some("bin") {
|
||||
continue;
|
||||
}
|
||||
let bytes = fs::read(&path)
|
||||
.with_context(|| format!("failed to read crossbar {}", path.display()))?;
|
||||
let mut crossbar = Crossbar::new(cols * 4, rows, CoreMemory::new());
|
||||
crossbar.execute_store(&bytes)?;
|
||||
core_crossbars.push(crossbar);
|
||||
}
|
||||
}
|
||||
owned_crossbars.push(core_crossbars);
|
||||
}
|
||||
Ok(owned_crossbars)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn json_folder_tester() {
|
||||
let examples = collect_json_from_subfolders("data").unwrap();
|
||||
for example in examples {
|
||||
let (config, cores) = example;
|
||||
json_to_executor::json_to_executor(config, cores.iter()).execute();
|
||||
let examples = collect_examples("tests/data").unwrap();
|
||||
for folder in examples {
|
||||
let config_path = folder.join("config.json");
|
||||
let config_file = File::open(&config_path).unwrap();
|
||||
let config: Value = serde_json::from_reader(BufReader::new(config_file)).unwrap();
|
||||
|
||||
let core_cnt = config.get("core_cnt").unwrap().as_i64().unwrap() as usize;
|
||||
let mut core_paths: Vec<_> = fs::read_dir(&folder)
|
||||
.unwrap()
|
||||
.map(|entry| entry.unwrap().path())
|
||||
.filter(|path| path.extension().and_then(|ext| ext.to_str()) == Some("json"))
|
||||
.filter(|path| path.file_name().unwrap() != "config.json")
|
||||
.collect();
|
||||
core_paths.sort_by_cached_key(|path| core_sort_key(path));
|
||||
assert_eq!(core_paths.len(), core_cnt);
|
||||
|
||||
let mut core_readers: Vec<_> = core_paths
|
||||
.into_iter()
|
||||
.map(|path| BufReader::new(File::open(path).unwrap()))
|
||||
.collect();
|
||||
|
||||
let owned_crossbars = load_crossbars(&folder, &config).unwrap();
|
||||
let crossbars = owned_crossbars
|
||||
.iter()
|
||||
.map(|core_crossbars| core_crossbars.iter().collect())
|
||||
.collect();
|
||||
|
||||
let mut executable =
|
||||
json_to_executor::json_to_executor(config, &mut core_readers, crossbars);
|
||||
let memory = fs::read(folder.join("memory.bin")).unwrap();
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.execute_store(0, &memory)
|
||||
.unwrap();
|
||||
executable.execute();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,110 +1,117 @@
|
||||
mod common;
|
||||
|
||||
use pimcore::{Executable, cpu::CPU, instruction_set::{InstructionType, InstructionsBuilder, instruction_data::InstructionDataBuilder, isa::*}};
|
||||
|
||||
use pimcore::{
|
||||
Executable,
|
||||
instruction_set::{
|
||||
InstructionType, InstructionsBuilder, instruction_data::InstructionDataBuilder, isa::*,
|
||||
},
|
||||
};
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "Function not found for the requested size") ]
|
||||
#[should_panic(expected = "Function not found for the requested size")]
|
||||
fn wrong_size_place_holder() {
|
||||
let cpu = CPU::new(0);
|
||||
let cpu = common::empty_cpu(0);
|
||||
let mut inst_builder = InstructionsBuilder::new();
|
||||
let mut idata_build = InstructionDataBuilder::new();
|
||||
idata_build.set_core_indx(0).fix_core_indx();
|
||||
inst_builder.make_inst(
|
||||
setbw,
|
||||
idata_build
|
||||
.set_ibiw_obiw(55, 55)
|
||||
.build(),
|
||||
);
|
||||
inst_builder.make_inst(setbw, idata_build.set_ibiw_obiw(55, 55).build());
|
||||
inst_builder.make_inst(
|
||||
vvadd,
|
||||
idata_build
|
||||
.set_rdr1r2(3, 1, 2)
|
||||
.set_imm_len(8 * size_of::<f32>() as i32)
|
||||
.build(),
|
||||
idata_build.set_rdr1r2(3, 1, 2).set_imm_len(8).build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
let mut executable = Executable::new(cpu, core_instruction);
|
||||
executable.execute();
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "Function not found for the requested size input:8 output:8")]
|
||||
fn unsupported_8_bit_vectors_do_not_alias_f32() {
|
||||
let mut inst_builder = InstructionsBuilder::new();
|
||||
let mut idata_build = InstructionDataBuilder::new();
|
||||
idata_build.set_core_indx(0).fix_core_indx();
|
||||
inst_builder.make_inst(setbw, idata_build.set_ibiw_obiw(8, 8).build());
|
||||
inst_builder.make_inst(
|
||||
vvadd,
|
||||
idata_build.set_rdr1r2(3, 1, 2).set_imm_len(8).build(),
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
fn place_holder(inst : InstructionType) {
|
||||
let mut cpu = CPU::new(0);
|
||||
fn place_holder(inst: InstructionType) {
|
||||
let mut cpu = common::empty_cpu(0);
|
||||
let mut idata_build = InstructionDataBuilder::new();
|
||||
idata_build.set_core_indx(0).fix_core_indx();
|
||||
inst(&mut cpu, idata_build.build()).unwrap();
|
||||
}
|
||||
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version") ]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version")]
|
||||
fn vvadd_placeholder() {
|
||||
place_holder(vvadd);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version") ]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version")]
|
||||
fn vvsub_placeholder() {
|
||||
place_holder(vvsub);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version") ]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version")]
|
||||
fn vvmul_placeholder() {
|
||||
place_holder(vvmul);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version") ]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version")]
|
||||
fn vvdmul_placeholder() {
|
||||
place_holder(vvdmul);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version") ]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version")]
|
||||
fn vvmax_placeholder() {
|
||||
place_holder(vvmax);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version") ]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version")]
|
||||
fn vavg_placeholder() {
|
||||
place_holder(vavg);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version") ]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version")]
|
||||
fn vrelu_placeholder() {
|
||||
place_holder(vrelu);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version") ]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version")]
|
||||
fn vtanh_placeholder() {
|
||||
place_holder(vtanh);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version") ]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version")]
|
||||
fn vsigm_placeholder() {
|
||||
place_holder(vsigm);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version") ]
|
||||
#[should_panic(expected = "You are calling a placeholder, the real call is the generic version")]
|
||||
fn mvmul_placeholder() {
|
||||
place_holder(mvmul);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic ]
|
||||
#[should_panic]
|
||||
fn vvsll_why_inst() {
|
||||
place_holder(vvsll);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[should_panic ]
|
||||
#[should_panic]
|
||||
fn vvsra_why_inst() {
|
||||
place_holder(vvsra);
|
||||
}
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
mod common;
|
||||
|
||||
use pimcore::{
|
||||
Executable,
|
||||
cpu::CPU,
|
||||
cpu::crossbar::Crossbar,
|
||||
instruction_set::{InstructionsBuilder, instruction_data::InstructionDataBuilder, isa::*},
|
||||
memory_manager::{MemoryStorable, type_traits::UpcastDestTraits},
|
||||
memory_manager::{CoreMemory, MemoryStorable, type_traits::UpcastDestTraits},
|
||||
};
|
||||
|
||||
/// VVADD Test
|
||||
@@ -11,7 +13,7 @@ where
|
||||
F: From<f32> + std::fmt::Debug + PartialEq<F> + MemoryStorable,
|
||||
T: From<f32> + std::fmt::Debug + PartialEq<T> + MemoryStorable,
|
||||
{
|
||||
let mut cpu = CPU::new(0);
|
||||
let mut cpu = common::empty_cpu(0);
|
||||
let buff: [F; _] = [
|
||||
1.0.into(),
|
||||
2.0.into(),
|
||||
@@ -51,10 +53,7 @@ where
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
vvadd,
|
||||
idata_build
|
||||
.set_rdr1r2(3, 1, 2)
|
||||
.set_imm_len(8 * size_of::<F>() as i32)
|
||||
.build(),
|
||||
idata_build.set_rdr1r2(3, 1, 2).set_imm_len(8).build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
let mut executable = Executable::new(cpu, core_instruction);
|
||||
@@ -65,7 +64,8 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(16 * size_of::<F>(), 8 * size_of::<T>()).unwrap()[0],
|
||||
.load::<T>(16 * size_of::<F>(), 8 * size_of::<T>())
|
||||
.unwrap()[0],
|
||||
vec![
|
||||
10.0.into(),
|
||||
12.0.into(),
|
||||
@@ -84,17 +84,22 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<F>(0, 16 * size_of::<F>()).unwrap()[0],
|
||||
.load::<F>(0, 16 * size_of::<F>())
|
||||
.unwrap()[0],
|
||||
&buff,
|
||||
"Altered first part for {}",
|
||||
err
|
||||
);
|
||||
//Check that later is 0
|
||||
assert_eq!(
|
||||
executable.cpu_mut().host().load::<i32>(
|
||||
16 * size_of::<F>() + 8 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
).unwrap()[0],
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<i32>(
|
||||
16 * size_of::<F>() + 8 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
)
|
||||
.unwrap()[0],
|
||||
[0, 0, 0, 0],
|
||||
"Altered first part for {}",
|
||||
err
|
||||
@@ -115,7 +120,7 @@ where
|
||||
F: From<f32> + std::fmt::Debug + PartialEq<F> + MemoryStorable,
|
||||
T: From<f32> + std::fmt::Debug + PartialEq<T> + MemoryStorable,
|
||||
{
|
||||
let mut cpu = CPU::new(0);
|
||||
let mut cpu = common::empty_cpu(0);
|
||||
let buff: [F; _] = [
|
||||
1.0.into(),
|
||||
2.0.into(),
|
||||
@@ -155,10 +160,7 @@ where
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
vvsub,
|
||||
idata_build
|
||||
.set_rdr1r2(3, 1, 2)
|
||||
.set_imm_len(8 * size_of::<F>() as i32)
|
||||
.build(),
|
||||
idata_build.set_rdr1r2(3, 1, 2).set_imm_len(8).build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
let mut executable = Executable::new(cpu, core_instruction);
|
||||
@@ -169,7 +171,8 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(16 * size_of::<F>(), 8 * size_of::<T>()).unwrap()[0],
|
||||
.load::<T>(16 * size_of::<F>(), 8 * size_of::<T>())
|
||||
.unwrap()[0],
|
||||
vec![
|
||||
(-8.0).into(),
|
||||
(-8.0).into(),
|
||||
@@ -188,17 +191,22 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<F>(0, 16 * size_of::<F>()).unwrap()[0],
|
||||
.load::<F>(0, 16 * size_of::<F>())
|
||||
.unwrap()[0],
|
||||
&buff,
|
||||
"Altered first part for {}",
|
||||
err
|
||||
);
|
||||
//Check that later is 0
|
||||
assert_eq!(
|
||||
executable.cpu_mut().host().load::<i32>(
|
||||
16 * size_of::<F>() + 8 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
).unwrap()[0],
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<i32>(
|
||||
16 * size_of::<F>() + 8 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
)
|
||||
.unwrap()[0],
|
||||
[0, 0, 0, 0],
|
||||
"Altered first part for {}",
|
||||
err
|
||||
@@ -219,7 +227,7 @@ where
|
||||
F: From<f32> + std::fmt::Debug + PartialEq<F> + MemoryStorable,
|
||||
T: From<f32> + std::fmt::Debug + PartialEq<T> + MemoryStorable,
|
||||
{
|
||||
let mut cpu = CPU::new(0);
|
||||
let mut cpu = common::empty_cpu(0);
|
||||
let buff: [F; _] = [
|
||||
1.0.into(),
|
||||
2.0.into(),
|
||||
@@ -259,10 +267,7 @@ where
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
vvmul,
|
||||
idata_build
|
||||
.set_rdr1r2(3, 1, 2)
|
||||
.set_imm_len(8 * size_of::<F>() as i32)
|
||||
.build(),
|
||||
idata_build.set_rdr1r2(3, 1, 2).set_imm_len(8).build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
let mut executable = Executable::new(cpu, core_instruction);
|
||||
@@ -273,7 +278,8 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(16 * size_of::<F>(), 8 * size_of::<T>()).unwrap()[0],
|
||||
.load::<T>(16 * size_of::<F>(), 8 * size_of::<T>())
|
||||
.unwrap()[0],
|
||||
vec![
|
||||
(9.0).into(),
|
||||
(20.0).into(),
|
||||
@@ -292,17 +298,22 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<F>(0, 16 * size_of::<F>()).unwrap()[0],
|
||||
.load::<F>(0, 16 * size_of::<F>())
|
||||
.unwrap()[0],
|
||||
&buff,
|
||||
"Altered first part for {}",
|
||||
err
|
||||
);
|
||||
//Check that later is 0
|
||||
assert_eq!(
|
||||
executable.cpu_mut().host().load::<i32>(
|
||||
16 * size_of::<F>() + 8 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
).unwrap()[0],
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<i32>(
|
||||
16 * size_of::<F>() + 8 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
)
|
||||
.unwrap()[0],
|
||||
[0, 0, 0, 0],
|
||||
"Altered first part for {}",
|
||||
err
|
||||
@@ -323,7 +334,7 @@ where
|
||||
F: From<f32> + std::fmt::Debug + PartialEq<F> + MemoryStorable,
|
||||
T: From<f32> + std::fmt::Debug + PartialEq<T> + MemoryStorable,
|
||||
{
|
||||
let mut cpu = CPU::new(0);
|
||||
let mut cpu = common::empty_cpu(0);
|
||||
let buff: [F; _] = [
|
||||
1.0.into(),
|
||||
2.0.into(),
|
||||
@@ -363,10 +374,7 @@ where
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
vvdmul,
|
||||
idata_build
|
||||
.set_rdr1r2(3, 1, 2)
|
||||
.set_imm_len(8 * size_of::<F>() as i32)
|
||||
.build(),
|
||||
idata_build.set_rdr1r2(3, 1, 2).set_imm_len(8).build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
let mut executable = Executable::new(cpu, core_instruction);
|
||||
@@ -377,10 +385,9 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(16 * size_of::<F>(), size_of::<T>()).unwrap()[0],
|
||||
vec![
|
||||
(492.0).into(),
|
||||
],
|
||||
.load::<T>(16 * size_of::<F>(), size_of::<T>())
|
||||
.unwrap()[0],
|
||||
vec![(492.0).into(),],
|
||||
"Wrong result for {}",
|
||||
err
|
||||
);
|
||||
@@ -389,17 +396,19 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<F>(0, 16 * size_of::<F>()).unwrap()[0],
|
||||
.load::<F>(0, 16 * size_of::<F>())
|
||||
.unwrap()[0],
|
||||
&buff,
|
||||
"Altered first part for {}",
|
||||
err
|
||||
);
|
||||
//Check that later is 0
|
||||
assert_eq!(
|
||||
executable.cpu_mut().host().load::<i32>(
|
||||
16 * size_of::<F>() + size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
).unwrap()[0],
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<i32>(16 * size_of::<F>() + size_of::<T>(), 4 * size_of::<i32>())
|
||||
.unwrap()[0],
|
||||
[0, 0, 0, 0],
|
||||
"Altered first part for {}",
|
||||
err
|
||||
@@ -420,7 +429,7 @@ where
|
||||
F: From<f32> + std::fmt::Debug + PartialEq<F> + MemoryStorable,
|
||||
T: From<f32> + std::fmt::Debug + PartialEq<T> + MemoryStorable,
|
||||
{
|
||||
let mut cpu = CPU::new(0);
|
||||
let mut cpu = common::empty_cpu(0);
|
||||
let buff: [F; _] = [
|
||||
9.0.into(),
|
||||
2.0.into(),
|
||||
@@ -460,10 +469,7 @@ where
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
vvmax,
|
||||
idata_build
|
||||
.set_rdr1r2(3, 1, 2)
|
||||
.set_imm_len(8 * size_of::<F>() as i32)
|
||||
.build(),
|
||||
idata_build.set_rdr1r2(3, 1, 2).set_imm_len(8).build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
let mut executable = Executable::new(cpu, core_instruction);
|
||||
@@ -474,16 +480,17 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(16 * size_of::<F>(), 8 * size_of::<T>()).unwrap()[0],
|
||||
.load::<T>(16 * size_of::<F>(), 8 * size_of::<T>())
|
||||
.unwrap()[0],
|
||||
vec![
|
||||
9.0.into(),
|
||||
10.0.into(),
|
||||
11.0.into(),
|
||||
12.0.into(),
|
||||
13.0.into(),
|
||||
14.0.into(),
|
||||
15.0.into(),
|
||||
16.0.into(),
|
||||
9.0.into(),
|
||||
10.0.into(),
|
||||
11.0.into(),
|
||||
12.0.into(),
|
||||
13.0.into(),
|
||||
14.0.into(),
|
||||
15.0.into(),
|
||||
16.0.into(),
|
||||
],
|
||||
"Wrong result for {}",
|
||||
err
|
||||
@@ -493,17 +500,22 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<F>(0, 16 * size_of::<F>()).unwrap()[0],
|
||||
.load::<F>(0, 16 * size_of::<F>())
|
||||
.unwrap()[0],
|
||||
&buff,
|
||||
"Altered first part for {}",
|
||||
err
|
||||
);
|
||||
//Check that later is 0
|
||||
assert_eq!(
|
||||
executable.cpu_mut().host().load::<i32>(
|
||||
16 * size_of::<F>() + 8 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
).unwrap()[0],
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<i32>(
|
||||
16 * size_of::<F>() + 8 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
)
|
||||
.unwrap()[0],
|
||||
[0, 0, 0, 0],
|
||||
"Altered first part for {}",
|
||||
err
|
||||
@@ -524,7 +536,7 @@ where
|
||||
F: From<f32> + std::fmt::Debug + PartialEq<F> + MemoryStorable,
|
||||
T: From<f32> + std::fmt::Debug + PartialEq<T> + MemoryStorable,
|
||||
{
|
||||
let mut cpu = CPU::new(0);
|
||||
let mut cpu = common::empty_cpu(0);
|
||||
let buff: [F; _] = [
|
||||
9.0.into(),
|
||||
2.0.into(),
|
||||
@@ -562,7 +574,8 @@ where
|
||||
vavg,
|
||||
idata_build
|
||||
.set_rdr1r2(3, 1, 1)
|
||||
.set_imm_len(8 * size_of::<F>() as i32)
|
||||
.set_offset_select(1)
|
||||
.set_imm_len(8)
|
||||
.build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
@@ -574,10 +587,9 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(16 * size_of::<F>(), size_of::<T>()).unwrap()[0],
|
||||
vec![
|
||||
7.5.into(),
|
||||
],
|
||||
.load::<T>(16 * size_of::<F>(), size_of::<T>())
|
||||
.unwrap()[0],
|
||||
vec![7.5.into(),],
|
||||
"Wrong result for {}",
|
||||
err
|
||||
);
|
||||
@@ -586,17 +598,19 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<F>(0, 16 * size_of::<F>()).unwrap()[0],
|
||||
.load::<F>(0, 16 * size_of::<F>())
|
||||
.unwrap()[0],
|
||||
&buff,
|
||||
"Altered first part for {}",
|
||||
err
|
||||
);
|
||||
//Check that later is 0
|
||||
assert_eq!(
|
||||
executable.cpu_mut().host().load::<i32>(
|
||||
16 * size_of::<F>() + size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
).unwrap()[0],
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<i32>(16 * size_of::<F>() + size_of::<T>(), 4 * size_of::<i32>())
|
||||
.unwrap()[0],
|
||||
[0, 0, 0, 0],
|
||||
"Altered first part for {}",
|
||||
err
|
||||
@@ -617,7 +631,7 @@ where
|
||||
F: From<f32> + std::fmt::Debug + PartialEq<F> + MemoryStorable,
|
||||
T: From<f32> + std::fmt::Debug + PartialEq<T> + MemoryStorable,
|
||||
{
|
||||
let mut cpu = CPU::new(0);
|
||||
let mut cpu = common::empty_cpu(0);
|
||||
let buff: [F; _] = [
|
||||
(-9.0).into(),
|
||||
2.0.into(),
|
||||
@@ -653,10 +667,7 @@ where
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
vrelu,
|
||||
idata_build
|
||||
.set_rdr1r2(3, 1, 1)
|
||||
.set_imm_len(8 * size_of::<F>() as i32)
|
||||
.build(),
|
||||
idata_build.set_rdr1r2(3, 1, 1).set_imm_len(8).build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
let mut executable = Executable::new(cpu, core_instruction);
|
||||
@@ -667,16 +678,17 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(16 * size_of::<F>(), 8*size_of::<T>()).unwrap()[0],
|
||||
.load::<T>(16 * size_of::<F>(), 8 * size_of::<T>())
|
||||
.unwrap()[0],
|
||||
vec![
|
||||
0.0.into(),
|
||||
2.0.into(),
|
||||
11.0.into(),
|
||||
0.0.into(),
|
||||
13.0.into(),
|
||||
0.0.into(),
|
||||
7.0.into(),
|
||||
0.0.into(),
|
||||
0.0.into(),
|
||||
2.0.into(),
|
||||
11.0.into(),
|
||||
0.0.into(),
|
||||
13.0.into(),
|
||||
0.0.into(),
|
||||
7.0.into(),
|
||||
0.0.into(),
|
||||
],
|
||||
"Wrong result for {}",
|
||||
err
|
||||
@@ -686,17 +698,22 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<F>(0, 16 * size_of::<F>()).unwrap()[0],
|
||||
.load::<F>(0, 16 * size_of::<F>())
|
||||
.unwrap()[0],
|
||||
&buff,
|
||||
"Altered first part for {}",
|
||||
err
|
||||
);
|
||||
//Check that later is 0
|
||||
assert_eq!(
|
||||
executable.cpu_mut().host().load::<i32>(
|
||||
16 * size_of::<F>() + 8*size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
).unwrap()[0],
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<i32>(
|
||||
16 * size_of::<F>() + 8 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
)
|
||||
.unwrap()[0],
|
||||
[0, 0, 0, 0],
|
||||
"Altered first part for {}",
|
||||
err
|
||||
@@ -717,7 +734,7 @@ where
|
||||
F: From<f32> + std::fmt::Debug + PartialEq<F> + MemoryStorable,
|
||||
T: From<f32> + std::fmt::Debug + PartialEq<T> + MemoryStorable + UpcastDestTraits<T>,
|
||||
{
|
||||
let mut cpu = CPU::new(0);
|
||||
let mut cpu = common::empty_cpu(0);
|
||||
let buff: [F; _] = [
|
||||
0.1.into(),
|
||||
0.2.into(),
|
||||
@@ -753,32 +770,32 @@ where
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
vtanh,
|
||||
idata_build
|
||||
.set_rdr1r2(3, 1, 1)
|
||||
.set_imm_len(8 * size_of::<F>() as i32)
|
||||
.build(),
|
||||
idata_build.set_rdr1r2(3, 1, 1).set_imm_len(8).build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
let mut executable = Executable::new(cpu, core_instruction);
|
||||
executable.execute();
|
||||
|
||||
// Check result correct
|
||||
|
||||
|
||||
assert!(
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(16 * size_of::<F>(), 8*size_of::<T>()).unwrap()[0].iter().zip(
|
||||
vec![
|
||||
T::from(0.1).tanh(),
|
||||
T::from(0.2).tanh(),
|
||||
T::from(0.3).tanh(),
|
||||
T::from(0.4).tanh(),
|
||||
T::from(0.5).tanh(),
|
||||
T::from(0.6).tanh(),
|
||||
T::from(0.7).tanh(),
|
||||
T::from(0.8).tanh(),
|
||||
]).all(|(&a,b) : (&T, T)| {a-b < 0.001.into()}),
|
||||
.load::<T>(16 * size_of::<F>(), 8 * size_of::<T>())
|
||||
.unwrap()[0]
|
||||
.iter()
|
||||
.zip(vec![
|
||||
T::from(0.1).tanh(),
|
||||
T::from(0.2).tanh(),
|
||||
T::from(0.3).tanh(),
|
||||
T::from(0.4).tanh(),
|
||||
T::from(0.5).tanh(),
|
||||
T::from(0.6).tanh(),
|
||||
T::from(0.7).tanh(),
|
||||
T::from(0.8).tanh(),
|
||||
])
|
||||
.all(|(&a, b): (&T, T)| { a - b < 0.001.into() }),
|
||||
"Wrong result for {}",
|
||||
err
|
||||
);
|
||||
@@ -787,17 +804,22 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<F>(0, 16 * size_of::<F>()).unwrap()[0],
|
||||
.load::<F>(0, 16 * size_of::<F>())
|
||||
.unwrap()[0],
|
||||
&buff,
|
||||
"Altered first part for {}",
|
||||
err
|
||||
);
|
||||
//Check that later is 0
|
||||
assert_eq!(
|
||||
executable.cpu_mut().host().load::<i32>(
|
||||
16 * size_of::<F>() + 8*size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
).unwrap()[0],
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<i32>(
|
||||
16 * size_of::<F>() + 8 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
)
|
||||
.unwrap()[0],
|
||||
[0, 0, 0, 0],
|
||||
"Altered first part for {}",
|
||||
err
|
||||
@@ -812,14 +834,13 @@ fn vtanh_test() {
|
||||
vtanh_test_generic::<f64, f64>("vtanh<f64,f64>");
|
||||
}
|
||||
|
||||
|
||||
/// vsigm Test
|
||||
fn vsigm_test_generic<F, T>(err: &str)
|
||||
where
|
||||
F: From<f32> + std::fmt::Debug + PartialEq<F> + MemoryStorable,
|
||||
T: From<f32> + std::fmt::Debug + PartialEq<T> + MemoryStorable + UpcastDestTraits<T>,
|
||||
{
|
||||
let mut cpu = CPU::new(0);
|
||||
let mut cpu = common::empty_cpu(0);
|
||||
let buff: [F; _] = [
|
||||
0.1.into(),
|
||||
0.2.into(),
|
||||
@@ -855,32 +876,32 @@ where
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
vsigm,
|
||||
idata_build
|
||||
.set_rdr1r2(3, 1, 1)
|
||||
.set_imm_len(8 * size_of::<F>() as i32)
|
||||
.build(),
|
||||
idata_build.set_rdr1r2(3, 1, 1).set_imm_len(8).build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
let mut executable = Executable::new(cpu, core_instruction);
|
||||
executable.execute();
|
||||
|
||||
// Check result correct
|
||||
|
||||
|
||||
assert!(
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(16 * size_of::<F>(), 8*size_of::<T>()).unwrap()[0].iter().zip(
|
||||
vec![
|
||||
T::from(0.1).sigm(),
|
||||
T::from(0.2).sigm(),
|
||||
T::from(0.3).sigm(),
|
||||
T::from(0.4).sigm(),
|
||||
T::from(0.5).sigm(),
|
||||
T::from(0.6).sigm(),
|
||||
T::from(0.7).sigm(),
|
||||
T::from(0.8).sigm(),
|
||||
]).all(|(&a,b) : (&T, T)| {a-b < 0.001.into()}),
|
||||
.load::<T>(16 * size_of::<F>(), 8 * size_of::<T>())
|
||||
.unwrap()[0]
|
||||
.iter()
|
||||
.zip(vec![
|
||||
T::from(0.1).sigm(),
|
||||
T::from(0.2).sigm(),
|
||||
T::from(0.3).sigm(),
|
||||
T::from(0.4).sigm(),
|
||||
T::from(0.5).sigm(),
|
||||
T::from(0.6).sigm(),
|
||||
T::from(0.7).sigm(),
|
||||
T::from(0.8).sigm(),
|
||||
])
|
||||
.all(|(&a, b): (&T, T)| { a - b < 0.001.into() }),
|
||||
"Wrong result for {}",
|
||||
err
|
||||
);
|
||||
@@ -889,17 +910,22 @@ where
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<F>(0, 16 * size_of::<F>()).unwrap()[0],
|
||||
.load::<F>(0, 16 * size_of::<F>())
|
||||
.unwrap()[0],
|
||||
&buff,
|
||||
"Altered first part for {}",
|
||||
err
|
||||
);
|
||||
//Check that later is 0
|
||||
assert_eq!(
|
||||
executable.cpu_mut().host().load::<i32>(
|
||||
16 * size_of::<F>() + 8*size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
).unwrap()[0],
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<i32>(
|
||||
16 * size_of::<F>() + 8 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
)
|
||||
.unwrap()[0],
|
||||
[0, 0, 0, 0],
|
||||
"Altered first part for {}",
|
||||
err
|
||||
@@ -914,18 +940,13 @@ fn vsigm_test() {
|
||||
vsigm_test_generic::<f64, f64>("vsigm<f64,f64>");
|
||||
}
|
||||
|
||||
|
||||
|
||||
/// mvmul Test
|
||||
fn mvmul_test_generic<F,M, T>(err: &str, relu:i32)
|
||||
fn mvmul_test_generic<F, M, T>(err: &str, relu: i32)
|
||||
where
|
||||
F: From<f32> + std::fmt::Debug + PartialEq<F> + MemoryStorable,
|
||||
M: From<f32> + std::fmt::Debug + PartialEq<M> + MemoryStorable,
|
||||
T: From<f32> + std::fmt::Debug + PartialEq<T> + MemoryStorable + UpcastDestTraits<T>,
|
||||
{
|
||||
let mut cpu = CPU::new(0);
|
||||
cpu.reserve_crossbar(1, 4 * size_of::<M>(), 4);
|
||||
let (memory, crossbars) = cpu.host().get_memory_crossbar();
|
||||
let matrix: [M; _] = [
|
||||
1.0.into(),
|
||||
2.0.into(),
|
||||
@@ -944,13 +965,11 @@ where
|
||||
15.0.into(),
|
||||
16.0.into(),
|
||||
];
|
||||
crossbars.get_mut(0).unwrap().execute_store( &matrix).unwrap();
|
||||
let vector: [F; _] = [
|
||||
1.0.into(),
|
||||
2.0.into(),
|
||||
3.0.into(),
|
||||
4.0.into(),
|
||||
];
|
||||
let mut crossbar = Crossbar::new(4 * size_of::<M>(), 4, CoreMemory::new());
|
||||
crossbar.execute_store(&matrix).unwrap();
|
||||
let mut cpu = pimcore::cpu::CPU::new(0, vec![vec![&crossbar]]);
|
||||
let (memory, _) = cpu.host().get_memory_crossbar();
|
||||
let vector: [F; _] = [1.0.into(), 2.0.into(), 3.0.into(), 4.0.into()];
|
||||
memory.execute_store(0, &vector).unwrap();
|
||||
|
||||
let mut inst_builder = InstructionsBuilder::new();
|
||||
@@ -971,7 +990,7 @@ where
|
||||
mvmul,
|
||||
idata_build
|
||||
.set_rdr1(3, 1)
|
||||
.set_mbiw_immrelu_immgroup(8*size_of::<M>() as i32, relu, 0)
|
||||
.set_mbiw_immrelu_immgroup(8 * size_of::<M>() as i32, relu, 0)
|
||||
.build(),
|
||||
);
|
||||
let core_instruction = vec![inst_builder.build().into()];
|
||||
@@ -979,54 +998,50 @@ where
|
||||
executable.execute();
|
||||
|
||||
// Check result correct
|
||||
if relu == 0 {
|
||||
assert_eq!(
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(4 * size_of::<F>(), 4*size_of::<T>()).unwrap()[0],
|
||||
vec![
|
||||
90.0.into(),
|
||||
(-24.0).into(),
|
||||
110.0.into(),
|
||||
120.0.into(),
|
||||
],
|
||||
"Wrong result for {}",
|
||||
err
|
||||
);
|
||||
}
|
||||
else {
|
||||
assert_eq!(
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(4 * size_of::<F>(), 4*size_of::<T>()).unwrap()[0],
|
||||
vec![
|
||||
90.0.into(),
|
||||
0.0.into(),
|
||||
110.0.into(),
|
||||
120.0.into(),
|
||||
],
|
||||
"Wrong result for {}",
|
||||
err
|
||||
);
|
||||
}
|
||||
if relu == 0 {
|
||||
assert_eq!(
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(4 * size_of::<F>(), 4 * size_of::<T>())
|
||||
.unwrap()[0],
|
||||
vec![90.0.into(), (-24.0).into(), 110.0.into(), 120.0.into(),],
|
||||
"Wrong result for {}",
|
||||
err
|
||||
);
|
||||
} else {
|
||||
assert_eq!(
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<T>(4 * size_of::<F>(), 4 * size_of::<T>())
|
||||
.unwrap()[0],
|
||||
vec![90.0.into(), 0.0.into(), 110.0.into(), 120.0.into(),],
|
||||
"Wrong result for {}",
|
||||
err
|
||||
);
|
||||
}
|
||||
// Check first part equal
|
||||
assert_eq!(
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<F>(0, 4 * size_of::<F>()).unwrap()[0],
|
||||
.load::<F>(0, 4 * size_of::<F>())
|
||||
.unwrap()[0],
|
||||
&vector,
|
||||
"Altered first part for {}",
|
||||
err
|
||||
);
|
||||
//Check that later is 0
|
||||
assert_eq!(
|
||||
executable.cpu_mut().host().load::<i32>(
|
||||
4 * size_of::<F>() + 4*size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
).unwrap()[0],
|
||||
executable
|
||||
.cpu_mut()
|
||||
.host()
|
||||
.load::<i32>(
|
||||
4 * size_of::<F>() + 4 * size_of::<T>(),
|
||||
4 * size_of::<i32>()
|
||||
)
|
||||
.unwrap()[0],
|
||||
[0, 0, 0, 0],
|
||||
"Altered first part for {}",
|
||||
err
|
||||
@@ -1035,24 +1050,21 @@ where
|
||||
|
||||
#[test]
|
||||
fn mvmul_test() {
|
||||
mvmul_test_generic::<f32,f32,f32>("mvmul<f32,f32,f32>",0);
|
||||
mvmul_test_generic::<f32,f32,f64>("mvmul<f32,f32,f64>",0);
|
||||
mvmul_test_generic::<f32,f64,f32>("mvmul<f32,f64,f32>",0);
|
||||
mvmul_test_generic::<f32,f64,f64>("mvmul<f32,f64,f64>",0);
|
||||
mvmul_test_generic::<f64,f32,f32>("mvmul<f64,f32,f32>",0);
|
||||
mvmul_test_generic::<f64,f32,f64>("mvmul<f64,f32,f64>",0);
|
||||
mvmul_test_generic::<f64,f64,f32>("mvmul<f64,f64,f32>",0);
|
||||
mvmul_test_generic::<f64,f64,f64>("mvmul<f64,f64,f64>",0);
|
||||
|
||||
mvmul_test_generic::<f32,f32,f32>("mvmul<f32,f32,f32>",1);
|
||||
mvmul_test_generic::<f32,f32,f64>("mvmul<f32,f32,f64>",1);
|
||||
mvmul_test_generic::<f32,f64,f32>("mvmul<f32,f64,f32>",1);
|
||||
mvmul_test_generic::<f32,f64,f64>("mvmul<f32,f64,f64>",1);
|
||||
mvmul_test_generic::<f64,f32,f32>("mvmul<f64,f32,f32>",1);
|
||||
mvmul_test_generic::<f64,f32,f64>("mvmul<f64,f32,f64>",1);
|
||||
mvmul_test_generic::<f64,f64,f32>("mvmul<f64,f64,f32>",1);
|
||||
mvmul_test_generic::<f64,f64,f64>("mvmul<f64,f64,f64>",1);
|
||||
mvmul_test_generic::<f32, f32, f32>("mvmul<f32,f32,f32>", 0);
|
||||
mvmul_test_generic::<f32, f32, f64>("mvmul<f32,f32,f64>", 0);
|
||||
mvmul_test_generic::<f32, f64, f32>("mvmul<f32,f64,f32>", 0);
|
||||
mvmul_test_generic::<f32, f64, f64>("mvmul<f32,f64,f64>", 0);
|
||||
mvmul_test_generic::<f64, f32, f32>("mvmul<f64,f32,f32>", 0);
|
||||
mvmul_test_generic::<f64, f32, f64>("mvmul<f64,f32,f64>", 0);
|
||||
mvmul_test_generic::<f64, f64, f32>("mvmul<f64,f64,f32>", 0);
|
||||
mvmul_test_generic::<f64, f64, f64>("mvmul<f64,f64,f64>", 0);
|
||||
|
||||
mvmul_test_generic::<f32, f32, f32>("mvmul<f32,f32,f32>", 1);
|
||||
mvmul_test_generic::<f32, f32, f64>("mvmul<f32,f32,f64>", 1);
|
||||
mvmul_test_generic::<f32, f64, f32>("mvmul<f32,f64,f32>", 1);
|
||||
mvmul_test_generic::<f32, f64, f64>("mvmul<f32,f64,f64>", 1);
|
||||
mvmul_test_generic::<f64, f32, f32>("mvmul<f64,f32,f32>", 1);
|
||||
mvmul_test_generic::<f64, f32, f64>("mvmul<f64,f32,f64>", 1);
|
||||
mvmul_test_generic::<f64, f64, f32>("mvmul<f64,f64,f32>", 1);
|
||||
mvmul_test_generic::<f64, f64, f64>("mvmul<f64,f64,f64>", 1);
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -1,12 +1,13 @@
|
||||
mod common;
|
||||
|
||||
use pimcore::{
|
||||
Executable, CoreInstructionsBuilder,
|
||||
cpu::CPU,
|
||||
CoreInstructionsBuilder, Executable,
|
||||
instruction_set::{InstructionsBuilder, instruction_data::InstructionDataBuilder, isa::*},
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn ld_test() {
|
||||
let mut cpu = CPU::new(1);
|
||||
let mut cpu = common::empty_cpu(1);
|
||||
let mut core_instruction_builder = CoreInstructionsBuilder::new(1);
|
||||
let buff: [f32; _] = [
|
||||
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0,
|
||||
@@ -41,7 +42,7 @@ fn ld_test() {
|
||||
|
||||
#[test]
|
||||
fn st_test() {
|
||||
let mut cpu = CPU::new(1);
|
||||
let mut cpu = common::empty_cpu(1);
|
||||
let mut core_instruction_builder = CoreInstructionsBuilder::new(1);
|
||||
let buff: [f32; _] = [
|
||||
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0,
|
||||
@@ -76,7 +77,7 @@ fn st_test() {
|
||||
|
||||
#[test]
|
||||
fn lldi_test() {
|
||||
let cpu = CPU::new(1);
|
||||
let cpu = common::empty_cpu(1);
|
||||
let mut core_instruction_builder = CoreInstructionsBuilder::new(1);
|
||||
let mut inst_builder = InstructionsBuilder::new();
|
||||
let mut idata_build = InstructionDataBuilder::new();
|
||||
@@ -106,7 +107,7 @@ fn lldi_test() {
|
||||
|
||||
#[test]
|
||||
fn lmv_test() {
|
||||
let mut cpu = CPU::new(1);
|
||||
let mut cpu = common::empty_cpu(1);
|
||||
let mut core_instruction_builder = CoreInstructionsBuilder::new(1);
|
||||
let buff: [f32; _] = [
|
||||
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0,
|
||||
@@ -148,7 +149,7 @@ fn lmv_test() {
|
||||
|
||||
#[test]
|
||||
fn simple_send_recv_test() {
|
||||
let mut cpu = CPU::new(2);
|
||||
let mut cpu = common::empty_cpu(2);
|
||||
let mut core_instruction_builder = CoreInstructionsBuilder::new(2);
|
||||
let buff: [f32; _] = [
|
||||
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0,
|
||||
@@ -157,7 +158,12 @@ fn simple_send_recv_test() {
|
||||
let mut inst_builder = InstructionsBuilder::new();
|
||||
let mut idata_build = InstructionDataBuilder::new();
|
||||
idata_build.set_core_indx(1).fix_core_indx();
|
||||
inst_builder.make_inst(sldi, idata_build.set_rdimm(1, 3*size_of::<f32>() as i32).build());
|
||||
inst_builder.make_inst(
|
||||
sldi,
|
||||
idata_build
|
||||
.set_rdimm(1, 3 * size_of::<f32>() as i32)
|
||||
.build(),
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
send,
|
||||
idata_build
|
||||
@@ -187,15 +193,11 @@ fn simple_send_recv_test() {
|
||||
|
||||
assert_eq!(
|
||||
res.unwrap()[0],
|
||||
[
|
||||
4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0
|
||||
],
|
||||
[4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0],
|
||||
"send_recv failed to store"
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
|
||||
// 1 -> 3
|
||||
// 2 -> 3
|
||||
// 3 <- 2
|
||||
@@ -207,76 +209,77 @@ fn simple_send_recv_test() {
|
||||
|
||||
#[test]
|
||||
fn multiple_send_recv_test() {
|
||||
let mut cpu = CPU::new(4);
|
||||
let mut cpu = common::empty_cpu(4);
|
||||
let mut core_instruction_builder = CoreInstructionsBuilder::new(4);
|
||||
let buff: [f32; _] = [
|
||||
1.0, 1.0, 1.0, 1.0, 1.0
|
||||
];
|
||||
let buff: [f32; _] = [1.0, 1.0, 1.0, 1.0, 1.0];
|
||||
cpu.core(1).execute_store(0, &buff).unwrap();
|
||||
let buff: [f32; _] = [
|
||||
2.0, 2.0, 2.0, 2.0, 2.0
|
||||
];
|
||||
let buff: [f32; _] = [2.0, 2.0, 2.0, 2.0, 2.0];
|
||||
cpu.core(2).execute_store(0, &buff).unwrap();
|
||||
let buff: [f32; _] = [
|
||||
3.0, 3.0, 3.0, 3.0, 3.0
|
||||
];
|
||||
let buff: [f32; _] = [3.0, 3.0, 3.0, 3.0, 3.0];
|
||||
cpu.core(3).execute_store(0, &buff).unwrap();
|
||||
let buff: [f32; _] = [
|
||||
4.0, 4.0, 4.0, 4.0, 4.0
|
||||
];
|
||||
let buff: [f32; _] = [4.0, 4.0, 4.0, 4.0, 4.0];
|
||||
cpu.core(4).execute_store(0, &buff).unwrap();
|
||||
|
||||
let send_inst = |cpu :&mut CPU, core_instruction_builder: &mut CoreInstructionsBuilder, inst_builder: &mut InstructionsBuilder, from : i32, to : i32| {
|
||||
let mut idata_build = InstructionDataBuilder::new();
|
||||
idata_build.set_core_indx(from).fix_core_indx();
|
||||
inst_builder.make_inst(sldi, idata_build.set_rdimm(1, from*size_of::<f32>() as i32).build());
|
||||
inst_builder.make_inst(
|
||||
send,
|
||||
idata_build
|
||||
.set_r1(1)
|
||||
.set_imm_core(to)
|
||||
.set_imm_len(size_of::<f32>() as i32)
|
||||
.build(),
|
||||
);
|
||||
let send_inst = |inst_builder: &mut InstructionsBuilder, from: i32, to: i32| {
|
||||
let mut idata_build = InstructionDataBuilder::new();
|
||||
idata_build.set_core_indx(from).fix_core_indx();
|
||||
inst_builder.make_inst(
|
||||
sldi,
|
||||
idata_build
|
||||
.set_rdimm(1, from * size_of::<f32>() as i32)
|
||||
.build(),
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
send,
|
||||
idata_build
|
||||
.set_r1(1)
|
||||
.set_imm_core(to)
|
||||
.set_imm_len(size_of::<f32>() as i32)
|
||||
.build(),
|
||||
);
|
||||
};
|
||||
|
||||
let recv_inst = |cpu :&mut CPU, core_instruction_builder: &mut CoreInstructionsBuilder, mut inst_builder: &mut InstructionsBuilder, to : i32, from : i32| {
|
||||
let mut idata_build = InstructionDataBuilder::new();
|
||||
idata_build.set_core_indx(to).fix_core_indx();
|
||||
inst_builder.make_inst(sldi, idata_build.set_rdimm(1, from*size_of::<f32>() as i32).build());
|
||||
inst_builder.make_inst(
|
||||
recv,
|
||||
idata_build
|
||||
.set_rd(1)
|
||||
.set_imm_core(from)
|
||||
.set_imm_len(size_of::<f32>() as i32)
|
||||
.build(),
|
||||
);
|
||||
let recv_inst = |inst_builder: &mut InstructionsBuilder, to: i32, from: i32| {
|
||||
let mut idata_build = InstructionDataBuilder::new();
|
||||
idata_build.set_core_indx(to).fix_core_indx();
|
||||
inst_builder.make_inst(
|
||||
sldi,
|
||||
idata_build
|
||||
.set_rdimm(1, from * size_of::<f32>() as i32)
|
||||
.build(),
|
||||
);
|
||||
inst_builder.make_inst(
|
||||
recv,
|
||||
idata_build
|
||||
.set_rd(1)
|
||||
.set_imm_core(from)
|
||||
.set_imm_len(size_of::<f32>() as i32)
|
||||
.build(),
|
||||
);
|
||||
};
|
||||
let mut inst_builder = InstructionsBuilder::new();
|
||||
|
||||
|
||||
// 1 -> 3
|
||||
send_inst(&mut cpu,&mut core_instruction_builder,&mut inst_builder,1, 3);
|
||||
send_inst(&mut inst_builder, 1, 3);
|
||||
core_instruction_builder.set_core(1, inst_builder.build());
|
||||
|
||||
// 2 -> 3
|
||||
// 2 <- 4
|
||||
send_inst(&mut cpu,&mut core_instruction_builder,&mut inst_builder,2, 3);
|
||||
recv_inst(&mut cpu,&mut core_instruction_builder,&mut inst_builder,2, 4);
|
||||
send_inst(&mut inst_builder, 2, 3);
|
||||
recv_inst(&mut inst_builder, 2, 4);
|
||||
core_instruction_builder.set_core(2, inst_builder.build());
|
||||
|
||||
// 3 <- 2
|
||||
// 3 <- 4
|
||||
// 3 <- 1
|
||||
recv_inst(&mut cpu,&mut core_instruction_builder,&mut inst_builder,3, 2);
|
||||
recv_inst(&mut cpu,&mut core_instruction_builder,&mut inst_builder,3, 4);
|
||||
recv_inst(&mut cpu,&mut core_instruction_builder,&mut inst_builder,3, 1);
|
||||
recv_inst(&mut inst_builder, 3, 2);
|
||||
recv_inst(&mut inst_builder, 3, 4);
|
||||
recv_inst(&mut inst_builder, 3, 1);
|
||||
core_instruction_builder.set_core(3, inst_builder.build());
|
||||
// 4 -> 2
|
||||
// 4 -> 3
|
||||
send_inst(&mut cpu,&mut core_instruction_builder,&mut inst_builder,4, 2);
|
||||
send_inst(&mut cpu,&mut core_instruction_builder,&mut inst_builder,4, 3);
|
||||
send_inst(&mut inst_builder, 4, 2);
|
||||
send_inst(&mut inst_builder, 4, 3);
|
||||
core_instruction_builder.set_core(4, inst_builder.build());
|
||||
|
||||
let mut executable = Executable::new(cpu, core_instruction_builder.build());
|
||||
@@ -288,7 +291,72 @@ fn multiple_send_recv_test() {
|
||||
|
||||
assert_eq!(
|
||||
res.unwrap()[0],
|
||||
[ 1.0, 2.0, 3.0, 4.0 ],
|
||||
[1.0, 2.0, 3.0, 4.0],
|
||||
"send_recv failed to store"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sync_wait_tokens_test() {
|
||||
let cpu = common::empty_cpu(2);
|
||||
let mut cores = CoreInstructionsBuilder::new(2);
|
||||
let mut instructions = InstructionsBuilder::new();
|
||||
let mut data = InstructionDataBuilder::new();
|
||||
|
||||
data.set_core_indx(1).fix_core_indx();
|
||||
for _ in 0..2 {
|
||||
instructions.make_inst(
|
||||
sync,
|
||||
data.set_imm_core(2).set_offset_select_value(0, 0).build(),
|
||||
);
|
||||
}
|
||||
cores.set_core(1, instructions.build());
|
||||
|
||||
data.set_core_indx(2).fix_core_indx();
|
||||
for _ in 0..2 {
|
||||
instructions.make_inst(wait, data.set_offset_select_value(0, 1).build());
|
||||
}
|
||||
cores.set_core(2, instructions.build());
|
||||
|
||||
Executable::new(cpu, cores.build()).execute().unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn blocked_transfers_do_not_starve_sync_producer() {
|
||||
let cpu = common::empty_cpu(4);
|
||||
let mut cores = CoreInstructionsBuilder::new(4);
|
||||
|
||||
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(recv, data.set_rd(1).set_imm_core(2).set_imm_len(1).build());
|
||||
instructions.make_inst(send, data.set_r1(1).set_imm_core(3).set_imm_len(1).build());
|
||||
cores.set_core(1, instructions.build());
|
||||
|
||||
let mut instructions = InstructionsBuilder::new();
|
||||
let mut data = InstructionDataBuilder::new();
|
||||
data.set_core_indx(2).fix_core_indx();
|
||||
instructions.make_inst(sldi, data.set_rdimm(1, 0).build());
|
||||
instructions.make_inst(wait, data.set_offset_select_value(0, 1).build());
|
||||
instructions.make_inst(send, data.set_r1(1).set_imm_core(1).set_imm_len(1).build());
|
||||
cores.set_core(2, instructions.build());
|
||||
|
||||
let mut instructions = InstructionsBuilder::new();
|
||||
let mut data = InstructionDataBuilder::new();
|
||||
data.set_core_indx(3).fix_core_indx();
|
||||
instructions.make_inst(sldi, data.set_rdimm(1, 0).build());
|
||||
instructions.make_inst(recv, data.set_rd(1).set_imm_core(1).set_imm_len(1).build());
|
||||
cores.set_core(3, instructions.build());
|
||||
|
||||
let mut instructions = InstructionsBuilder::new();
|
||||
let mut data = InstructionDataBuilder::new();
|
||||
data.set_core_indx(4).fix_core_indx();
|
||||
instructions.make_inst(
|
||||
sync,
|
||||
data.set_imm_core(2).set_offset_select_value(0, 0).build(),
|
||||
);
|
||||
cores.set_core(4, instructions.build());
|
||||
|
||||
Executable::new(cpu, cores.build()).execute().unwrap();
|
||||
}
|
||||
|
||||
Submodule backend-simulators/pim/pimsim-nn updated: 3e3442b663...6a3832525b
@@ -0,0 +1,5 @@
|
||||
colorama>=0.4.6,<1
|
||||
numpy>=1.26.4,<3
|
||||
onnx>=1.17,<2
|
||||
onnxsim>=0.6.5,<1
|
||||
-e ./tools/raptor_graph_explorer[dev]
|
||||
+57
-2
@@ -10,6 +10,56 @@ set(PIM_INCLUDE_PATH ${CMAKE_INCLUDE_OUTPUT_DIRECTORY})
|
||||
set(PIM_ONNX_MLIR_SRC_ROOT ${ONNX_MLIR_SRC_ROOT})
|
||||
set(PIM_ONNX_MLIR_BIN_ROOT ${ONNX_MLIR_BIN_ROOT})
|
||||
|
||||
set(PIM_GENERATED_PATH_SHIM_TARGET "")
|
||||
get_filename_component(PIM_BIN_ROOT_NAME "${PIM_BIN_ROOT}" NAME)
|
||||
if (PIM_BIN_ROOT_NAME STREQUAL "raptor-external")
|
||||
get_filename_component(PIM_GENERATED_PATH_SHIM_ROOT "${PIM_BIN_ROOT}" DIRECTORY)
|
||||
set(PIM_GENERATED_PATH_SHIM_OUTPUTS)
|
||||
|
||||
function(add_pim_generated_path_shim relative_path)
|
||||
set(real_file "${PIM_BIN_ROOT}/${relative_path}")
|
||||
set(shim_file "${PIM_GENERATED_PATH_SHIM_ROOT}/${relative_path}")
|
||||
get_filename_component(shim_dir "${shim_file}" DIRECTORY)
|
||||
|
||||
add_custom_command(
|
||||
OUTPUT "${shim_file}"
|
||||
DEPENDS "${real_file}"
|
||||
COMMAND "${CMAKE_COMMAND}" -E make_directory "${shim_dir}"
|
||||
COMMAND "${CMAKE_COMMAND}" -E rm -f "${shim_file}"
|
||||
COMMAND "${CMAKE_COMMAND}" -E create_symlink "${real_file}" "${shim_file}"
|
||||
VERBATIM
|
||||
)
|
||||
|
||||
list(APPEND PIM_GENERATED_PATH_SHIM_OUTPUTS "${shim_file}")
|
||||
set(PIM_GENERATED_PATH_SHIM_OUTPUTS "${PIM_GENERATED_PATH_SHIM_OUTPUTS}" PARENT_SCOPE)
|
||||
endfunction()
|
||||
|
||||
file(GLOB_RECURSE pim_generated_path_scan_sources
|
||||
CONFIGURE_DEPENDS
|
||||
"${PIM_SRC_ROOT}/*.cpp"
|
||||
"${PIM_SRC_ROOT}/*.hpp"
|
||||
)
|
||||
|
||||
set(pim_generated_path_shims)
|
||||
foreach (source_file IN LISTS pim_generated_path_scan_sources)
|
||||
file(READ "${source_file}" source_contents)
|
||||
string(REGEX MATCHALL "#include \"src/Accelerators/PIM/[^\"]+\\.inc\"" source_inc_matches "${source_contents}")
|
||||
|
||||
foreach (inc_match IN LISTS source_inc_matches)
|
||||
string(REGEX REPLACE "^#include \"src/Accelerators/PIM/(.+)\"$" "\\1" relative_inc_path "${inc_match}")
|
||||
list(APPEND pim_generated_path_shims "${relative_inc_path}")
|
||||
endforeach ()
|
||||
endforeach ()
|
||||
|
||||
list(REMOVE_DUPLICATES pim_generated_path_shims)
|
||||
foreach (relative_inc_path IN LISTS pim_generated_path_shims)
|
||||
add_pim_generated_path_shim("${relative_inc_path}")
|
||||
endforeach ()
|
||||
|
||||
add_custom_target(OMPimGeneratedPathShims DEPENDS ${PIM_GENERATED_PATH_SHIM_OUTPUTS})
|
||||
set(PIM_GENERATED_PATH_SHIM_TARGET OMPimGeneratedPathShims)
|
||||
endif ()
|
||||
|
||||
set(PIM_PUBLIC_INCLUDE_DIRS
|
||||
${ONNX_MLIR_SRC_ROOT}/include
|
||||
${ONNX_MLIR_SRC_ROOT}
|
||||
@@ -37,11 +87,14 @@ set(PIM_GENERATED_INCLUDE_DIRS
|
||||
|
||||
function(add_pim_library name)
|
||||
add_onnx_mlir_library(${name} STATIC ${ARGN})
|
||||
if (PIM_GENERATED_PATH_SHIM_TARGET)
|
||||
add_dependencies(${name} ${PIM_GENERATED_PATH_SHIM_TARGET})
|
||||
endif ()
|
||||
endfunction()
|
||||
|
||||
add_subdirectory(Dialect)
|
||||
add_subdirectory(Common)
|
||||
add_subdirectory(Pass)
|
||||
add_subdirectory(Passes)
|
||||
add_subdirectory(Compiler)
|
||||
add_subdirectory(Conversion)
|
||||
|
||||
@@ -64,9 +117,11 @@ add_pim_library(OMPIMAccel
|
||||
SpatialOps
|
||||
PimOps
|
||||
OMONNXToSpatial
|
||||
OMSpatialToGraphviz
|
||||
OMSpatialToPim
|
||||
OMPimCommon
|
||||
OMPimBufferization
|
||||
OMPimHostConstantFolding
|
||||
OMPimInstructionSelection
|
||||
OMPimVerification
|
||||
MLIRTensorInferTypeOpInterfaceImpl
|
||||
)
|
||||
|
||||
@@ -1,12 +1,24 @@
|
||||
add_pim_library(OMPimCommon
|
||||
IR/AffineUtils.cpp
|
||||
IR/AddressAnalysis.cpp
|
||||
IR/BatchCoreUtils.cpp
|
||||
IR/ConstantUtils.cpp
|
||||
IR/CoreBlockUtils.cpp
|
||||
IR/EntryPointUtils.cpp
|
||||
IR/IndexingUtils.cpp
|
||||
IR/LoopUtils.cpp
|
||||
IR/ShapeUtils.cpp
|
||||
IR/ShapingUtils.cpp
|
||||
IR/StaticIntSequence.cpp
|
||||
IR/StaticIntGrid.cpp
|
||||
IR/SubviewUtils.cpp
|
||||
IR/TensorSliceUtils.cpp
|
||||
IR/WeightUtils.cpp
|
||||
Support/CheckedArithmetic.cpp
|
||||
Support/DebugDump.cpp
|
||||
Support/Diagnostics.cpp
|
||||
Support/FileSystemUtils.cpp
|
||||
Support/ReportUtils.cpp
|
||||
|
||||
EXCLUDE_FROM_OM_LIBS
|
||||
|
||||
@@ -14,6 +26,8 @@ add_pim_library(OMPimCommon
|
||||
${PIM_PUBLIC_INCLUDE_DIRS}
|
||||
|
||||
LINK_LIBS PUBLIC
|
||||
MLIRLinalgDialect
|
||||
MLIRSCFDialect
|
||||
onnx
|
||||
SpatialOps
|
||||
PimOps
|
||||
|
||||
@@ -1,9 +1,19 @@
|
||||
#include "mlir/Dialect/Affine/IR/AffineOps.h"
|
||||
#include "mlir/Dialect/Arith/IR/Arith.h"
|
||||
#include "mlir/Dialect/Bufferization/IR/Bufferization.h"
|
||||
#include "mlir/Dialect/MemRef/IR/MemRef.h"
|
||||
#include "mlir/Dialect/SCF/IR/SCF.h"
|
||||
#include "mlir/IR/BuiltinAttributes.h"
|
||||
#include "mlir/Interfaces/DestinationStyleOpInterface.h"
|
||||
|
||||
#include "llvm/ADT/SmallPtrSet.h"
|
||||
|
||||
#include <limits>
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/AddressAnalysis.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/AffineUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
@@ -27,33 +37,378 @@ mlir::Value resolveAlias(mlir::Value value, const StaticValueKnowledge* knowledg
|
||||
return value;
|
||||
}
|
||||
|
||||
mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value, const StaticValueKnowledge* knowledge) {
|
||||
value = resolveAlias(value, knowledge);
|
||||
llvm::FailureOr<CompiledIndexExpr> compileIndexValueImpl(mlir::Value value);
|
||||
llvm::FailureOr<CompiledAddressExpr> compileContiguousAddressExprImpl(mlir::Value value);
|
||||
using AliasResolutionSet = llvm::SmallPtrSet<mlir::Value, 8>;
|
||||
mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value,
|
||||
const StaticValueKnowledge* knowledge,
|
||||
AliasResolutionSet& visited);
|
||||
mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value, const StaticValueKnowledge* knowledge);
|
||||
|
||||
if (mlir::isa<mlir::BlockArgument>(value))
|
||||
template <typename... Args>
|
||||
CompiledIndexExpr makeCompiledIndexExpr(Args&&... args) {
|
||||
return CompiledIndexExpr(std::make_shared<CompiledIndexExprNode>(std::forward<Args>(args)...));
|
||||
}
|
||||
|
||||
static mlir::Value resolveForYieldedAliasToInit(mlir::scf::ForOp forOp,
|
||||
mlir::Value yieldedValue,
|
||||
const StaticValueKnowledge* knowledge,
|
||||
AliasResolutionSet& visited) {
|
||||
yieldedValue = resolveLoopCarriedAliasImpl(yieldedValue, knowledge, visited);
|
||||
if (auto blockArgument = mlir::dyn_cast<mlir::BlockArgument>(yieldedValue)) {
|
||||
if (blockArgument.getOwner() == forOp.getBody() && blockArgument.getArgNumber() > 0
|
||||
&& static_cast<unsigned>(blockArgument.getArgNumber() - 1) < forOp.getInitArgs().size())
|
||||
return resolveLoopCarriedAliasImpl(forOp.getInitArgs()[blockArgument.getArgNumber() - 1], knowledge, visited);
|
||||
}
|
||||
return yieldedValue;
|
||||
}
|
||||
|
||||
mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value,
|
||||
const StaticValueKnowledge* knowledge,
|
||||
AliasResolutionSet& visited) {
|
||||
value = resolveAlias(value, knowledge);
|
||||
if (!value || !visited.insert(value).second)
|
||||
return value;
|
||||
|
||||
if (auto blockArgument = mlir::dyn_cast<mlir::BlockArgument>(value)) {
|
||||
auto forOp = mlir::dyn_cast_or_null<mlir::scf::ForOp>(blockArgument.getOwner()->getParentOp());
|
||||
if (forOp && blockArgument.getArgNumber() > 0) {
|
||||
const unsigned iterArgIndex = blockArgument.getArgNumber() - 1;
|
||||
auto yieldOp = mlir::dyn_cast<mlir::scf::YieldOp>(forOp.getBody()->getTerminator());
|
||||
if (iterArgIndex < forOp.getInitArgs().size() && yieldOp
|
||||
&& iterArgIndex < yieldOp.getNumOperands()) {
|
||||
mlir::Value yieldedValue = resolveAlias(yieldOp.getOperand(iterArgIndex), knowledge);
|
||||
if (yieldedValue == blockArgument
|
||||
|| (yieldedValue && resolveLoopCarriedAliasImpl(yieldedValue, knowledge, visited) == blockArgument))
|
||||
return resolveLoopCarriedAliasImpl(forOp.getInitArgs()[iterArgIndex], knowledge, visited);
|
||||
}
|
||||
}
|
||||
return value;
|
||||
}
|
||||
|
||||
mlir::Operation* definingOp = value.getDefiningOp();
|
||||
if (!definingOp)
|
||||
return value;
|
||||
|
||||
if (auto toBufferOp = mlir::dyn_cast<mlir::bufferization::ToBufferOp>(definingOp))
|
||||
return resolveLoopCarriedAliasImpl(toBufferOp.getTensor(), knowledge, visited);
|
||||
if (auto toTensorOp = mlir::dyn_cast<mlir::bufferization::ToTensorOp>(definingOp))
|
||||
return resolveLoopCarriedAliasImpl(toTensorOp.getBuffer(), knowledge, visited);
|
||||
|
||||
if (auto dpsDefiningOp = mlir::dyn_cast<mlir::DestinationStyleOpInterface>(definingOp)) {
|
||||
if (auto result = mlir::dyn_cast<mlir::OpResult>(value))
|
||||
if (mlir::OpOperand* tiedOperand = dpsDefiningOp.getTiedOpOperand(result))
|
||||
return resolveLoopCarriedAliasImpl(tiedOperand->get(), knowledge);
|
||||
return resolveLoopCarriedAliasImpl(tiedOperand->get(), knowledge, visited);
|
||||
}
|
||||
|
||||
if (auto forOp = mlir::dyn_cast<mlir::scf::ForOp>(definingOp)) {
|
||||
auto result = mlir::dyn_cast<mlir::OpResult>(value);
|
||||
if (result) {
|
||||
auto yieldOp = mlir::dyn_cast<mlir::scf::YieldOp>(forOp.getBody()->getTerminator());
|
||||
if (yieldOp && result.getResultNumber() < yieldOp.getNumOperands())
|
||||
return resolveForYieldedAliasToInit(
|
||||
forOp, yieldOp.getOperand(result.getResultNumber()), knowledge, visited);
|
||||
}
|
||||
}
|
||||
|
||||
if (auto castOp = mlir::dyn_cast<mlir::memref::CastOp>(definingOp))
|
||||
return resolveLoopCarriedAliasImpl(castOp.getSource(), knowledge);
|
||||
return resolveLoopCarriedAliasImpl(castOp.getSource(), knowledge, visited);
|
||||
if (auto collapseOp = mlir::dyn_cast<mlir::memref::CollapseShapeOp>(definingOp))
|
||||
return resolveLoopCarriedAliasImpl(collapseOp.getSrc(), knowledge);
|
||||
return resolveLoopCarriedAliasImpl(collapseOp.getSrc(), knowledge, visited);
|
||||
if (auto expandOp = mlir::dyn_cast<mlir::memref::ExpandShapeOp>(definingOp))
|
||||
return resolveLoopCarriedAliasImpl(expandOp.getSrc(), knowledge);
|
||||
return resolveLoopCarriedAliasImpl(expandOp.getSrc(), knowledge, visited);
|
||||
|
||||
return value;
|
||||
}
|
||||
|
||||
mlir::Value resolveLoopCarriedAliasImpl(mlir::Value value, const StaticValueKnowledge* knowledge) {
|
||||
AliasResolutionSet visited;
|
||||
return resolveLoopCarriedAliasImpl(value, knowledge, visited);
|
||||
}
|
||||
|
||||
llvm::FailureOr<int64_t> resolveOpFoldResult(mlir::OpFoldResult ofr, const StaticValueKnowledge* knowledge);
|
||||
llvm::FailureOr<int64_t> resolveIndexValueImpl(mlir::Value value, const StaticValueKnowledge* knowledge);
|
||||
|
||||
static llvm::FailureOr<llvm::SmallVector<int64_t>> getStaticMemRefStrides(mlir::MemRefType type) {
|
||||
llvm::SmallVector<int64_t> strides;
|
||||
int64_t offset = 0;
|
||||
if (failed(type.getStridesAndOffset(strides, offset)))
|
||||
return mlir::failure();
|
||||
if (llvm::any_of(strides, mlir::ShapedType::isDynamic))
|
||||
return mlir::failure();
|
||||
return strides;
|
||||
}
|
||||
|
||||
static llvm::FailureOr<int64_t> resolveConstantGlobalLoad(mlir::memref::LoadOp loadOp,
|
||||
const StaticValueKnowledge* knowledge) {
|
||||
auto getGlobalOp = loadOp.getMemRef().getDefiningOp<mlir::memref::GetGlobalOp>();
|
||||
if (!getGlobalOp)
|
||||
return mlir::failure();
|
||||
|
||||
auto moduleOp = loadOp->getParentOfType<mlir::ModuleOp>();
|
||||
auto globalOp = lookupGlobalForGetGlobal(moduleOp, getGlobalOp);
|
||||
if (!globalOp || !globalOp.getConstant() || !globalOp.getInitialValue())
|
||||
return mlir::failure();
|
||||
|
||||
auto denseAttr = mlir::dyn_cast<mlir::DenseElementsAttr>(*globalOp.getInitialValue());
|
||||
auto globalType = mlir::dyn_cast<mlir::MemRefType>(getGlobalOp.getType());
|
||||
if (!denseAttr || !globalType || !globalType.hasStaticShape())
|
||||
return mlir::failure();
|
||||
|
||||
auto elementType = denseAttr.getElementType();
|
||||
if (!elementType.isIndex() && !elementType.isInteger())
|
||||
return mlir::failure();
|
||||
|
||||
llvm::SmallVector<int64_t> indices;
|
||||
indices.reserve(loadOp.getIndices().size());
|
||||
for (mlir::Value index : loadOp.getIndices()) {
|
||||
auto resolvedIndex = resolveIndexValueImpl(index, knowledge);
|
||||
if (failed(resolvedIndex))
|
||||
return mlir::failure();
|
||||
indices.push_back(*resolvedIndex);
|
||||
}
|
||||
|
||||
if (indices.size() != static_cast<size_t>(globalType.getRank()))
|
||||
return mlir::failure();
|
||||
|
||||
auto strides = computeRowMajorStrides(globalType.getShape());
|
||||
int64_t linearIndex = linearizeIndex(indices, strides);
|
||||
if (linearIndex < 0 || linearIndex >= globalType.getNumElements())
|
||||
return mlir::failure();
|
||||
|
||||
return denseAttr.getValues<llvm::APInt>()[linearIndex].getSExtValue();
|
||||
}
|
||||
|
||||
static bool evaluateCmpPredicate(mlir::arith::CmpIPredicate predicate, int64_t lhs, int64_t rhs) {
|
||||
switch (predicate) {
|
||||
case mlir::arith::CmpIPredicate::eq: return lhs == rhs;
|
||||
case mlir::arith::CmpIPredicate::ne: return lhs != rhs;
|
||||
case mlir::arith::CmpIPredicate::slt: return lhs < rhs;
|
||||
case mlir::arith::CmpIPredicate::sle: return lhs <= rhs;
|
||||
case mlir::arith::CmpIPredicate::sgt: return lhs > rhs;
|
||||
case mlir::arith::CmpIPredicate::sge: return lhs >= rhs;
|
||||
case mlir::arith::CmpIPredicate::ult: return static_cast<uint64_t>(lhs) < static_cast<uint64_t>(rhs);
|
||||
case mlir::arith::CmpIPredicate::ule: return static_cast<uint64_t>(lhs) <= static_cast<uint64_t>(rhs);
|
||||
case mlir::arith::CmpIPredicate::ugt: return static_cast<uint64_t>(lhs) > static_cast<uint64_t>(rhs);
|
||||
case mlir::arith::CmpIPredicate::uge: return static_cast<uint64_t>(lhs) >= static_cast<uint64_t>(rhs);
|
||||
}
|
||||
|
||||
llvm_unreachable("unknown cmpi predicate");
|
||||
}
|
||||
|
||||
llvm::FailureOr<int64_t> evaluateCompiledIndexExpr(const CompiledIndexExpr& expr,
|
||||
const StaticValueKnowledge& knowledge) {
|
||||
if (!expr.node)
|
||||
return mlir::failure();
|
||||
|
||||
switch (expr.node->kind) {
|
||||
case CompiledIndexExprNode::Kind::Constant: return expr.node->constant;
|
||||
case CompiledIndexExprNode::Kind::Symbol: {
|
||||
auto value = resolveAlias(expr.node->symbol, &knowledge);
|
||||
auto iter = knowledge.indexValues.find(value);
|
||||
if (iter != knowledge.indexValues.end())
|
||||
return iter->second;
|
||||
return mlir::failure();
|
||||
}
|
||||
case CompiledIndexExprNode::Kind::Add:
|
||||
case CompiledIndexExprNode::Kind::Sub:
|
||||
case CompiledIndexExprNode::Kind::Mul:
|
||||
case CompiledIndexExprNode::Kind::DivUI:
|
||||
case CompiledIndexExprNode::Kind::DivSI:
|
||||
case CompiledIndexExprNode::Kind::RemUI:
|
||||
case CompiledIndexExprNode::Kind::RemSI:
|
||||
case CompiledIndexExprNode::Kind::MinUI:
|
||||
case CompiledIndexExprNode::Kind::CmpI: {
|
||||
auto lhs = evaluateCompiledIndexExpr(expr.node->operands[0], knowledge);
|
||||
auto rhs = evaluateCompiledIndexExpr(expr.node->operands[1], knowledge);
|
||||
if (failed(lhs) || failed(rhs))
|
||||
return mlir::failure();
|
||||
|
||||
switch (expr.node->kind) {
|
||||
case CompiledIndexExprNode::Kind::Add: return *lhs + *rhs;
|
||||
case CompiledIndexExprNode::Kind::Sub: return *lhs - *rhs;
|
||||
case CompiledIndexExprNode::Kind::Mul: return *lhs * *rhs;
|
||||
case CompiledIndexExprNode::Kind::DivUI:
|
||||
if (*rhs == 0)
|
||||
return mlir::failure();
|
||||
return static_cast<int64_t>(static_cast<uint64_t>(*lhs) / static_cast<uint64_t>(*rhs));
|
||||
case CompiledIndexExprNode::Kind::DivSI:
|
||||
if (*rhs == 0 || (*lhs == std::numeric_limits<int64_t>::min() && *rhs == -1))
|
||||
return mlir::failure();
|
||||
return *lhs / *rhs;
|
||||
case CompiledIndexExprNode::Kind::RemUI:
|
||||
if (*rhs == 0)
|
||||
return mlir::failure();
|
||||
return static_cast<int64_t>(static_cast<uint64_t>(*lhs) % static_cast<uint64_t>(*rhs));
|
||||
case CompiledIndexExprNode::Kind::RemSI:
|
||||
if (*rhs == 0)
|
||||
return mlir::failure();
|
||||
if (*lhs == std::numeric_limits<int64_t>::min() && *rhs == -1)
|
||||
return 0;
|
||||
return *lhs % *rhs;
|
||||
case CompiledIndexExprNode::Kind::MinUI:
|
||||
return static_cast<int64_t>(std::min(static_cast<uint64_t>(*lhs), static_cast<uint64_t>(*rhs)));
|
||||
case CompiledIndexExprNode::Kind::CmpI: return evaluateCmpPredicate(expr.node->predicate, *lhs, *rhs) ? 1 : 0;
|
||||
default: llvm_unreachable("unexpected binary compiled index kind");
|
||||
}
|
||||
}
|
||||
case CompiledIndexExprNode::Kind::Select: {
|
||||
auto condition = evaluateCompiledIndexExpr(expr.node->operands[0], knowledge);
|
||||
if (failed(condition))
|
||||
return mlir::failure();
|
||||
return evaluateCompiledIndexExpr(*condition != 0 ? expr.node->operands[1] : expr.node->operands[2], knowledge);
|
||||
}
|
||||
case CompiledIndexExprNode::Kind::ConstantGlobalLoad: {
|
||||
if (!expr.node->globalOp || !expr.node->globalOp.getInitialValue())
|
||||
return mlir::failure();
|
||||
|
||||
auto denseAttr = mlir::dyn_cast<mlir::DenseElementsAttr>(*expr.node->globalOp.getInitialValue());
|
||||
auto globalType = mlir::dyn_cast<mlir::MemRefType>(expr.node->globalOp.getType());
|
||||
if (!denseAttr || !globalType)
|
||||
return mlir::failure();
|
||||
|
||||
llvm::SmallVector<int64_t> indices;
|
||||
indices.reserve(expr.node->operands.size());
|
||||
for (const CompiledIndexExpr& operand : expr.node->operands) {
|
||||
auto resolvedIndex = evaluateCompiledIndexExpr(operand, knowledge);
|
||||
if (failed(resolvedIndex))
|
||||
return mlir::failure();
|
||||
indices.push_back(*resolvedIndex);
|
||||
}
|
||||
|
||||
int64_t linearIndex = linearizeIndex(indices, expr.node->globalStrides);
|
||||
if (linearIndex < 0 || linearIndex >= globalType.getNumElements())
|
||||
return mlir::failure();
|
||||
return denseAttr.getValues<llvm::APInt>()[linearIndex].getSExtValue();
|
||||
}
|
||||
}
|
||||
|
||||
llvm_unreachable("unknown compiled index kind");
|
||||
}
|
||||
|
||||
llvm::FailureOr<CompiledIndexExpr> compileConstantGlobalLoad(mlir::memref::LoadOp loadOp) {
|
||||
auto getGlobalOp = loadOp.getMemRef().getDefiningOp<mlir::memref::GetGlobalOp>();
|
||||
if (!getGlobalOp)
|
||||
return mlir::failure();
|
||||
|
||||
auto moduleOp = loadOp->getParentOfType<mlir::ModuleOp>();
|
||||
auto globalOp = lookupGlobalForGetGlobal(moduleOp, getGlobalOp);
|
||||
if (!globalOp || !globalOp.getConstant() || !globalOp.getInitialValue())
|
||||
return mlir::failure();
|
||||
|
||||
auto denseAttr = mlir::dyn_cast<mlir::DenseElementsAttr>(*globalOp.getInitialValue());
|
||||
auto globalType = mlir::dyn_cast<mlir::MemRefType>(getGlobalOp.getType());
|
||||
if (!denseAttr || !globalType || !globalType.hasStaticShape())
|
||||
return mlir::failure();
|
||||
|
||||
auto elementType = denseAttr.getElementType();
|
||||
if (!elementType.isIndex() && !elementType.isInteger())
|
||||
return mlir::failure();
|
||||
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::ConstantGlobalLoad;
|
||||
expr.globalOp = globalOp;
|
||||
expr.globalStrides = computeRowMajorStrides(globalType.getShape());
|
||||
expr.operands.reserve(loadOp.getIndices().size());
|
||||
for (mlir::Value index : loadOp.getIndices()) {
|
||||
auto compiledIndex = compileIndexValueImpl(index);
|
||||
if (failed(compiledIndex))
|
||||
return mlir::failure();
|
||||
expr.operands.push_back(*compiledIndex);
|
||||
}
|
||||
return makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
|
||||
llvm::FailureOr<CompiledIndexExpr> compileIndexValueImpl(mlir::Value value) {
|
||||
if (auto constantOp = value.getDefiningOp<mlir::arith::ConstantOp>()) {
|
||||
if (auto integerAttr = mlir::dyn_cast<mlir::IntegerAttr>(constantOp.getValue())) {
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Constant;
|
||||
expr.constant = integerAttr.getInt();
|
||||
return makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
}
|
||||
|
||||
mlir::Operation* definingOp = value.getDefiningOp();
|
||||
if (!definingOp) {
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Symbol;
|
||||
expr.symbol = value;
|
||||
return makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
|
||||
auto buildBinaryExpr = [&](CompiledIndexExprNode::Kind kind, mlir::Value lhsValue, mlir::Value rhsValue) {
|
||||
auto lhs = compileIndexValueImpl(lhsValue);
|
||||
auto rhs = compileIndexValueImpl(rhsValue);
|
||||
if (failed(lhs) || failed(rhs))
|
||||
return llvm::FailureOr<CompiledIndexExpr>(mlir::failure());
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = kind;
|
||||
expr.operands = {*lhs, *rhs};
|
||||
return llvm::FailureOr<CompiledIndexExpr>(makeCompiledIndexExpr(std::move(expr)));
|
||||
};
|
||||
|
||||
if (auto indexCastOp = mlir::dyn_cast<mlir::arith::IndexCastOp>(definingOp))
|
||||
return compileIndexValueImpl(indexCastOp.getIn());
|
||||
if (auto addOp = mlir::dyn_cast<mlir::arith::AddIOp>(definingOp))
|
||||
return buildBinaryExpr(CompiledIndexExprNode::Kind::Add, addOp.getLhs(), addOp.getRhs());
|
||||
if (auto subOp = mlir::dyn_cast<mlir::arith::SubIOp>(definingOp))
|
||||
return buildBinaryExpr(CompiledIndexExprNode::Kind::Sub, subOp.getLhs(), subOp.getRhs());
|
||||
if (auto mulOp = mlir::dyn_cast<mlir::arith::MulIOp>(definingOp))
|
||||
return buildBinaryExpr(CompiledIndexExprNode::Kind::Mul, mulOp.getLhs(), mulOp.getRhs());
|
||||
if (auto divOp = mlir::dyn_cast<mlir::arith::DivUIOp>(definingOp))
|
||||
return buildBinaryExpr(CompiledIndexExprNode::Kind::DivUI, divOp.getLhs(), divOp.getRhs());
|
||||
if (auto divOp = mlir::dyn_cast<mlir::arith::DivSIOp>(definingOp))
|
||||
return buildBinaryExpr(CompiledIndexExprNode::Kind::DivSI, divOp.getLhs(), divOp.getRhs());
|
||||
if (auto minOp = mlir::dyn_cast<mlir::arith::MinUIOp>(definingOp))
|
||||
return buildBinaryExpr(CompiledIndexExprNode::Kind::MinUI, minOp.getLhs(), minOp.getRhs());
|
||||
if (auto remOp = mlir::dyn_cast<mlir::arith::RemUIOp>(definingOp))
|
||||
return buildBinaryExpr(CompiledIndexExprNode::Kind::RemUI, remOp.getLhs(), remOp.getRhs());
|
||||
if (auto remOp = mlir::dyn_cast<mlir::arith::RemSIOp>(definingOp))
|
||||
return buildBinaryExpr(CompiledIndexExprNode::Kind::RemSI, remOp.getLhs(), remOp.getRhs());
|
||||
if (auto cmpOp = mlir::dyn_cast<mlir::arith::CmpIOp>(definingOp)) {
|
||||
auto expr = buildBinaryExpr(CompiledIndexExprNode::Kind::CmpI, cmpOp.getLhs(), cmpOp.getRhs());
|
||||
if (failed(expr))
|
||||
return mlir::failure();
|
||||
auto exprNode = std::make_shared<CompiledIndexExprNode>(*expr->node);
|
||||
exprNode->predicate = cmpOp.getPredicate();
|
||||
return CompiledIndexExpr(exprNode);
|
||||
}
|
||||
if (auto maxOp = mlir::dyn_cast<mlir::arith::MaxUIOp>(definingOp)) {
|
||||
auto lhs = compileIndexValueImpl(maxOp.getLhs());
|
||||
auto rhs = compileIndexValueImpl(maxOp.getRhs());
|
||||
if (failed(lhs) || failed(rhs))
|
||||
return mlir::failure();
|
||||
|
||||
CompiledIndexExprNode cmpExpr;
|
||||
cmpExpr.kind = CompiledIndexExprNode::Kind::CmpI;
|
||||
cmpExpr.predicate = mlir::arith::CmpIPredicate::uge;
|
||||
cmpExpr.operands = {*lhs, *rhs};
|
||||
|
||||
CompiledIndexExprNode selectExpr;
|
||||
selectExpr.kind = CompiledIndexExprNode::Kind::Select;
|
||||
selectExpr.operands = {makeCompiledIndexExpr(std::move(cmpExpr)), *lhs, *rhs};
|
||||
return makeCompiledIndexExpr(std::move(selectExpr));
|
||||
}
|
||||
if (auto selectOp = mlir::dyn_cast<mlir::arith::SelectOp>(definingOp)) {
|
||||
auto condition = compileIndexValueImpl(selectOp.getCondition());
|
||||
auto trueValue = compileIndexValueImpl(selectOp.getTrueValue());
|
||||
auto falseValue = compileIndexValueImpl(selectOp.getFalseValue());
|
||||
if (failed(condition) || failed(trueValue) || failed(falseValue))
|
||||
return mlir::failure();
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Select;
|
||||
expr.operands = {*condition, *trueValue, *falseValue};
|
||||
return makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
if (auto loadOp = mlir::dyn_cast<mlir::memref::LoadOp>(definingOp))
|
||||
return compileConstantGlobalLoad(loadOp);
|
||||
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Symbol;
|
||||
expr.symbol = value;
|
||||
return makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
|
||||
llvm::FailureOr<int64_t> resolveIndexValueImpl(mlir::Value value, const StaticValueKnowledge* knowledge) {
|
||||
value = resolveAlias(value, knowledge);
|
||||
@@ -74,6 +429,11 @@ llvm::FailureOr<int64_t> resolveIndexValueImpl(mlir::Value value, const StaticVa
|
||||
if (!definingOp)
|
||||
return mlir::failure();
|
||||
|
||||
if (auto affineApplyOp = mlir::dyn_cast<mlir::affine::AffineApplyOp>(definingOp))
|
||||
return evaluateAffineApply(affineApplyOp, [&](mlir::Value operand) {
|
||||
return resolveIndexValueImpl(operand, knowledge);
|
||||
});
|
||||
|
||||
if (auto indexCastOp = mlir::dyn_cast<mlir::arith::IndexCastOp>(definingOp))
|
||||
return resolveIndexValueImpl(indexCastOp.getIn(), knowledge);
|
||||
|
||||
@@ -109,6 +469,24 @@ llvm::FailureOr<int64_t> resolveIndexValueImpl(mlir::Value value, const StaticVa
|
||||
return static_cast<int64_t>(static_cast<uint64_t>(*lhs) / static_cast<uint64_t>(*rhs));
|
||||
}
|
||||
|
||||
if (auto divOp = mlir::dyn_cast<mlir::arith::DivSIOp>(definingOp)) {
|
||||
auto lhs = resolveIndexValueImpl(divOp.getLhs(), knowledge);
|
||||
auto rhs = resolveIndexValueImpl(divOp.getRhs(), knowledge);
|
||||
if (failed(lhs) || failed(rhs) || *rhs == 0)
|
||||
return mlir::failure();
|
||||
if (*lhs == std::numeric_limits<int64_t>::min() && *rhs == -1)
|
||||
return mlir::failure();
|
||||
return *lhs / *rhs;
|
||||
}
|
||||
|
||||
if (auto minOp = mlir::dyn_cast<mlir::arith::MinUIOp>(definingOp)) {
|
||||
auto lhs = resolveIndexValueImpl(minOp.getLhs(), knowledge);
|
||||
auto rhs = resolveIndexValueImpl(minOp.getRhs(), knowledge);
|
||||
if (failed(lhs) || failed(rhs))
|
||||
return mlir::failure();
|
||||
return static_cast<int64_t>(std::min(static_cast<uint64_t>(*lhs), static_cast<uint64_t>(*rhs)));
|
||||
}
|
||||
|
||||
if (auto remOp = mlir::dyn_cast<mlir::arith::RemUIOp>(definingOp)) {
|
||||
auto lhs = resolveIndexValueImpl(remOp.getLhs(), knowledge);
|
||||
auto rhs = resolveIndexValueImpl(remOp.getRhs(), knowledge);
|
||||
@@ -117,6 +495,34 @@ llvm::FailureOr<int64_t> resolveIndexValueImpl(mlir::Value value, const StaticVa
|
||||
return static_cast<int64_t>(static_cast<uint64_t>(*lhs) % static_cast<uint64_t>(*rhs));
|
||||
}
|
||||
|
||||
if (auto remOp = mlir::dyn_cast<mlir::arith::RemSIOp>(definingOp)) {
|
||||
auto lhs = resolveIndexValueImpl(remOp.getLhs(), knowledge);
|
||||
auto rhs = resolveIndexValueImpl(remOp.getRhs(), knowledge);
|
||||
if (failed(lhs) || failed(rhs) || *rhs == 0)
|
||||
return mlir::failure();
|
||||
if (*lhs == std::numeric_limits<int64_t>::min() && *rhs == -1)
|
||||
return 0;
|
||||
return *lhs % *rhs;
|
||||
}
|
||||
|
||||
if (auto cmpOp = mlir::dyn_cast<mlir::arith::CmpIOp>(definingOp)) {
|
||||
auto lhs = resolveIndexValueImpl(cmpOp.getLhs(), knowledge);
|
||||
auto rhs = resolveIndexValueImpl(cmpOp.getRhs(), knowledge);
|
||||
if (failed(lhs) || failed(rhs))
|
||||
return mlir::failure();
|
||||
return evaluateCmpPredicate(cmpOp.getPredicate(), *lhs, *rhs) ? 1 : 0;
|
||||
}
|
||||
|
||||
if (auto selectOp = mlir::dyn_cast<mlir::arith::SelectOp>(definingOp)) {
|
||||
auto condition = resolveIndexValueImpl(selectOp.getCondition(), knowledge);
|
||||
if (failed(condition))
|
||||
return mlir::failure();
|
||||
return resolveIndexValueImpl(*condition != 0 ? selectOp.getTrueValue() : selectOp.getFalseValue(), knowledge);
|
||||
}
|
||||
|
||||
if (auto loadOp = mlir::dyn_cast<mlir::memref::LoadOp>(definingOp))
|
||||
return resolveConstantGlobalLoad(loadOp, knowledge);
|
||||
|
||||
return mlir::failure();
|
||||
}
|
||||
|
||||
@@ -144,6 +550,15 @@ llvm::FailureOr<ResolvedContiguousAddress> resolveContiguousAddressImpl(mlir::Va
|
||||
if (!definingOp)
|
||||
return mlir::failure();
|
||||
|
||||
if (auto toBufferOp = mlir::dyn_cast<mlir::bufferization::ToBufferOp>(definingOp)) {
|
||||
value = resolveAlias(toBufferOp.getTensor(), knowledge);
|
||||
continue;
|
||||
}
|
||||
if (auto toTensorOp = mlir::dyn_cast<mlir::bufferization::ToTensorOp>(definingOp)) {
|
||||
value = resolveAlias(toTensorOp.getBuffer(), knowledge);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto dpsDefiningOp = mlir::dyn_cast<mlir::DestinationStyleOpInterface>(definingOp)) {
|
||||
mlir::OpOperand* tiedOperand = dpsDefiningOp.getTiedOpOperand(mlir::dyn_cast<mlir::OpResult>(value));
|
||||
if (!tiedOperand)
|
||||
@@ -158,16 +573,27 @@ llvm::FailureOr<ResolvedContiguousAddress> resolveContiguousAddressImpl(mlir::Va
|
||||
return mlir::failure();
|
||||
|
||||
auto yieldOp = mlir::cast<mlir::scf::YieldOp>(forOp.getBody()->getTerminator());
|
||||
mlir::Value yieldedValue = resolveLoopCarriedAliasImpl(yieldOp.getOperand(result.getResultNumber()), knowledge);
|
||||
if (auto blockArgument = mlir::dyn_cast<mlir::BlockArgument>(yieldedValue)) {
|
||||
if (blockArgument.getOwner() == forOp.getBody() && blockArgument.getArgNumber() > 0
|
||||
&& static_cast<unsigned>(blockArgument.getArgNumber() - 1) < forOp.getInitArgs().size()) {
|
||||
value = resolveAlias(forOp.getInitArgs()[blockArgument.getArgNumber() - 1], knowledge);
|
||||
continue;
|
||||
}
|
||||
}
|
||||
AliasResolutionSet visited;
|
||||
value = resolveForYieldedAliasToInit(
|
||||
forOp, yieldOp.getOperand(result.getResultNumber()), knowledge, visited);
|
||||
continue;
|
||||
}
|
||||
|
||||
value = yieldedValue;
|
||||
if (auto ifOp = mlir::dyn_cast<mlir::scf::IfOp>(definingOp)) {
|
||||
auto result = mlir::dyn_cast<mlir::OpResult>(value);
|
||||
if (!result)
|
||||
return mlir::failure();
|
||||
|
||||
auto condition = resolveIndexValueImpl(ifOp.getCondition(), knowledge);
|
||||
if (failed(condition))
|
||||
return mlir::failure();
|
||||
|
||||
mlir::Region& selectedRegion = *condition != 0 ? ifOp.getThenRegion() : ifOp.getElseRegion();
|
||||
auto yieldOp = mlir::dyn_cast<mlir::scf::YieldOp>(selectedRegion.front().getTerminator());
|
||||
if (!yieldOp || result.getResultNumber() >= yieldOp.getNumOperands())
|
||||
return mlir::failure();
|
||||
|
||||
value = resolveLoopCarriedAliasImpl(yieldOp.getOperand(result.getResultNumber()), knowledge);
|
||||
continue;
|
||||
}
|
||||
|
||||
@@ -208,8 +634,10 @@ llvm::FailureOr<ResolvedContiguousAddress> resolveContiguousAddressImpl(mlir::Va
|
||||
if (!isMemoryContiguous(sourceType.getShape(), offsets, sizes, strides))
|
||||
return mlir::failure();
|
||||
|
||||
auto sourceStrides = computeRowMajorStrides(sourceType.getShape());
|
||||
byteOffset += linearizeIndex(offsets, sourceStrides) * subviewType.getElementTypeBitWidth() / 8;
|
||||
auto sourceStrides = getStaticMemRefStrides(sourceType);
|
||||
if (failed(sourceStrides))
|
||||
return mlir::failure();
|
||||
byteOffset += linearizeIndex(offsets, *sourceStrides) * getElementTypeSizeInBytes(subviewType.getElementType());
|
||||
value = resolveAlias(subviewOp.getSource(), knowledge);
|
||||
continue;
|
||||
}
|
||||
@@ -234,17 +662,235 @@ llvm::FailureOr<ResolvedContiguousAddress> resolveContiguousAddressImpl(mlir::Va
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
llvm::FailureOr<CompiledAddressExpr> compileContiguousAddressExprImpl(mlir::Value value) {
|
||||
int64_t constantByteOffset = 0;
|
||||
CompiledIndexExpr byteOffsetExpr;
|
||||
{
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Constant;
|
||||
expr.constant = 0;
|
||||
byteOffsetExpr = makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
|
||||
llvm::FailureOr<int64_t> resolveIndexValue(mlir::Value value) { return resolveIndexValueImpl(value, nullptr); }
|
||||
while (true) {
|
||||
if (mlir::isa<mlir::BlockArgument>(value))
|
||||
return CompiledAddressExpr {value, byteOffsetExpr};
|
||||
|
||||
mlir::Operation* definingOp = value.getDefiningOp();
|
||||
if (!definingOp)
|
||||
return mlir::failure();
|
||||
|
||||
if (auto toBufferOp = mlir::dyn_cast<mlir::bufferization::ToBufferOp>(definingOp)) {
|
||||
value = toBufferOp.getTensor();
|
||||
continue;
|
||||
}
|
||||
if (auto toTensorOp = mlir::dyn_cast<mlir::bufferization::ToTensorOp>(definingOp)) {
|
||||
value = toTensorOp.getBuffer();
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto dpsDefiningOp = mlir::dyn_cast<mlir::DestinationStyleOpInterface>(definingOp)) {
|
||||
mlir::OpOperand* tiedOperand = dpsDefiningOp.getTiedOpOperand(mlir::dyn_cast<mlir::OpResult>(value));
|
||||
if (!tiedOperand)
|
||||
return mlir::failure();
|
||||
value = tiedOperand->get();
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto forOp = mlir::dyn_cast<mlir::scf::ForOp>(definingOp)) {
|
||||
auto result = mlir::dyn_cast<mlir::OpResult>(value);
|
||||
if (!result)
|
||||
return mlir::failure();
|
||||
|
||||
auto yieldOp = mlir::cast<mlir::scf::YieldOp>(forOp.getBody()->getTerminator());
|
||||
AliasResolutionSet visited;
|
||||
value = resolveForYieldedAliasToInit(
|
||||
forOp, yieldOp.getOperand(result.getResultNumber()), nullptr, visited);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto ifOp = mlir::dyn_cast<mlir::scf::IfOp>(definingOp)) {
|
||||
auto result = mlir::dyn_cast<mlir::OpResult>(value);
|
||||
if (!result)
|
||||
return mlir::failure();
|
||||
|
||||
auto thenYield = mlir::dyn_cast<mlir::scf::YieldOp>(ifOp.getThenRegion().front().getTerminator());
|
||||
auto elseYield = mlir::dyn_cast<mlir::scf::YieldOp>(ifOp.getElseRegion().front().getTerminator());
|
||||
if (!thenYield || !elseYield || result.getResultNumber() >= thenYield.getNumOperands()
|
||||
|| result.getResultNumber() >= elseYield.getNumOperands()) {
|
||||
return mlir::failure();
|
||||
}
|
||||
|
||||
auto thenAddress = compileContiguousAddressExprImpl(thenYield.getOperand(result.getResultNumber()));
|
||||
auto elseAddress = compileContiguousAddressExprImpl(elseYield.getOperand(result.getResultNumber()));
|
||||
if (failed(thenAddress) || failed(elseAddress) || thenAddress->base != elseAddress->base)
|
||||
return mlir::failure();
|
||||
|
||||
auto condition = compileIndexValueImpl(ifOp.getCondition());
|
||||
if (failed(condition))
|
||||
return mlir::failure();
|
||||
|
||||
CompiledIndexExprNode selectExpr;
|
||||
selectExpr.kind = CompiledIndexExprNode::Kind::Select;
|
||||
selectExpr.operands = {*condition, thenAddress->byteOffset, elseAddress->byteOffset};
|
||||
return CompiledAddressExpr {thenAddress->base, makeCompiledIndexExpr(std::move(selectExpr))};
|
||||
}
|
||||
|
||||
if (auto subviewOp = mlir::dyn_cast<mlir::memref::SubViewOp>(definingOp)) {
|
||||
auto sourceType = mlir::dyn_cast<mlir::MemRefType>(subviewOp.getSource().getType());
|
||||
auto subviewType = mlir::dyn_cast<mlir::MemRefType>(subviewOp.getType());
|
||||
if (!sourceType || !subviewType || !sourceType.hasStaticShape() || !subviewType.hasStaticShape())
|
||||
return mlir::failure();
|
||||
|
||||
llvm::SmallVector<int64_t> staticSizes;
|
||||
staticSizes.reserve(subviewOp.getMixedSizes().size());
|
||||
llvm::SmallVector<int64_t> staticStrides;
|
||||
staticStrides.reserve(subviewOp.getMixedStrides().size());
|
||||
llvm::SmallVector<int64_t> staticOffsets;
|
||||
staticOffsets.reserve(subviewOp.getMixedOffsets().size());
|
||||
bool hasOnlyStaticOffsets = true;
|
||||
|
||||
for (mlir::OpFoldResult offset : subviewOp.getMixedOffsets())
|
||||
if (auto attr = mlir::dyn_cast<mlir::Attribute>(offset))
|
||||
staticOffsets.push_back(mlir::cast<mlir::IntegerAttr>(attr).getInt());
|
||||
else
|
||||
hasOnlyStaticOffsets = false;
|
||||
for (mlir::OpFoldResult size : subviewOp.getMixedSizes()) {
|
||||
auto attr = mlir::dyn_cast<mlir::Attribute>(size);
|
||||
if (!attr)
|
||||
return mlir::failure();
|
||||
staticSizes.push_back(mlir::cast<mlir::IntegerAttr>(attr).getInt());
|
||||
}
|
||||
for (mlir::OpFoldResult stride : subviewOp.getMixedStrides()) {
|
||||
auto attr = mlir::dyn_cast<mlir::Attribute>(stride);
|
||||
if (!attr)
|
||||
return mlir::failure();
|
||||
staticStrides.push_back(mlir::cast<mlir::IntegerAttr>(attr).getInt());
|
||||
}
|
||||
|
||||
if (!isContiguousSubviewWithDynamicOffsets(
|
||||
sourceType.getShape(), subviewOp.getMixedOffsets(), staticSizes, staticStrides)) {
|
||||
return mlir::failure();
|
||||
}
|
||||
|
||||
if (hasOnlyStaticOffsets) {
|
||||
if (!isMemoryContiguous(sourceType.getShape(), staticOffsets, staticSizes, staticStrides))
|
||||
return mlir::failure();
|
||||
|
||||
auto sourceStrides = getStaticMemRefStrides(sourceType);
|
||||
if (failed(sourceStrides))
|
||||
return mlir::failure();
|
||||
constantByteOffset +=
|
||||
linearizeIndex(staticOffsets, *sourceStrides) * getElementTypeSizeInBytes(subviewType.getElementType());
|
||||
}
|
||||
else {
|
||||
auto sourceStrides = getStaticMemRefStrides(sourceType);
|
||||
if (failed(sourceStrides))
|
||||
return mlir::failure();
|
||||
CompiledIndexExpr offsetExpr;
|
||||
{
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Constant;
|
||||
expr.constant = 0;
|
||||
offsetExpr = makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
|
||||
for (auto [mixedOffset, sourceStride] : llvm::zip_equal(subviewOp.getMixedOffsets(), *sourceStrides)) {
|
||||
CompiledIndexExpr operandExpr;
|
||||
if (auto attr = mlir::dyn_cast<mlir::Attribute>(mixedOffset)) {
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Constant;
|
||||
expr.constant = mlir::cast<mlir::IntegerAttr>(attr).getInt() * sourceStride
|
||||
* getElementTypeSizeInBytes(subviewType.getElementType());
|
||||
operandExpr = makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
else {
|
||||
auto compiledOffset = compileIndexValueImpl(mlir::cast<mlir::Value>(mixedOffset));
|
||||
if (failed(compiledOffset))
|
||||
return mlir::failure();
|
||||
CompiledIndexExpr scaleExpr;
|
||||
{
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Constant;
|
||||
expr.constant = sourceStride * getElementTypeSizeInBytes(subviewType.getElementType());
|
||||
scaleExpr = makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Mul;
|
||||
expr.operands = {*compiledOffset, scaleExpr};
|
||||
operandExpr = makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Add;
|
||||
expr.operands = {offsetExpr, operandExpr};
|
||||
offsetExpr = makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
|
||||
CompiledIndexExpr constantExpr;
|
||||
{
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Constant;
|
||||
expr.constant = constantByteOffset;
|
||||
constantExpr = makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Add;
|
||||
expr.operands = {constantExpr, offsetExpr};
|
||||
byteOffsetExpr = makeCompiledIndexExpr(std::move(expr));
|
||||
constantByteOffset = 0;
|
||||
}
|
||||
|
||||
value = subviewOp.getSource();
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto castOp = mlir::dyn_cast<mlir::memref::CastOp>(definingOp)) {
|
||||
value = castOp.getSource();
|
||||
continue;
|
||||
}
|
||||
if (auto collapseOp = mlir::dyn_cast<mlir::memref::CollapseShapeOp>(definingOp)) {
|
||||
value = collapseOp.getSrc();
|
||||
continue;
|
||||
}
|
||||
if (auto expandOp = mlir::dyn_cast<mlir::memref::ExpandShapeOp>(definingOp)) {
|
||||
value = expandOp.getSrc();
|
||||
continue;
|
||||
}
|
||||
|
||||
if (mlir::isa<mlir::memref::AllocOp, mlir::memref::GetGlobalOp>(definingOp)) {
|
||||
if (constantByteOffset != 0) {
|
||||
CompiledIndexExpr constantExpr;
|
||||
{
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Constant;
|
||||
expr.constant = constantByteOffset;
|
||||
constantExpr = makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
if (byteOffsetExpr.node->kind == CompiledIndexExprNode::Kind::Constant && byteOffsetExpr.node->constant == 0)
|
||||
byteOffsetExpr = constantExpr;
|
||||
else {
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Add;
|
||||
expr.operands = {constantExpr, byteOffsetExpr};
|
||||
byteOffsetExpr = makeCompiledIndexExpr(std::move(expr));
|
||||
}
|
||||
}
|
||||
return CompiledAddressExpr {value, byteOffsetExpr};
|
||||
}
|
||||
|
||||
return mlir::failure();
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
llvm::FailureOr<int64_t> resolveIndexValue(mlir::Value value, const StaticValueKnowledge& knowledge) {
|
||||
return resolveIndexValueImpl(value, &knowledge);
|
||||
}
|
||||
|
||||
llvm::FailureOr<ResolvedContiguousAddress> resolveContiguousAddress(mlir::Value value) {
|
||||
return resolveContiguousAddressImpl(value, nullptr);
|
||||
}
|
||||
llvm::FailureOr<CompiledIndexExpr> compileIndexExpr(mlir::Value value) { return compileIndexValueImpl(value); }
|
||||
|
||||
llvm::FailureOr<ResolvedContiguousAddress> resolveContiguousAddress(mlir::Value value,
|
||||
const StaticValueKnowledge& knowledge) {
|
||||
@@ -255,4 +901,21 @@ mlir::Value resolveLoopCarriedAlias(mlir::Value value, const StaticValueKnowledg
|
||||
return resolveLoopCarriedAliasImpl(value, &knowledge);
|
||||
}
|
||||
|
||||
llvm::FailureOr<CompiledAddressExpr> compileContiguousAddressExpr(mlir::Value value) {
|
||||
return compileContiguousAddressExprImpl(value);
|
||||
}
|
||||
|
||||
llvm::FailureOr<int64_t> CompiledIndexExpr::evaluate(const StaticValueKnowledge& knowledge) const {
|
||||
return evaluateCompiledIndexExpr(*this, knowledge);
|
||||
}
|
||||
|
||||
llvm::FailureOr<ResolvedContiguousAddress> CompiledAddressExpr::evaluate(const StaticValueKnowledge& knowledge,
|
||||
std::optional<unsigned> lane) const {
|
||||
(void) lane;
|
||||
auto resolvedOffset = byteOffset.evaluate(knowledge);
|
||||
if (failed(resolvedOffset))
|
||||
return mlir::failure();
|
||||
return ResolvedContiguousAddress {resolveAlias(base, &knowledge), *resolvedOffset};
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -1,10 +1,14 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/Dialect/Arith/IR/Arith.h"
|
||||
#include "mlir/Dialect/MemRef/IR/MemRef.h"
|
||||
#include "mlir/IR/Value.h"
|
||||
|
||||
#include "llvm/ADT/DenseMap.h"
|
||||
|
||||
#include <memory>
|
||||
#include <optional>
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
/// Describes a value as a base addressable object plus a statically known
|
||||
@@ -23,21 +27,68 @@ struct StaticValueKnowledge {
|
||||
StaticValueKnowledge() {}
|
||||
};
|
||||
|
||||
struct CompiledIndexExprNode;
|
||||
|
||||
struct CompiledIndexExpr {
|
||||
std::shared_ptr<CompiledIndexExprNode> node;
|
||||
|
||||
CompiledIndexExpr() = default;
|
||||
explicit CompiledIndexExpr(std::shared_ptr<CompiledIndexExprNode> node)
|
||||
: node(std::move(node)) {}
|
||||
|
||||
llvm::FailureOr<int64_t> evaluate(const StaticValueKnowledge& knowledge) const;
|
||||
};
|
||||
|
||||
struct CompiledIndexExprNode {
|
||||
enum class Kind {
|
||||
Constant,
|
||||
Symbol,
|
||||
Add,
|
||||
Sub,
|
||||
Mul,
|
||||
DivUI,
|
||||
DivSI,
|
||||
RemUI,
|
||||
RemSI,
|
||||
MinUI,
|
||||
CmpI,
|
||||
Select,
|
||||
ConstantGlobalLoad
|
||||
};
|
||||
|
||||
Kind kind = Kind::Constant;
|
||||
int64_t constant = 0;
|
||||
mlir::Value symbol;
|
||||
mlir::arith::CmpIPredicate predicate = mlir::arith::CmpIPredicate::eq;
|
||||
mlir::memref::GlobalOp globalOp;
|
||||
llvm::SmallVector<int64_t, 4> globalStrides;
|
||||
llvm::SmallVector<CompiledIndexExpr, 4> operands;
|
||||
};
|
||||
|
||||
struct CompiledAddressExpr {
|
||||
mlir::Value base;
|
||||
CompiledIndexExpr byteOffset;
|
||||
|
||||
llvm::FailureOr<ResolvedContiguousAddress> evaluate(const StaticValueKnowledge& knowledge,
|
||||
std::optional<unsigned> lane) const;
|
||||
};
|
||||
|
||||
mlir::memref::GlobalOp lookupGlobalForGetGlobal(mlir::ModuleOp moduleOp, mlir::memref::GetGlobalOp getGlobalOp);
|
||||
|
||||
/// Resolves a value to contiguous backing storage when that storage can be
|
||||
/// proven statically from aliases, DPS ties, casts, and subviews.
|
||||
llvm::FailureOr<ResolvedContiguousAddress> resolveContiguousAddress(mlir::Value value);
|
||||
llvm::FailureOr<ResolvedContiguousAddress> resolveContiguousAddress(mlir::Value value,
|
||||
const StaticValueKnowledge& knowledge);
|
||||
const StaticValueKnowledge& knowledge = {});
|
||||
|
||||
/// Statically evaluates index-like SSA values, including simple integer
|
||||
/// arithmetic and loop facts recorded in `knowledge`.
|
||||
llvm::FailureOr<int64_t> resolveIndexValue(mlir::Value value);
|
||||
llvm::FailureOr<int64_t> resolveIndexValue(mlir::Value value, const StaticValueKnowledge& knowledge);
|
||||
llvm::FailureOr<int64_t> resolveIndexValue(mlir::Value value, const StaticValueKnowledge& knowledge = {});
|
||||
llvm::FailureOr<CompiledIndexExpr> compileIndexExpr(mlir::Value value);
|
||||
|
||||
/// Follows alias, view, and DPS chains to recover the backing value of a
|
||||
/// loop-carried memref/result.
|
||||
mlir::Value resolveLoopCarriedAlias(mlir::Value value, const StaticValueKnowledge& knowledge);
|
||||
|
||||
llvm::FailureOr<CompiledAddressExpr> compileContiguousAddressExpr(mlir::Value value);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -0,0 +1,219 @@
|
||||
#include "mlir/Dialect/Arith/IR/Arith.h"
|
||||
#include "mlir/IR/Matchers.h"
|
||||
|
||||
#include "AffineUtils.hpp"
|
||||
#include "ConstantUtils.hpp"
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
static FailureOr<int64_t> floorDivSigned(int64_t lhs, int64_t rhs) {
|
||||
if (rhs <= 0)
|
||||
return failure();
|
||||
|
||||
int64_t quotient = lhs / rhs;
|
||||
int64_t remainder = lhs % rhs;
|
||||
if (remainder != 0 && lhs < 0)
|
||||
--quotient;
|
||||
return quotient;
|
||||
}
|
||||
|
||||
static FailureOr<int64_t> ceilDivSigned(int64_t lhs, int64_t rhs) {
|
||||
if (rhs <= 0)
|
||||
return failure();
|
||||
|
||||
int64_t quotient = lhs / rhs;
|
||||
int64_t remainder = lhs % rhs;
|
||||
if (remainder != 0 && lhs > 0)
|
||||
++quotient;
|
||||
return quotient;
|
||||
}
|
||||
|
||||
Value createOrFoldAffineApply(
|
||||
OpBuilder& builder, Location loc, AffineMap map, ValueRange operands, Operation* constantAnchor) {
|
||||
assert(constantAnchor && "expected a valid constant anchor");
|
||||
assert(map.getNumResults() == 1 && "affine.apply expects a single-result affine map");
|
||||
|
||||
SmallVector<Attribute> operandConstants;
|
||||
operandConstants.reserve(operands.size());
|
||||
for (Value operand : operands) {
|
||||
std::optional<int64_t> constantValue = matchConstantIndexValue(operand);
|
||||
if (!constantValue)
|
||||
return affine::AffineApplyOp::create(builder, loc, map, operands).getResult();
|
||||
operandConstants.push_back(builder.getIndexAttr(*constantValue));
|
||||
}
|
||||
|
||||
SmallVector<Attribute> foldedResults;
|
||||
if (succeeded(map.constantFold(operandConstants, foldedResults)) && foldedResults.size() == 1)
|
||||
if (auto constantResult = dyn_cast<IntegerAttr>(foldedResults.front()))
|
||||
return getOrCreateIndexConstant(builder, constantAnchor, constantResult.getInt());
|
||||
|
||||
return affine::AffineApplyOp::create(builder, loc, map, operands).getResult();
|
||||
}
|
||||
|
||||
Value createOrFoldAffineApply(
|
||||
OpBuilder& builder, Location loc, AffineExpr expr, ValueRange dims, Operation* constantAnchor) {
|
||||
AffineMap map = AffineMap::get(/*dimCount=*/dims.size(), /*symbolCount=*/0, expr);
|
||||
return createOrFoldAffineApply(builder, loc, map, dims, constantAnchor);
|
||||
}
|
||||
|
||||
Value affineMulConst(OpBuilder& builder, Location loc, Value value, int64_t multiplier, Operation* constantAnchor) {
|
||||
assert(constantAnchor && "expected a valid constant anchor");
|
||||
if (multiplier == 0)
|
||||
return getOrCreateIndexConstant(builder, constantAnchor, 0);
|
||||
if (multiplier == 1)
|
||||
return value;
|
||||
|
||||
AffineExpr d0 = getAffineDimExpr(0, builder.getContext());
|
||||
return createOrFoldAffineApply(builder, loc, d0 * multiplier, ValueRange {value}, constantAnchor);
|
||||
}
|
||||
|
||||
Value affineAddConst(OpBuilder& builder, Location loc, Value value, int64_t offset, Operation* constantAnchor) {
|
||||
assert(constantAnchor && "expected a valid constant anchor");
|
||||
if (offset == 0)
|
||||
return value;
|
||||
|
||||
AffineExpr d0 = getAffineDimExpr(0, builder.getContext());
|
||||
return createOrFoldAffineApply(builder, loc, d0 + offset, ValueRange {value}, constantAnchor);
|
||||
}
|
||||
|
||||
Value affineModConst(OpBuilder& builder, Location loc, Value value, int64_t divisor, Operation* constantAnchor) {
|
||||
assert(constantAnchor && "expected a valid constant anchor");
|
||||
assert(divisor > 0 && "expected a positive affine.mod divisor");
|
||||
if (divisor == 1)
|
||||
return getOrCreateIndexConstant(builder, constantAnchor, 0);
|
||||
|
||||
AffineExpr d0 = getAffineDimExpr(0, builder.getContext());
|
||||
return createOrFoldAffineApply(builder, loc, d0 % divisor, ValueRange {value}, constantAnchor);
|
||||
}
|
||||
|
||||
Value affineFloorDivConst(
|
||||
OpBuilder& builder, Location loc, Value value, int64_t divisor, Operation* constantAnchor) {
|
||||
assert(constantAnchor && "expected a valid constant anchor");
|
||||
assert(divisor > 0 && "expected a positive affine.floor_div divisor");
|
||||
if (divisor == 1)
|
||||
return value;
|
||||
|
||||
AffineExpr d0 = getAffineDimExpr(0, builder.getContext());
|
||||
return createOrFoldAffineApply(builder, loc, d0.floorDiv(divisor), ValueRange {value}, constantAnchor);
|
||||
}
|
||||
|
||||
Value affineAddModConst(
|
||||
OpBuilder& builder, Location loc, Value value, int64_t offset, int64_t divisor, Operation* constantAnchor) {
|
||||
assert(constantAnchor && "expected a valid constant anchor");
|
||||
assert(divisor > 0 && "expected a positive affine.mod divisor");
|
||||
if (divisor == 1)
|
||||
return getOrCreateIndexConstant(builder, constantAnchor, 0);
|
||||
|
||||
AffineExpr d0 = getAffineDimExpr(0, builder.getContext());
|
||||
AffineExpr expr = d0;
|
||||
if (offset != 0)
|
||||
expr = expr + offset;
|
||||
return createOrFoldAffineApply(builder, loc, expr % divisor, ValueRange {value}, constantAnchor);
|
||||
}
|
||||
|
||||
Value affineAddFloorDivConst(
|
||||
OpBuilder& builder, Location loc, Value value, int64_t offset, int64_t divisor, Operation* constantAnchor) {
|
||||
assert(constantAnchor && "expected a valid constant anchor");
|
||||
assert(divisor > 0 && "expected a positive affine.floor_div divisor");
|
||||
if (divisor == 1)
|
||||
return offset == 0 ? value : affineAddConst(builder, loc, value, offset, constantAnchor);
|
||||
|
||||
AffineExpr d0 = getAffineDimExpr(0, builder.getContext());
|
||||
AffineExpr expr = d0;
|
||||
if (offset != 0)
|
||||
expr = expr + offset;
|
||||
return createOrFoldAffineApply(builder, loc, expr.floorDiv(divisor), ValueRange {value}, constantAnchor);
|
||||
}
|
||||
|
||||
FailureOr<int64_t> evaluateAffineExpr(AffineExpr expr, ArrayRef<int64_t> dims, ArrayRef<int64_t> symbols) {
|
||||
if (auto constant = dyn_cast<AffineConstantExpr>(expr))
|
||||
return constant.getValue();
|
||||
if (auto dim = dyn_cast<AffineDimExpr>(expr)) {
|
||||
unsigned position = dim.getPosition();
|
||||
if (position >= dims.size())
|
||||
return failure();
|
||||
return dims[position];
|
||||
}
|
||||
if (auto symbol = dyn_cast<AffineSymbolExpr>(expr)) {
|
||||
unsigned position = symbol.getPosition();
|
||||
if (position >= symbols.size())
|
||||
return failure();
|
||||
return symbols[position];
|
||||
}
|
||||
|
||||
auto binary = dyn_cast<AffineBinaryOpExpr>(expr);
|
||||
if (!binary)
|
||||
return failure();
|
||||
|
||||
FailureOr<int64_t> lhs = evaluateAffineExpr(binary.getLHS(), dims, symbols);
|
||||
FailureOr<int64_t> rhs = evaluateAffineExpr(binary.getRHS(), dims, symbols);
|
||||
if (failed(lhs) || failed(rhs))
|
||||
return failure();
|
||||
|
||||
switch (binary.getKind()) {
|
||||
case AffineExprKind::Add: return *lhs + *rhs;
|
||||
case AffineExprKind::Mul: return *lhs * *rhs;
|
||||
case AffineExprKind::FloorDiv: return floorDivSigned(*lhs, *rhs);
|
||||
case AffineExprKind::CeilDiv: return ceilDivSigned(*lhs, *rhs);
|
||||
case AffineExprKind::Mod: {
|
||||
FailureOr<int64_t> div = floorDivSigned(*lhs, *rhs);
|
||||
if (failed(div))
|
||||
return failure();
|
||||
return *lhs - *div * *rhs;
|
||||
}
|
||||
default: return failure();
|
||||
}
|
||||
}
|
||||
|
||||
FailureOr<int64_t> evaluateSingleResultAffineMap(AffineMap map, ArrayRef<int64_t> operands) {
|
||||
if (map.getNumResults() != 1 || operands.size() != map.getNumInputs())
|
||||
return failure();
|
||||
|
||||
ArrayRef<int64_t> dims(operands.data(), map.getNumDims());
|
||||
ArrayRef<int64_t> symbols(operands.data() + map.getNumDims(), map.getNumSymbols());
|
||||
return evaluateAffineExpr(map.getResult(0), dims, symbols);
|
||||
}
|
||||
|
||||
FailureOr<int64_t> evaluateAffineApply(affine::AffineApplyOp affineApply, IndexValueResolver resolver) {
|
||||
SmallVector<int64_t, 4> operands;
|
||||
operands.reserve(affineApply.getMapOperands().size());
|
||||
for (Value operand : affineApply.getMapOperands()) {
|
||||
FailureOr<int64_t> folded = resolver(operand);
|
||||
if (failed(folded))
|
||||
return failure();
|
||||
operands.push_back(*folded);
|
||||
}
|
||||
|
||||
return evaluateSingleResultAffineMap(affineApply.getAffineMap(), operands);
|
||||
}
|
||||
|
||||
bool isSingleResultSymbolFreeAffineMap(AffineMap map) { return map.getNumResults() == 1 && map.getNumSymbols() == 0; }
|
||||
|
||||
bool isDimAndConstantAffineExpr(AffineExpr expr) {
|
||||
switch (expr.getKind()) {
|
||||
case AffineExprKind::Constant:
|
||||
case AffineExprKind::DimId: return true;
|
||||
case AffineExprKind::SymbolId: return false;
|
||||
case AffineExprKind::Add: {
|
||||
auto binaryExpr = cast<AffineBinaryOpExpr>(expr);
|
||||
return isDimAndConstantAffineExpr(binaryExpr.getLHS()) && isDimAndConstantAffineExpr(binaryExpr.getRHS());
|
||||
}
|
||||
case AffineExprKind::Mul: {
|
||||
auto binaryExpr = cast<AffineBinaryOpExpr>(expr);
|
||||
return (isa<AffineConstantExpr>(binaryExpr.getLHS()) && isDimAndConstantAffineExpr(binaryExpr.getRHS()))
|
||||
|| (isa<AffineConstantExpr>(binaryExpr.getRHS()) && isDimAndConstantAffineExpr(binaryExpr.getLHS()));
|
||||
}
|
||||
case AffineExprKind::FloorDiv:
|
||||
case AffineExprKind::CeilDiv:
|
||||
case AffineExprKind::Mod: {
|
||||
auto binaryExpr = cast<AffineBinaryOpExpr>(expr);
|
||||
return isa<AffineConstantExpr>(binaryExpr.getRHS()) && isDimAndConstantAffineExpr(binaryExpr.getLHS());
|
||||
}
|
||||
}
|
||||
|
||||
llvm_unreachable("unexpected affine expression kind");
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,75 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/Dialect/Affine/IR/AffineOps.h"
|
||||
#include "mlir/IR/AffineExpr.h"
|
||||
#include "mlir/IR/PatternMatch.h"
|
||||
#include "mlir/Support/LogicalResult.h"
|
||||
|
||||
#include "llvm/ADT/FunctionExtras.h"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
using IndexValueResolver = llvm::function_ref<llvm::FailureOr<int64_t>(mlir::Value)>;
|
||||
|
||||
mlir::Value createOrFoldAffineApply(mlir::OpBuilder& builder,
|
||||
mlir::Location loc,
|
||||
mlir::AffineMap map,
|
||||
mlir::ValueRange operands,
|
||||
mlir::Operation* constantAnchor);
|
||||
|
||||
mlir::Value createOrFoldAffineApply(mlir::OpBuilder& builder,
|
||||
mlir::Location loc,
|
||||
mlir::AffineExpr expr,
|
||||
mlir::ValueRange dims,
|
||||
mlir::Operation* constantAnchor);
|
||||
|
||||
mlir::Value affineMulConst(mlir::OpBuilder& builder,
|
||||
mlir::Location loc,
|
||||
mlir::Value value,
|
||||
int64_t multiplier,
|
||||
mlir::Operation* constantAnchor);
|
||||
|
||||
mlir::Value affineAddConst(mlir::OpBuilder& builder,
|
||||
mlir::Location loc,
|
||||
mlir::Value value,
|
||||
int64_t offset,
|
||||
mlir::Operation* constantAnchor);
|
||||
|
||||
mlir::Value affineModConst(mlir::OpBuilder& builder,
|
||||
mlir::Location loc,
|
||||
mlir::Value value,
|
||||
int64_t divisor,
|
||||
mlir::Operation* constantAnchor);
|
||||
|
||||
mlir::Value affineFloorDivConst(mlir::OpBuilder& builder,
|
||||
mlir::Location loc,
|
||||
mlir::Value value,
|
||||
int64_t divisor,
|
||||
mlir::Operation* constantAnchor);
|
||||
|
||||
mlir::Value affineAddModConst(mlir::OpBuilder& builder,
|
||||
mlir::Location loc,
|
||||
mlir::Value value,
|
||||
int64_t offset,
|
||||
int64_t divisor,
|
||||
mlir::Operation* constantAnchor);
|
||||
|
||||
mlir::Value affineAddFloorDivConst(mlir::OpBuilder& builder,
|
||||
mlir::Location loc,
|
||||
mlir::Value value,
|
||||
int64_t offset,
|
||||
int64_t divisor,
|
||||
mlir::Operation* constantAnchor);
|
||||
|
||||
llvm::FailureOr<int64_t>
|
||||
evaluateAffineExpr(mlir::AffineExpr expr, llvm::ArrayRef<int64_t> dims, llvm::ArrayRef<int64_t> symbols = {});
|
||||
|
||||
llvm::FailureOr<int64_t> evaluateSingleResultAffineMap(mlir::AffineMap map, llvm::ArrayRef<int64_t> operands);
|
||||
|
||||
llvm::FailureOr<int64_t> evaluateAffineApply(mlir::affine::AffineApplyOp affineApply, IndexValueResolver resolver);
|
||||
|
||||
bool isSingleResultSymbolFreeAffineMap(mlir::AffineMap map);
|
||||
|
||||
bool isDimAndConstantAffineExpr(mlir::AffineExpr expr);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,92 @@
|
||||
#include "src/Accelerators/PIM/Common/IR/BatchCoreUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/PimCommon.hpp"
|
||||
#include "src/Accelerators/PIM/Common/Support/CheckedArithmetic.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
llvm::SmallVector<int32_t> getBatchCoreIds(pim::PimCoreBatchOp coreBatchOp) {
|
||||
auto coreIdsAttr = coreBatchOp->getAttrOfType<mlir::DenseI32ArrayAttr>(onnx_mlir::kCoreIdsAttrName);
|
||||
assert(coreIdsAttr && "pim.core_batch requires coreIds array attribute");
|
||||
return llvm::SmallVector<int32_t>(coreIdsAttr.asArrayRef().begin(), coreIdsAttr.asArrayRef().end());
|
||||
}
|
||||
|
||||
mlir::FailureOr<std::optional<int32_t>>
|
||||
getOptionalScheduledCoreId(spatial::SpatScheduledCompute computeOp, llvm::StringRef fieldName) {
|
||||
auto coreIdAttr = computeOp->getAttrOfType<mlir::IntegerAttr>(onnx_mlir::kCoreIdAttrName);
|
||||
if (!coreIdAttr)
|
||||
return std::optional<int32_t> {};
|
||||
if (coreIdAttr.getInt() < 0) {
|
||||
computeOp.emitOpError() << fieldName << " must be non-negative";
|
||||
return mlir::failure();
|
||||
}
|
||||
auto checkedCoreId = pim::checkedI32(coreIdAttr.getInt(), computeOp, fieldName);
|
||||
if (mlir::failed(checkedCoreId))
|
||||
return mlir::failure();
|
||||
return std::optional<int32_t> {*checkedCoreId};
|
||||
}
|
||||
|
||||
mlir::FailureOr<int32_t> getRequiredScheduledCoreId(spatial::SpatScheduledCompute computeOp, llvm::StringRef fieldName) {
|
||||
auto coreId = getOptionalScheduledCoreId(computeOp, fieldName);
|
||||
if (mlir::failed(coreId))
|
||||
return mlir::failure();
|
||||
if (!*coreId) {
|
||||
computeOp.emitOpError() << "missing required " << onnx_mlir::kCoreIdAttrName;
|
||||
return mlir::failure();
|
||||
}
|
||||
return **coreId;
|
||||
}
|
||||
|
||||
mlir::FailureOr<std::optional<llvm::SmallVector<int32_t>>>
|
||||
getOptionalScheduledBatchCoreIds(spatial::SpatScheduledComputeBatch computeBatchOp, llvm::StringRef fieldName) {
|
||||
auto coreIdsAttr = computeBatchOp->getAttrOfType<mlir::DenseI32ArrayAttr>(onnx_mlir::kCoreIdsAttrName);
|
||||
if (!coreIdsAttr)
|
||||
return std::optional<llvm::SmallVector<int32_t>> {};
|
||||
|
||||
llvm::SmallVector<int32_t> coreIds;
|
||||
coreIds.reserve(coreIdsAttr.size());
|
||||
for (int32_t coreId : coreIdsAttr.asArrayRef()) {
|
||||
if (coreId < 0) {
|
||||
computeBatchOp.emitOpError() << fieldName << " values must be non-negative";
|
||||
return mlir::failure();
|
||||
}
|
||||
auto checkedCoreId = pim::checkedI32(static_cast<int64_t>(coreId), computeBatchOp, fieldName);
|
||||
if (mlir::failed(checkedCoreId))
|
||||
return mlir::failure();
|
||||
coreIds.push_back(*checkedCoreId);
|
||||
}
|
||||
return std::optional<llvm::SmallVector<int32_t>> {std::move(coreIds)};
|
||||
}
|
||||
|
||||
mlir::FailureOr<llvm::SmallVector<int32_t>>
|
||||
getRequiredScheduledBatchCoreIds(spatial::SpatScheduledComputeBatch computeBatchOp, llvm::StringRef fieldName) {
|
||||
auto coreIds = getOptionalScheduledBatchCoreIds(computeBatchOp, fieldName);
|
||||
if (mlir::failed(coreIds))
|
||||
return mlir::failure();
|
||||
if (!*coreIds) {
|
||||
computeBatchOp.emitOpError() << "missing required " << onnx_mlir::kCoreIdsAttrName;
|
||||
return mlir::failure();
|
||||
}
|
||||
return std::move(**coreIds);
|
||||
}
|
||||
|
||||
llvm::SmallVector<int32_t> getLaneChunkCoreIds(llvm::ArrayRef<int32_t> coreIds, size_t laneCount, unsigned lane) {
|
||||
llvm::SmallVector<int32_t> laneCoreIds;
|
||||
laneCoreIds.reserve(coreIds.size() / laneCount);
|
||||
for (size_t chunkIndex = 0; chunkIndex < coreIds.size() / laneCount; ++chunkIndex)
|
||||
laneCoreIds.push_back(coreIds[chunkIndex * laneCount + lane]);
|
||||
return laneCoreIds;
|
||||
}
|
||||
|
||||
bool isExplicitHostMemCopyOperand(mlir::Operation* op, unsigned operandIndex) {
|
||||
if (mlir::isa<pim::PimMemCopyHostToDevOp>(op))
|
||||
return operandIndex == 3;
|
||||
if (mlir::isa<pim::PimMemCopyDevToHostOp>(op))
|
||||
return operandIndex == 2;
|
||||
return false;
|
||||
}
|
||||
|
||||
bool isExplicitDevToHostTargetOperand(mlir::Operation* op, unsigned operandIndex) {
|
||||
return mlir::isa<pim::PimMemCopyDevToHostOp>(op) && operandIndex == 2;
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,32 @@
|
||||
#pragma once
|
||||
|
||||
#include "llvm/ADT/ArrayRef.h"
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
|
||||
#include <optional>
|
||||
|
||||
#include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
llvm::SmallVector<int32_t> getBatchCoreIds(pim::PimCoreBatchOp coreBatchOp);
|
||||
|
||||
mlir::FailureOr<std::optional<int32_t>>
|
||||
getOptionalScheduledCoreId(spatial::SpatScheduledCompute computeOp, llvm::StringRef fieldName);
|
||||
|
||||
mlir::FailureOr<int32_t> getRequiredScheduledCoreId(spatial::SpatScheduledCompute computeOp, llvm::StringRef fieldName);
|
||||
|
||||
mlir::FailureOr<std::optional<llvm::SmallVector<int32_t>>>
|
||||
getOptionalScheduledBatchCoreIds(spatial::SpatScheduledComputeBatch computeBatchOp, llvm::StringRef fieldName);
|
||||
|
||||
mlir::FailureOr<llvm::SmallVector<int32_t>>
|
||||
getRequiredScheduledBatchCoreIds(spatial::SpatScheduledComputeBatch computeBatchOp, llvm::StringRef fieldName);
|
||||
|
||||
llvm::SmallVector<int32_t> getLaneChunkCoreIds(llvm::ArrayRef<int32_t> coreIds, size_t laneCount, unsigned lane);
|
||||
|
||||
bool isExplicitHostMemCopyOperand(mlir::Operation* op, unsigned operandIndex);
|
||||
|
||||
bool isExplicitDevToHostTargetOperand(mlir::Operation* op, unsigned operandIndex);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,794 @@
|
||||
#ifndef ONNX_MLIR_PIM_COMPACT_ASM_UTILS_HPP
|
||||
#define ONNX_MLIR_PIM_COMPACT_ASM_UTILS_HPP
|
||||
|
||||
#include "mlir/IR/Builders.h"
|
||||
#include "mlir/IR/OpImplementation.h"
|
||||
#include "mlir/IR/Value.h"
|
||||
#include "mlir/Support/LLVM.h"
|
||||
|
||||
#include "llvm/ADT/ArrayRef.h"
|
||||
#include "llvm/ADT/STLExtras.h"
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
#include "llvm/ADT/Twine.h"
|
||||
#include "llvm/Support/LogicalResult.h"
|
||||
|
||||
#include <optional>
|
||||
|
||||
namespace onnx_mlir {
|
||||
namespace compact_asm {
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
enum class ListDelimiter {
|
||||
Square,
|
||||
Paren
|
||||
};
|
||||
|
||||
struct NumericSuffixRange {
|
||||
StringRef prefix;
|
||||
unsigned first = 0;
|
||||
unsigned last = 0;
|
||||
};
|
||||
|
||||
inline std::optional<NumericSuffixRange> getNumericSuffixRange(StringRef firstName, StringRef lastName) {
|
||||
auto split = [](StringRef name) -> std::optional<std::pair<StringRef, unsigned>> {
|
||||
size_t suffixStart = name.size();
|
||||
while (suffixStart > 0 && name[suffixStart - 1] >= '0' && name[suffixStart - 1] <= '9')
|
||||
--suffixStart;
|
||||
if (suffixStart == name.size())
|
||||
return std::nullopt;
|
||||
|
||||
unsigned number = 0;
|
||||
if (name.drop_front(suffixStart).getAsInteger(10, number))
|
||||
return std::nullopt;
|
||||
return std::pair(name.take_front(suffixStart), number);
|
||||
};
|
||||
|
||||
auto first = split(firstName);
|
||||
auto last = split(lastName);
|
||||
if (!first || !last || first->first != last->first || first->second > last->second)
|
||||
return std::nullopt;
|
||||
return NumericSuffixRange {first->first, first->second, last->second};
|
||||
}
|
||||
|
||||
inline StringRef getInternedNumberedName(OpAsmParser& parser, StringRef prefix, unsigned number) {
|
||||
return parser.getBuilder().getStringAttr((Twine(prefix) + Twine(number)).str()).getValue();
|
||||
}
|
||||
|
||||
inline ParseResult parseOptionalRepeatCount(OpAsmParser& parser, int64_t& repeatCount) {
|
||||
repeatCount = 1;
|
||||
if (succeeded(parser.parseOptionalKeyword("x"))) {
|
||||
if (parser.parseInteger(repeatCount) || repeatCount <= 0)
|
||||
return parser.emitError(parser.getCurrentLocation(), "repeat count after 'x' must be positive");
|
||||
return success();
|
||||
}
|
||||
|
||||
const char* tokenStart = parser.getCurrentLocation().getPointer();
|
||||
if (!tokenStart || tokenStart[0] != 'x' || tokenStart[1] < '0' || tokenStart[1] > '9')
|
||||
return success();
|
||||
|
||||
StringRef fusedRepeat;
|
||||
if (failed(parser.parseOptionalKeyword(&fusedRepeat)))
|
||||
return failure();
|
||||
if (!fusedRepeat.consume_front("x") || fusedRepeat.empty()
|
||||
|| fusedRepeat.getAsInteger(10, repeatCount) || repeatCount <= 0) {
|
||||
return parser.emitError(parser.getCurrentLocation(), "repeat count after 'x' must be positive");
|
||||
}
|
||||
return success();
|
||||
}
|
||||
|
||||
inline ParseResult parseOpenDelimiter(OpAsmParser& parser, ListDelimiter delimiter) {
|
||||
if (delimiter == ListDelimiter::Square)
|
||||
return parser.parseLSquare();
|
||||
return parser.parseLParen();
|
||||
}
|
||||
|
||||
inline ParseResult parseOptionalCloseDelimiter(OpAsmParser& parser, ListDelimiter delimiter) {
|
||||
if (delimiter == ListDelimiter::Square)
|
||||
return parser.parseOptionalRSquare();
|
||||
return parser.parseOptionalRParen();
|
||||
}
|
||||
|
||||
template <typename StreamT>
|
||||
inline void printOpenDelimiter(StreamT& stream, ListDelimiter delimiter) {
|
||||
stream << (delimiter == ListDelimiter::Square ? "[" : "(");
|
||||
}
|
||||
|
||||
template <typename StreamT>
|
||||
inline void printCloseDelimiter(StreamT& stream, ListDelimiter delimiter) {
|
||||
stream << (delimiter == ListDelimiter::Square ? "]" : ")");
|
||||
}
|
||||
|
||||
template <typename EntryT, typename ParseEntryFn>
|
||||
inline ParseResult parseCompressedRepeatedList(OpAsmParser& parser,
|
||||
ListDelimiter delimiter,
|
||||
SmallVectorImpl<EntryT>& entries,
|
||||
ParseEntryFn parseEntry) {
|
||||
if (parseOpenDelimiter(parser, delimiter))
|
||||
return failure();
|
||||
if (succeeded(parseOptionalCloseDelimiter(parser, delimiter)))
|
||||
return success();
|
||||
|
||||
while (true) {
|
||||
EntryT entry;
|
||||
if (parseEntry(entry))
|
||||
return failure();
|
||||
|
||||
int64_t repeatCount = 1;
|
||||
if (parseOptionalRepeatCount(parser, repeatCount))
|
||||
return failure();
|
||||
for (int64_t index = 0; index < repeatCount; ++index)
|
||||
entries.push_back(entry);
|
||||
|
||||
if (succeeded(parseOptionalCloseDelimiter(parser, delimiter)))
|
||||
return success();
|
||||
if (parser.parseComma())
|
||||
return failure();
|
||||
}
|
||||
}
|
||||
|
||||
template <typename IntT>
|
||||
inline ParseResult
|
||||
parseCompressedIntegerEntries(OpAsmParser& parser, ListDelimiter delimiter, SmallVectorImpl<IntT>& values) {
|
||||
if (succeeded(parseOptionalCloseDelimiter(parser, delimiter)))
|
||||
return success();
|
||||
|
||||
while (true) {
|
||||
if (succeeded(parser.parseOptionalLParen())) {
|
||||
SmallVector<IntT> subgroup;
|
||||
if (parseCompressedIntegerEntries(parser, ListDelimiter::Paren, subgroup))
|
||||
return failure();
|
||||
|
||||
int64_t repeatCount = 1;
|
||||
if (parseOptionalRepeatCount(parser, repeatCount))
|
||||
return failure();
|
||||
for (int64_t repeat = 0; repeat < repeatCount; ++repeat)
|
||||
llvm::append_range(values, subgroup);
|
||||
}
|
||||
else {
|
||||
int64_t first = 0;
|
||||
if (parser.parseInteger(first))
|
||||
return failure();
|
||||
|
||||
if (succeeded(parser.parseOptionalKeyword("to"))) {
|
||||
int64_t last = 0;
|
||||
if (parser.parseInteger(last) || last < first)
|
||||
return parser.emitError(parser.getCurrentLocation(), "invalid ascending range");
|
||||
|
||||
int64_t step = 1;
|
||||
if (succeeded(parser.parseOptionalKeyword("by"))) {
|
||||
if (parser.parseInteger(step) || step <= 0)
|
||||
return parser.emitError(parser.getCurrentLocation(), "step after 'by' must be positive");
|
||||
}
|
||||
int64_t repeatCount = 1;
|
||||
if (parseOptionalRepeatCount(parser, repeatCount))
|
||||
return failure();
|
||||
if ((last - first) % step != 0) {
|
||||
return parser.emitError(parser.getCurrentLocation(),
|
||||
"range end must be reachable from start using the given step");
|
||||
}
|
||||
|
||||
for (int64_t value = first; value <= last; value += step)
|
||||
for (int64_t index = 0; index < repeatCount; ++index)
|
||||
values.push_back(static_cast<IntT>(value));
|
||||
}
|
||||
else {
|
||||
int64_t repeatCount = 1;
|
||||
if (parseOptionalRepeatCount(parser, repeatCount))
|
||||
return failure();
|
||||
for (int64_t index = 0; index < repeatCount; ++index)
|
||||
values.push_back(static_cast<IntT>(first));
|
||||
}
|
||||
}
|
||||
|
||||
if (succeeded(parseOptionalCloseDelimiter(parser, delimiter)))
|
||||
return success();
|
||||
if (parser.parseComma())
|
||||
return failure();
|
||||
}
|
||||
}
|
||||
|
||||
template <typename IntT>
|
||||
inline ParseResult
|
||||
parseCompressedIntegerSequence(OpAsmParser& parser, ListDelimiter delimiter, SmallVectorImpl<IntT>& values) {
|
||||
if (parseOpenDelimiter(parser, delimiter))
|
||||
return failure();
|
||||
return parseCompressedIntegerEntries(parser, delimiter, values);
|
||||
}
|
||||
|
||||
template <typename RangeT, typename PrintEntryFn>
|
||||
inline void printCompressedEqualRuns(OpAsmPrinter& printer, RangeT entries, PrintEntryFn printEntry) {
|
||||
for (size_t index = 0; index < entries.size();) {
|
||||
size_t runEnd = index + 1;
|
||||
while (runEnd < entries.size() && entries[runEnd] == entries[index])
|
||||
++runEnd;
|
||||
|
||||
if (index != 0)
|
||||
printer << ", ";
|
||||
printEntry(entries[index]);
|
||||
size_t runLength = runEnd - index;
|
||||
if (runLength > 1)
|
||||
printer << " x" << runLength;
|
||||
index = runEnd;
|
||||
}
|
||||
}
|
||||
|
||||
template <typename StreamT, typename IntT>
|
||||
inline void printCompressedIntegerEntries(StreamT& stream, ArrayRef<IntT> values) {
|
||||
struct FlatCompression {
|
||||
enum class Kind {
|
||||
Single,
|
||||
EqualRun,
|
||||
Progression
|
||||
};
|
||||
|
||||
Kind kind = Kind::Single;
|
||||
size_t covered = 1;
|
||||
size_t repeatCount = 1;
|
||||
size_t progressionValueCount = 1;
|
||||
int64_t step = 1;
|
||||
IntT firstValue {};
|
||||
IntT lastValue {};
|
||||
};
|
||||
|
||||
auto computeFlatCompression = [&](size_t start) {
|
||||
FlatCompression compression;
|
||||
compression.firstValue = values[start];
|
||||
compression.lastValue = values[start];
|
||||
|
||||
auto findEqualRunEnd = [&](size_t runStart) {
|
||||
size_t runEnd = runStart + 1;
|
||||
while (runEnd < values.size() && values[runEnd] == values[runStart])
|
||||
++runEnd;
|
||||
return runEnd;
|
||||
};
|
||||
|
||||
size_t firstRunEnd = findEqualRunEnd(start);
|
||||
compression.repeatCount = firstRunEnd - start;
|
||||
size_t progressionEnd = firstRunEnd;
|
||||
int64_t step = 0;
|
||||
IntT lastValue = values[start];
|
||||
|
||||
if (firstRunEnd < values.size()) {
|
||||
size_t secondRunEnd = findEqualRunEnd(firstRunEnd);
|
||||
step = static_cast<int64_t>(values[firstRunEnd]) - static_cast<int64_t>(values[start]);
|
||||
if (step > 0 && secondRunEnd - firstRunEnd == compression.repeatCount) {
|
||||
progressionEnd = secondRunEnd;
|
||||
lastValue = values[firstRunEnd];
|
||||
size_t currentRunStart = secondRunEnd;
|
||||
while (currentRunStart < values.size()) {
|
||||
size_t currentRunEnd = findEqualRunEnd(currentRunStart);
|
||||
if (currentRunEnd - currentRunStart != compression.repeatCount)
|
||||
break;
|
||||
if (static_cast<int64_t>(values[currentRunStart]) != static_cast<int64_t>(lastValue) + step)
|
||||
break;
|
||||
lastValue = values[currentRunStart];
|
||||
progressionEnd = currentRunEnd;
|
||||
currentRunStart = currentRunEnd;
|
||||
}
|
||||
}
|
||||
else {
|
||||
step = 0;
|
||||
}
|
||||
}
|
||||
|
||||
compression.covered = 1;
|
||||
if (progressionEnd > firstRunEnd) {
|
||||
size_t progressionValueCount = (progressionEnd - start) / compression.repeatCount;
|
||||
if (progressionValueCount >= 3) {
|
||||
compression.kind = FlatCompression::Kind::Progression;
|
||||
compression.covered = progressionEnd - start;
|
||||
compression.progressionValueCount = progressionValueCount;
|
||||
compression.step = step;
|
||||
compression.lastValue = lastValue;
|
||||
return compression;
|
||||
}
|
||||
}
|
||||
|
||||
if (compression.repeatCount > 1) {
|
||||
compression.kind = FlatCompression::Kind::EqualRun;
|
||||
compression.covered = compression.repeatCount;
|
||||
return compression;
|
||||
}
|
||||
|
||||
return compression;
|
||||
};
|
||||
|
||||
auto findRepeatedSublist = [&](size_t start) {
|
||||
size_t bestLength = 0;
|
||||
size_t bestRepeatCount = 1;
|
||||
size_t remaining = values.size() - start;
|
||||
|
||||
for (size_t length = 2; length * 2 <= remaining; ++length) {
|
||||
size_t repeatCount = 1;
|
||||
ArrayRef<IntT> candidate = values.slice(start, length);
|
||||
while (start + (repeatCount + 1) * length <= values.size()
|
||||
&& llvm::equal(candidate, values.slice(start + repeatCount * length, length))) {
|
||||
++repeatCount;
|
||||
}
|
||||
|
||||
if (repeatCount <= 1)
|
||||
continue;
|
||||
|
||||
size_t covered = length * repeatCount;
|
||||
size_t bestCovered = bestLength * bestRepeatCount;
|
||||
if (covered > bestCovered || (covered == bestCovered && length < bestLength)) {
|
||||
bestLength = length;
|
||||
bestRepeatCount = repeatCount;
|
||||
}
|
||||
}
|
||||
|
||||
return std::pair(bestLength, bestRepeatCount);
|
||||
};
|
||||
|
||||
for (size_t index = 0; index < values.size();) {
|
||||
if (index != 0)
|
||||
stream << ", ";
|
||||
|
||||
FlatCompression flat = computeFlatCompression(index);
|
||||
auto [sublistLength, sublistRepeatCount] = findRepeatedSublist(index);
|
||||
size_t repeatedSublistCoverage = sublistLength * sublistRepeatCount;
|
||||
if (sublistRepeatCount > 1 && sublistLength > 1 && repeatedSublistCoverage > flat.covered) {
|
||||
printOpenDelimiter(stream, ListDelimiter::Paren);
|
||||
printCompressedIntegerEntries(stream, values.slice(index, sublistLength));
|
||||
printCloseDelimiter(stream, ListDelimiter::Paren);
|
||||
stream << " x" << sublistRepeatCount;
|
||||
index += repeatedSublistCoverage;
|
||||
continue;
|
||||
}
|
||||
|
||||
switch (flat.kind) {
|
||||
case FlatCompression::Kind::Progression:
|
||||
stream << flat.firstValue << " to " << flat.lastValue;
|
||||
if (flat.step != 1)
|
||||
stream << " by " << flat.step;
|
||||
if (flat.repeatCount > 1)
|
||||
stream << " x" << flat.repeatCount;
|
||||
index += flat.covered;
|
||||
break;
|
||||
case FlatCompression::Kind::EqualRun:
|
||||
stream << flat.firstValue << " x" << flat.repeatCount;
|
||||
index += flat.covered;
|
||||
break;
|
||||
case FlatCompression::Kind::Single:
|
||||
stream << flat.firstValue;
|
||||
index += flat.covered;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <typename StreamT, typename IntT>
|
||||
inline void printCompressedIntegerSequence(StreamT& stream, ArrayRef<IntT> values, ListDelimiter delimiter) {
|
||||
printOpenDelimiter(stream, delimiter);
|
||||
printCompressedIntegerEntries(stream, values);
|
||||
printCloseDelimiter(stream, delimiter);
|
||||
}
|
||||
|
||||
template <typename IntT>
|
||||
inline ParseResult parseCompressedIntegerList(OpAsmParser& parser, SmallVectorImpl<IntT>& values) {
|
||||
return parseCompressedIntegerSequence(parser, ListDelimiter::Square, values);
|
||||
}
|
||||
|
||||
template <typename IntT>
|
||||
inline void printCompressedIntegerList(OpAsmPrinter& printer, ArrayRef<IntT> values) {
|
||||
printCompressedIntegerSequence(printer, values, ListDelimiter::Square);
|
||||
}
|
||||
|
||||
inline void printCompressedValueSequence(OpAsmPrinter& printer, ValueRange values) {
|
||||
for (size_t index = 0; index < values.size();) {
|
||||
size_t equalRunEnd = index + 1;
|
||||
while (equalRunEnd < values.size() && values[equalRunEnd] == values[index])
|
||||
++equalRunEnd;
|
||||
|
||||
if (index != 0)
|
||||
printer << ", ";
|
||||
if (equalRunEnd - index > 1) {
|
||||
printer.printOperand(values[index]);
|
||||
printer << " x" << (equalRunEnd - index);
|
||||
index = equalRunEnd;
|
||||
continue;
|
||||
}
|
||||
|
||||
size_t rangeEnd = index + 1;
|
||||
if (auto firstResult = dyn_cast<OpResult>(values[index])) {
|
||||
while (rangeEnd < values.size()) {
|
||||
auto nextResult = dyn_cast<OpResult>(values[rangeEnd]);
|
||||
if (!nextResult || nextResult.getOwner() != firstResult.getOwner()
|
||||
|| nextResult.getResultNumber() != firstResult.getResultNumber() + (rangeEnd - index)) {
|
||||
break;
|
||||
}
|
||||
++rangeEnd;
|
||||
}
|
||||
}
|
||||
else if (auto firstArg = dyn_cast<BlockArgument>(values[index])) {
|
||||
while (rangeEnd < values.size()) {
|
||||
auto nextArg = dyn_cast<BlockArgument>(values[rangeEnd]);
|
||||
if (!nextArg || nextArg.getOwner() != firstArg.getOwner()
|
||||
|| nextArg.getArgNumber() != firstArg.getArgNumber() + (rangeEnd - index)) {
|
||||
break;
|
||||
}
|
||||
++rangeEnd;
|
||||
}
|
||||
}
|
||||
|
||||
printer.printOperand(values[index]);
|
||||
if (rangeEnd - index >= 3) {
|
||||
printer << " to ";
|
||||
printer.printOperand(values[rangeEnd - 1]);
|
||||
}
|
||||
else if (rangeEnd - index == 2) {
|
||||
printer << ", ";
|
||||
printer.printOperand(values[index + 1]);
|
||||
}
|
||||
index = rangeEnd;
|
||||
}
|
||||
}
|
||||
|
||||
inline void printCompressedTypeSequence(OpAsmPrinter& printer, TypeRange types) {
|
||||
printCompressedEqualRuns(printer, types, [&](Type type) { printer.printType(type); });
|
||||
}
|
||||
|
||||
inline void printCompressedValueList(OpAsmPrinter& printer, ValueRange values, ListDelimiter delimiter) {
|
||||
printOpenDelimiter(printer, delimiter);
|
||||
printCompressedValueSequence(printer, values);
|
||||
printCloseDelimiter(printer, delimiter);
|
||||
}
|
||||
|
||||
inline void printCompressedTypeList(OpAsmPrinter& printer, TypeRange types, ListDelimiter delimiter) {
|
||||
printOpenDelimiter(printer, delimiter);
|
||||
printCompressedTypeSequence(printer, types);
|
||||
printCloseDelimiter(printer, delimiter);
|
||||
}
|
||||
|
||||
inline ParseResult parseCompressedTypeSequence(OpAsmParser& parser, SmallVectorImpl<Type>& types, bool allowEmpty) {
|
||||
Type firstType;
|
||||
OptionalParseResult firstTypeResult = parser.parseOptionalType(firstType);
|
||||
if (!firstTypeResult.has_value()) {
|
||||
if (allowEmpty)
|
||||
return success();
|
||||
return parser.emitError(parser.getCurrentLocation(), "expected type");
|
||||
}
|
||||
if (failed(*firstTypeResult))
|
||||
return failure();
|
||||
|
||||
auto appendType = [&](Type type) -> ParseResult {
|
||||
int64_t repeatCount = 1;
|
||||
if (parseOptionalRepeatCount(parser, repeatCount))
|
||||
return failure();
|
||||
for (int64_t index = 0; index < repeatCount; ++index)
|
||||
types.push_back(type);
|
||||
return success();
|
||||
};
|
||||
|
||||
if (appendType(firstType))
|
||||
return failure();
|
||||
while (succeeded(parser.parseOptionalComma())) {
|
||||
Type nextType;
|
||||
if (parser.parseType(nextType) || appendType(nextType))
|
||||
return failure();
|
||||
}
|
||||
return success();
|
||||
}
|
||||
|
||||
inline ParseResult parseCompressedOperandEntryWithFirst(OpAsmParser& parser,
|
||||
OpAsmParser::UnresolvedOperand firstOperand,
|
||||
SmallVectorImpl<OpAsmParser::UnresolvedOperand>& operands) {
|
||||
if (succeeded(parser.parseOptionalKeyword("to"))) {
|
||||
OpAsmParser::UnresolvedOperand lastOperand;
|
||||
if (parser.parseOperand(lastOperand))
|
||||
return failure();
|
||||
if (firstOperand.name == lastOperand.name && firstOperand.number <= lastOperand.number) {
|
||||
for (unsigned number = firstOperand.number; number <= lastOperand.number; ++number)
|
||||
operands.push_back({firstOperand.location, firstOperand.name, number});
|
||||
return success();
|
||||
}
|
||||
|
||||
auto range = getNumericSuffixRange(firstOperand.name, lastOperand.name);
|
||||
if (!range || firstOperand.number != 0 || lastOperand.number != 0)
|
||||
return parser.emitError(parser.getCurrentLocation(), "invalid operand range");
|
||||
for (unsigned number = range->first; number <= range->last; ++number)
|
||||
operands.push_back({firstOperand.location, getInternedNumberedName(parser, range->prefix, number), 0});
|
||||
return success();
|
||||
}
|
||||
|
||||
int64_t repeatCount = 1;
|
||||
if (parseOptionalRepeatCount(parser, repeatCount))
|
||||
return failure();
|
||||
for (int64_t index = 0; index < repeatCount; ++index)
|
||||
operands.push_back(firstOperand);
|
||||
return success();
|
||||
}
|
||||
|
||||
inline ParseResult parseOneCompressedOperandEntry(OpAsmParser& parser,
|
||||
SmallVectorImpl<OpAsmParser::UnresolvedOperand>& operands) {
|
||||
OpAsmParser::UnresolvedOperand firstOperand;
|
||||
if (parser.parseOperand(firstOperand))
|
||||
return failure();
|
||||
return parseCompressedOperandEntryWithFirst(parser, firstOperand, operands);
|
||||
}
|
||||
|
||||
inline ParseResult parseCompressedOperandList(OpAsmParser& parser,
|
||||
ListDelimiter delimiter,
|
||||
SmallVectorImpl<OpAsmParser::UnresolvedOperand>& operands) {
|
||||
if (parseOpenDelimiter(parser, delimiter))
|
||||
return failure();
|
||||
if (succeeded(parseOptionalCloseDelimiter(parser, delimiter)))
|
||||
return success();
|
||||
|
||||
while (true) {
|
||||
if (parseOneCompressedOperandEntry(parser, operands))
|
||||
return failure();
|
||||
if (succeeded(parseOptionalCloseDelimiter(parser, delimiter)))
|
||||
return success();
|
||||
if (parser.parseComma())
|
||||
return failure();
|
||||
}
|
||||
}
|
||||
|
||||
inline ParseResult parseCompressedOperandSequence(OpAsmParser& parser,
|
||||
SmallVectorImpl<OpAsmParser::UnresolvedOperand>& operands) {
|
||||
if (parseOneCompressedOperandEntry(parser, operands))
|
||||
return failure();
|
||||
while (succeeded(parser.parseOptionalComma()))
|
||||
if (parseOneCompressedOperandEntry(parser, operands))
|
||||
return failure();
|
||||
return success();
|
||||
}
|
||||
|
||||
inline ParseResult parseCompressedTypeList(OpAsmParser& parser, ListDelimiter delimiter, SmallVectorImpl<Type>& types) {
|
||||
if (parseOpenDelimiter(parser, delimiter))
|
||||
return failure();
|
||||
if (succeeded(parseOptionalCloseDelimiter(parser, delimiter)))
|
||||
return success();
|
||||
|
||||
if (parseCompressedTypeSequence(parser, types, /*allowEmpty=*/false))
|
||||
return failure();
|
||||
return parseOptionalCloseDelimiter(parser, delimiter);
|
||||
}
|
||||
|
||||
inline bool hasRepeatedTuple(ValueRange values, size_t tupleSize) {
|
||||
if (tupleSize == 0 || values.empty() || values.size() % tupleSize != 0)
|
||||
return false;
|
||||
|
||||
SmallVector<Value> valueVec(values.begin(), values.end());
|
||||
ArrayRef<Value> tuple(valueVec.data(), tupleSize);
|
||||
for (size_t index = tupleSize; index < values.size(); index += tupleSize)
|
||||
if (!llvm::equal(tuple, ArrayRef<Value>(valueVec).slice(index, tupleSize)))
|
||||
return false;
|
||||
return true;
|
||||
}
|
||||
|
||||
inline bool hasRepeatedTuple(TypeRange types, size_t tupleSize) {
|
||||
if (tupleSize == 0 || types.empty() || types.size() % tupleSize != 0)
|
||||
return false;
|
||||
|
||||
SmallVector<Type> typeVec(types.begin(), types.end());
|
||||
ArrayRef<Type> tuple(typeVec.data(), tupleSize);
|
||||
for (size_t index = tupleSize; index < types.size(); index += tupleSize)
|
||||
if (!llvm::equal(tuple, ArrayRef<Type>(typeVec).slice(index, tupleSize)))
|
||||
return false;
|
||||
return true;
|
||||
}
|
||||
|
||||
inline void printValueTupleRun(OpAsmPrinter& printer, ValueRange values, size_t tupleSize, ListDelimiter delimiter) {
|
||||
printOpenDelimiter(printer, delimiter);
|
||||
printOpenDelimiter(printer, ListDelimiter::Paren);
|
||||
for (size_t index = 0; index < tupleSize; ++index) {
|
||||
if (index != 0)
|
||||
printer << ", ";
|
||||
printer.printOperand(values[index]);
|
||||
}
|
||||
printCloseDelimiter(printer, ListDelimiter::Paren);
|
||||
printer << " x" << (values.size() / tupleSize);
|
||||
printCloseDelimiter(printer, delimiter);
|
||||
}
|
||||
|
||||
inline void printTypeTupleRun(OpAsmPrinter& printer, TypeRange types, size_t tupleSize, ListDelimiter delimiter) {
|
||||
printOpenDelimiter(printer, delimiter);
|
||||
printOpenDelimiter(printer, ListDelimiter::Paren);
|
||||
for (size_t index = 0; index < tupleSize; ++index) {
|
||||
if (index != 0)
|
||||
printer << ", ";
|
||||
printer.printType(types[index]);
|
||||
}
|
||||
printCloseDelimiter(printer, ListDelimiter::Paren);
|
||||
printer << " x" << (types.size() / tupleSize);
|
||||
printCloseDelimiter(printer, delimiter);
|
||||
}
|
||||
|
||||
inline ParseResult parseCompressedOrTupleOperandList(OpAsmParser& parser,
|
||||
ListDelimiter delimiter,
|
||||
SmallVectorImpl<OpAsmParser::UnresolvedOperand>& operands) {
|
||||
if (parseOpenDelimiter(parser, delimiter))
|
||||
return failure();
|
||||
if (succeeded(parseOptionalCloseDelimiter(parser, delimiter)))
|
||||
return success();
|
||||
|
||||
if (succeeded(parser.parseOptionalLParen())) {
|
||||
SmallVector<OpAsmParser::UnresolvedOperand> tupleOperands;
|
||||
if (parseCompressedOperandSequence(parser, tupleOperands) || parser.parseRParen())
|
||||
return failure();
|
||||
|
||||
int64_t repeatCount = 1;
|
||||
if (parseOptionalRepeatCount(parser, repeatCount))
|
||||
return failure();
|
||||
for (int64_t repeat = 0; repeat < repeatCount; ++repeat)
|
||||
llvm::append_range(operands, tupleOperands);
|
||||
|
||||
while (succeeded(parser.parseOptionalComma())) {
|
||||
if (parser.parseLParen())
|
||||
return failure();
|
||||
tupleOperands.clear();
|
||||
if (parseCompressedOperandSequence(parser, tupleOperands) || parser.parseRParen())
|
||||
return failure();
|
||||
|
||||
if (parseOptionalRepeatCount(parser, repeatCount))
|
||||
return failure();
|
||||
for (int64_t repeat = 0; repeat < repeatCount; ++repeat)
|
||||
llvm::append_range(operands, tupleOperands);
|
||||
}
|
||||
return parseOptionalCloseDelimiter(parser, delimiter);
|
||||
}
|
||||
|
||||
while (true) {
|
||||
if (parseOneCompressedOperandEntry(parser, operands))
|
||||
return failure();
|
||||
if (succeeded(parseOptionalCloseDelimiter(parser, delimiter)))
|
||||
return success();
|
||||
if (parser.parseComma())
|
||||
return failure();
|
||||
}
|
||||
}
|
||||
|
||||
inline ParseResult
|
||||
parseCompressedOrTupleTypeList(OpAsmParser& parser, ListDelimiter delimiter, SmallVectorImpl<Type>& types) {
|
||||
if (parseOpenDelimiter(parser, delimiter))
|
||||
return failure();
|
||||
if (succeeded(parseOptionalCloseDelimiter(parser, delimiter)))
|
||||
return success();
|
||||
|
||||
if (succeeded(parser.parseOptionalLParen())) {
|
||||
SmallVector<Type> tupleTypes;
|
||||
if (parseCompressedTypeSequence(parser, tupleTypes, /*allowEmpty=*/false) || parser.parseRParen())
|
||||
return failure();
|
||||
|
||||
int64_t repeatCount = 1;
|
||||
if (parseOptionalRepeatCount(parser, repeatCount))
|
||||
return failure();
|
||||
for (int64_t repeat = 0; repeat < repeatCount; ++repeat)
|
||||
llvm::append_range(types, tupleTypes);
|
||||
|
||||
while (succeeded(parser.parseOptionalComma())) {
|
||||
if (parser.parseLParen())
|
||||
return failure();
|
||||
tupleTypes.clear();
|
||||
if (parseCompressedTypeSequence(parser, tupleTypes, /*allowEmpty=*/false) || parser.parseRParen())
|
||||
return failure();
|
||||
|
||||
if (parseOptionalRepeatCount(parser, repeatCount))
|
||||
return failure();
|
||||
for (int64_t repeat = 0; repeat < repeatCount; ++repeat)
|
||||
llvm::append_range(types, tupleTypes);
|
||||
}
|
||||
return parseOptionalCloseDelimiter(parser, delimiter);
|
||||
}
|
||||
|
||||
while (true) {
|
||||
Type type;
|
||||
if (parser.parseType(type))
|
||||
return failure();
|
||||
|
||||
int64_t repeatCount = 1;
|
||||
if (parseOptionalRepeatCount(parser, repeatCount))
|
||||
return failure();
|
||||
for (int64_t repeat = 0; repeat < repeatCount; ++repeat)
|
||||
types.push_back(type);
|
||||
|
||||
if (succeeded(parseOptionalCloseDelimiter(parser, delimiter)))
|
||||
return success();
|
||||
if (parser.parseComma())
|
||||
return failure();
|
||||
}
|
||||
}
|
||||
|
||||
inline void printArgumentBindings(OpAsmPrinter& printer, Block& block, ValueRange operands) {
|
||||
if (block.getNumArguments() == 0) {
|
||||
printer << "() = ()";
|
||||
return;
|
||||
}
|
||||
|
||||
if (block.getNumArguments() == 1) {
|
||||
printer.printOperand(block.getArgument(0));
|
||||
printer << " = ";
|
||||
printCompressedValueList(printer, operands, ListDelimiter::Paren);
|
||||
return;
|
||||
}
|
||||
|
||||
printCompressedValueList(printer, ValueRange(block.getArguments()), ListDelimiter::Paren);
|
||||
printer << " = ";
|
||||
printCompressedValueList(printer, operands, ListDelimiter::Paren);
|
||||
}
|
||||
|
||||
inline ParseResult parseCompressedArgumentEntryWithFirst(OpAsmParser& parser,
|
||||
OpAsmParser::Argument firstArgument,
|
||||
SmallVectorImpl<OpAsmParser::Argument>& arguments) {
|
||||
if (succeeded(parser.parseOptionalKeyword("to"))) {
|
||||
OpAsmParser::Argument lastArgument;
|
||||
if (parser.parseArgument(lastArgument))
|
||||
return failure();
|
||||
if (firstArgument.ssaName.name == lastArgument.ssaName.name
|
||||
&& firstArgument.ssaName.number <= lastArgument.ssaName.number) {
|
||||
for (unsigned number = firstArgument.ssaName.number; number <= lastArgument.ssaName.number; ++number) {
|
||||
OpAsmParser::Argument argument;
|
||||
argument.ssaName = {firstArgument.ssaName.location, firstArgument.ssaName.name, number};
|
||||
arguments.push_back(argument);
|
||||
}
|
||||
return success();
|
||||
}
|
||||
|
||||
auto range = getNumericSuffixRange(firstArgument.ssaName.name, lastArgument.ssaName.name);
|
||||
if (!range || firstArgument.ssaName.number != 0 || lastArgument.ssaName.number != 0)
|
||||
return parser.emitError(parser.getCurrentLocation(), "invalid argument range");
|
||||
for (unsigned number = range->first; number <= range->last; ++number) {
|
||||
OpAsmParser::Argument argument;
|
||||
argument.ssaName = {firstArgument.ssaName.location, getInternedNumberedName(parser, range->prefix, number), 0};
|
||||
arguments.push_back(argument);
|
||||
}
|
||||
return success();
|
||||
}
|
||||
|
||||
arguments.push_back(firstArgument);
|
||||
return success();
|
||||
}
|
||||
|
||||
inline ParseResult parseOneCompressedArgumentEntry(OpAsmParser& parser,
|
||||
SmallVectorImpl<OpAsmParser::Argument>& arguments) {
|
||||
OpAsmParser::Argument firstArgument;
|
||||
if (parser.parseArgument(firstArgument))
|
||||
return failure();
|
||||
return parseCompressedArgumentEntryWithFirst(parser, firstArgument, arguments);
|
||||
}
|
||||
|
||||
inline void applyArgumentTypes(ArrayRef<Type> inputTypes, SmallVectorImpl<OpAsmParser::Argument>& arguments) {
|
||||
for (auto [argument, inputType] : llvm::zip_equal(arguments, inputTypes))
|
||||
argument.type = inputType;
|
||||
}
|
||||
|
||||
inline ParseResult parseArgumentBindings(OpAsmParser& parser,
|
||||
SmallVectorImpl<OpAsmParser::Argument>& arguments,
|
||||
SmallVectorImpl<OpAsmParser::UnresolvedOperand>& operands) {
|
||||
if (succeeded(parser.parseOptionalLParen())) {
|
||||
if (succeeded(parser.parseOptionalRParen())) {
|
||||
if (parser.parseEqual() || parseCompressedOperandList(parser, ListDelimiter::Paren, operands))
|
||||
return failure();
|
||||
return success();
|
||||
}
|
||||
|
||||
OpAsmParser::Argument firstArgument;
|
||||
if (parser.parseArgument(firstArgument) || parseCompressedArgumentEntryWithFirst(parser, firstArgument, arguments))
|
||||
return failure();
|
||||
while (succeeded(parser.parseOptionalComma()))
|
||||
if (parseOneCompressedArgumentEntry(parser, arguments))
|
||||
return failure();
|
||||
if (parser.parseRParen() || parser.parseEqual()
|
||||
|| parseCompressedOperandList(parser, ListDelimiter::Paren, operands)) {
|
||||
return failure();
|
||||
}
|
||||
return success();
|
||||
}
|
||||
|
||||
OpAsmParser::Argument argument;
|
||||
if (parser.parseArgument(argument) || parser.parseEqual()
|
||||
|| parseCompressedOperandList(parser, ListDelimiter::Paren, operands)) {
|
||||
return failure();
|
||||
}
|
||||
arguments.push_back(argument);
|
||||
return success();
|
||||
}
|
||||
|
||||
} // namespace compact_asm
|
||||
} // namespace onnx_mlir
|
||||
|
||||
#endif
|
||||
@@ -0,0 +1,189 @@
|
||||
#include "mlir/Dialect/Arith/IR/Arith.h"
|
||||
#include "mlir/Dialect/Func/IR/FuncOps.h"
|
||||
#include "mlir/IR/Builders.h"
|
||||
#include "mlir/IR/BuiltinOps.h"
|
||||
#include "mlir/IR/Dialect.h"
|
||||
|
||||
#include "ConstantUtils.hpp"
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
ConstantPool::ConstantPool(Operation *constantAnchor, OpBuilder &builder)
|
||||
: anchor(constantAnchor), block(getConstantInsertionBlock(constantAnchor)),
|
||||
builder(builder) {
|
||||
for (Operation &op : *block)
|
||||
if (auto constant = dyn_cast<arith::ConstantOp>(&op))
|
||||
cache.try_emplace(
|
||||
std::make_pair(constant.getType(), constant.getValue()),
|
||||
constant.getResult());
|
||||
}
|
||||
|
||||
Value ConstantPool::getIndex(int64_t value) {
|
||||
return get(builder.getIndexType(), builder.getIndexAttr(value));
|
||||
}
|
||||
|
||||
Value ConstantPool::get(Type type, Attribute value) {
|
||||
auto key = std::make_pair(type, value);
|
||||
if (Value existing = cache.lookup(key))
|
||||
return existing;
|
||||
OpBuilder::InsertionGuard guard(builder);
|
||||
builder.setInsertionPointToStart(block);
|
||||
Value constant = arith::ConstantOp::create(
|
||||
builder, anchor->getLoc(), type, cast<TypedAttr>(value)).getResult();
|
||||
cache.try_emplace(key, constant);
|
||||
return constant;
|
||||
}
|
||||
|
||||
static std::optional<int64_t> getIndexConstantValue(arith::ConstantOp constantOp) {
|
||||
if (!constantOp.getType().isIndex())
|
||||
return std::nullopt;
|
||||
|
||||
auto intAttr = dyn_cast<IntegerAttr>(constantOp.getValue());
|
||||
if (!intAttr || !intAttr.getType().isIndex())
|
||||
return std::nullopt;
|
||||
|
||||
return intAttr.getInt();
|
||||
}
|
||||
|
||||
Block* getConstantInsertionBlock(Operation* anchorOp) {
|
||||
assert(anchorOp && "expected a valid anchor operation");
|
||||
|
||||
if (auto funcOp = dyn_cast<func::FuncOp>(anchorOp))
|
||||
return &funcOp.getBody().front();
|
||||
if (auto funcOp = anchorOp->getParentOfType<func::FuncOp>())
|
||||
return &funcOp.getBody().front();
|
||||
if (auto moduleOp = dyn_cast<ModuleOp>(anchorOp))
|
||||
return moduleOp.getBody();
|
||||
if (auto moduleOp = anchorOp->getParentOfType<ModuleOp>())
|
||||
return moduleOp.getBody();
|
||||
return anchorOp->getBlock();
|
||||
}
|
||||
|
||||
Value getOrCreateConstant(OperationFolder& folder, Operation* anchorOp, Attribute value, Type type) {
|
||||
assert(anchorOp && "expected a valid anchor operation");
|
||||
Block* hostBlock = getConstantInsertionBlock(anchorOp);
|
||||
for (Operation& op : *hostBlock) {
|
||||
auto constantOp = dyn_cast<arith::ConstantOp>(&op);
|
||||
if (!constantOp || constantOp.getType() != type || constantOp.getValue() != value)
|
||||
continue;
|
||||
return constantOp.getResult();
|
||||
}
|
||||
|
||||
auto* arithDialect = anchorOp->getContext()->getOrLoadDialect<arith::ArithDialect>();
|
||||
return folder.getOrCreateConstant(hostBlock, arithDialect, value, type);
|
||||
}
|
||||
|
||||
Value getOrCreateConstant(OpBuilder& builder, Operation* anchorOp, Attribute value, Type type) {
|
||||
assert(anchorOp && "expected a valid anchor operation");
|
||||
Block* hostBlock = getConstantInsertionBlock(anchorOp);
|
||||
for (Operation& op : *hostBlock) {
|
||||
auto constantOp = dyn_cast<arith::ConstantOp>(&op);
|
||||
if (!constantOp || constantOp.getType() != type || constantOp.getValue() != value)
|
||||
continue;
|
||||
return constantOp.getResult();
|
||||
}
|
||||
|
||||
OpBuilder::InsertionGuard guard(builder);
|
||||
builder.setInsertionPointToStart(hostBlock);
|
||||
return arith::ConstantOp::create(builder, anchorOp->getLoc(), type, cast<TypedAttr>(value)).getResult();
|
||||
}
|
||||
|
||||
Value createConstantAtHostBlockStart(OpBuilder& builder, Operation* anchorOp, TypedAttr value) {
|
||||
assert(anchorOp && "expected a valid anchor operation");
|
||||
OpBuilder::InsertionGuard guard(builder);
|
||||
builder.setInsertionPointToStart(getConstantInsertionBlock(anchorOp));
|
||||
return arith::ConstantOp::create(builder, anchorOp->getLoc(), value).getResult();
|
||||
}
|
||||
|
||||
Value getOrCreateConstantLike(OperationFolder& folder, arith::ConstantOp constantOp) {
|
||||
return getOrCreateConstant(folder, constantOp.getOperation(), constantOp.getValue(), constantOp.getType());
|
||||
}
|
||||
|
||||
Value getOrCreateIndexConstant(OperationFolder& folder, Operation* anchorOp, int64_t value) {
|
||||
Builder builder(anchorOp->getContext());
|
||||
return getOrCreateConstant(folder, anchorOp, builder.getIndexAttr(value), builder.getIndexType());
|
||||
}
|
||||
|
||||
Value getOrCreateIndexConstant(OpBuilder& builder, Operation* anchorOp, int64_t value) {
|
||||
return getOrCreateConstant(builder, anchorOp, builder.getIndexAttr(value), builder.getIndexType());
|
||||
}
|
||||
|
||||
void hoistAndUniquifyIndexConstants(func::FuncOp funcOp, RewriterBase& rewriter) {
|
||||
if (funcOp.getBody().empty())
|
||||
return;
|
||||
|
||||
Block& entryBlock = funcOp.getBody().front();
|
||||
DenseMap<int64_t, Value> canonicalByValue;
|
||||
SmallVector<arith::ConstantOp> constants;
|
||||
|
||||
funcOp.walk([&](arith::ConstantOp constantOp) {
|
||||
if (!getIndexConstantValue(constantOp))
|
||||
return;
|
||||
constants.push_back(constantOp);
|
||||
});
|
||||
|
||||
for (arith::ConstantOp constantOp : constants) {
|
||||
auto value = getIndexConstantValue(constantOp);
|
||||
if (!value || constantOp->getBlock() != &entryBlock)
|
||||
continue;
|
||||
canonicalByValue.try_emplace(*value, constantOp.getResult());
|
||||
}
|
||||
|
||||
for (arith::ConstantOp constantOp : constants) {
|
||||
auto value = getIndexConstantValue(constantOp);
|
||||
if (!value)
|
||||
continue;
|
||||
|
||||
Value canonical = canonicalByValue.lookup(*value);
|
||||
if (!canonical) {
|
||||
OpBuilder::InsertionGuard guard(rewriter);
|
||||
rewriter.setInsertionPointToStart(&entryBlock);
|
||||
Builder builder(funcOp.getContext());
|
||||
canonical =
|
||||
arith::ConstantOp::create(rewriter, constantOp.getLoc(), builder.getIndexType(), builder.getIndexAttr(*value));
|
||||
canonicalByValue[*value] = canonical;
|
||||
}
|
||||
|
||||
if (constantOp.getResult() == canonical)
|
||||
continue;
|
||||
|
||||
constantOp.getResult().replaceAllUsesWith(canonical);
|
||||
}
|
||||
|
||||
for (arith::ConstantOp constantOp : llvm::reverse(constants)) {
|
||||
auto value = getIndexConstantValue(constantOp);
|
||||
if (!value)
|
||||
continue;
|
||||
if (constantOp.getResult() == canonicalByValue.lookup(*value))
|
||||
continue;
|
||||
if (constantOp.use_empty())
|
||||
rewriter.eraseOp(constantOp);
|
||||
}
|
||||
}
|
||||
|
||||
std::optional<int64_t> matchConstantIndexValue(Value value) {
|
||||
if (!value || !value.getType().isIndex())
|
||||
return std::nullopt;
|
||||
|
||||
if (auto constant = value.getDefiningOp<arith::ConstantIndexOp>())
|
||||
return constant.value();
|
||||
|
||||
if (auto constant = value.getDefiningOp<arith::ConstantOp>())
|
||||
if (auto intAttr = dyn_cast<IntegerAttr>(constant.getValue()); intAttr && intAttr.getType().isIndex())
|
||||
return intAttr.getInt();
|
||||
|
||||
return std::nullopt;
|
||||
}
|
||||
|
||||
std::optional<int64_t> matchConstantIndexValue(OpFoldResult value) {
|
||||
if (auto attr = dyn_cast<Attribute>(value))
|
||||
if (auto intAttr = dyn_cast<IntegerAttr>(attr); intAttr && intAttr.getType().isIndex())
|
||||
return intAttr.getInt();
|
||||
if (auto operand = dyn_cast<Value>(value))
|
||||
return matchConstantIndexValue(operand);
|
||||
return std::nullopt;
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,52 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/Dialect/Arith/IR/Arith.h"
|
||||
#include "mlir/Dialect/Func/IR/FuncOps.h"
|
||||
#include "mlir/IR/PatternMatch.h"
|
||||
#include "mlir/IR/Value.h"
|
||||
#include "mlir/Transforms/FoldUtils.h"
|
||||
|
||||
#include "llvm/ADT/DenseMap.h"
|
||||
|
||||
#include <optional>
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
class ConstantPool {
|
||||
public:
|
||||
ConstantPool(mlir::Operation *constantAnchor, mlir::OpBuilder &builder);
|
||||
|
||||
mlir::Value getIndex(int64_t value);
|
||||
mlir::Value get(mlir::Type type, mlir::Attribute value);
|
||||
|
||||
private:
|
||||
mlir::Operation *anchor;
|
||||
mlir::Block *block;
|
||||
mlir::OpBuilder &builder;
|
||||
llvm::DenseMap<std::pair<mlir::Type, mlir::Attribute>, mlir::Value> cache;
|
||||
};
|
||||
|
||||
mlir::Block* getConstantInsertionBlock(mlir::Operation* anchorOp);
|
||||
|
||||
mlir::Value
|
||||
getOrCreateConstant(mlir::OperationFolder& folder, mlir::Operation* anchorOp, mlir::Attribute value, mlir::Type type);
|
||||
|
||||
mlir::Value
|
||||
getOrCreateConstant(mlir::OpBuilder& builder, mlir::Operation* anchorOp, mlir::Attribute value, mlir::Type type);
|
||||
|
||||
mlir::Value
|
||||
createConstantAtHostBlockStart(mlir::OpBuilder& builder, mlir::Operation* anchorOp, mlir::TypedAttr value);
|
||||
|
||||
mlir::Value getOrCreateConstantLike(mlir::OperationFolder& folder, mlir::arith::ConstantOp constantOp);
|
||||
|
||||
mlir::Value getOrCreateIndexConstant(mlir::OperationFolder& folder, mlir::Operation* anchorOp, int64_t value);
|
||||
|
||||
mlir::Value getOrCreateIndexConstant(mlir::OpBuilder& builder, mlir::Operation* anchorOp, int64_t value);
|
||||
|
||||
void hoistAndUniquifyIndexConstants(mlir::func::FuncOp funcOp, mlir::RewriterBase& rewriter);
|
||||
|
||||
std::optional<int64_t> matchConstantIndexValue(mlir::Value value);
|
||||
|
||||
std::optional<int64_t> matchConstantIndexValue(mlir::OpFoldResult value);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -1,67 +1,211 @@
|
||||
#include "mlir/Dialect/Affine/IR/AffineOps.h"
|
||||
#include "mlir/Dialect/Arith/IR/Arith.h"
|
||||
#include "mlir/Dialect/MemRef/IR/MemRef.h"
|
||||
#include "mlir/Dialect/SCF/IR/SCF.h"
|
||||
|
||||
#include "llvm/ADT/DenseSet.h"
|
||||
#include "llvm/ADT/STLExtras.h"
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/CoreBlockUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
bool isCoreStaticAddressOp(mlir::Operation* op) {
|
||||
return mlir::isa<mlir::arith::ConstantOp,
|
||||
mlir::arith::AddIOp,
|
||||
mlir::arith::SubIOp,
|
||||
mlir::arith::MulIOp,
|
||||
mlir::arith::DivUIOp,
|
||||
mlir::arith::RemUIOp,
|
||||
mlir::arith::IndexCastOp,
|
||||
mlir::memref::AllocOp,
|
||||
mlir::memref::SubViewOp,
|
||||
mlir::memref::CastOp,
|
||||
mlir::memref::CollapseShapeOp,
|
||||
mlir::memref::ExpandShapeOp>(op);
|
||||
if (mlir::isa<mlir::affine::AffineApplyOp,
|
||||
mlir::arith::ConstantOp,
|
||||
mlir::arith::AddIOp,
|
||||
mlir::arith::SubIOp,
|
||||
mlir::arith::MulIOp,
|
||||
mlir::arith::DivUIOp,
|
||||
mlir::arith::DivSIOp,
|
||||
mlir::arith::MinUIOp,
|
||||
mlir::arith::RemUIOp,
|
||||
mlir::arith::RemSIOp,
|
||||
mlir::arith::IndexCastOp,
|
||||
mlir::arith::CmpIOp,
|
||||
mlir::memref::AllocOp,
|
||||
mlir::memref::SubViewOp,
|
||||
mlir::memref::CastOp,
|
||||
mlir::memref::CollapseShapeOp,
|
||||
mlir::memref::ExpandShapeOp>(op))
|
||||
return true;
|
||||
auto selectOp = mlir::dyn_cast<mlir::arith::SelectOp>(op);
|
||||
return selectOp && selectOp.getType().isIntOrIndex();
|
||||
}
|
||||
|
||||
mlir::LogicalResult
|
||||
walkPimCoreBlock(mlir::Block& block,
|
||||
const StaticValueKnowledge& knowledge,
|
||||
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)> callback) {
|
||||
namespace {
|
||||
|
||||
enum class CoreWalkMode {
|
||||
ExecuteCommunication,
|
||||
StructuralExtremes
|
||||
};
|
||||
using CoreWalkCallback = llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)>;
|
||||
|
||||
static void propagateRegionResults(mlir::ValueRange results, mlir::Region& region, StaticValueKnowledge& knowledge) {
|
||||
if (region.empty())
|
||||
return;
|
||||
auto yield = mlir::cast<mlir::scf::YieldOp>(region.front().getTerminator());
|
||||
for (auto [result, yielded] : llvm::zip(results, yield.getOperands()))
|
||||
knowledge.aliases[result] = resolveLoopCarriedAlias(yielded, knowledge);
|
||||
}
|
||||
|
||||
static mlir::LogicalResult walkPimCoreBlockImpl(mlir::Block& block,
|
||||
const StaticValueKnowledge& initialKnowledge,
|
||||
CoreWalkMode mode,
|
||||
const PimCoreCommunicationPlan* communicationPlan,
|
||||
CoreWalkCallback callback) {
|
||||
bool hasFailure = false;
|
||||
for (mlir::Operation& op : block) {
|
||||
StaticValueKnowledge knowledge = initialKnowledge;
|
||||
llvm::StringRef purpose = mode == CoreWalkMode::ExecuteCommunication ? "communication verification" : "verification";
|
||||
llvm::SmallVector<mlir::Operation*, 0> structuralOperations;
|
||||
llvm::ArrayRef<mlir::Operation*> operations;
|
||||
if (communicationPlan) {
|
||||
auto it = communicationPlan->find(&block);
|
||||
if (it != communicationPlan->end())
|
||||
operations = it->second;
|
||||
}
|
||||
else {
|
||||
structuralOperations.reserve(block.getOperations().size());
|
||||
for (mlir::Operation& op : block)
|
||||
structuralOperations.push_back(&op);
|
||||
operations = structuralOperations;
|
||||
}
|
||||
|
||||
for (mlir::Operation* operation : operations) {
|
||||
mlir::Operation& op = *operation;
|
||||
if (mlir::isa<pim::PimHaltOp, mlir::scf::YieldOp>(op) || isCoreStaticAddressOp(&op))
|
||||
continue;
|
||||
|
||||
if (auto loadOp = mlir::dyn_cast<mlir::memref::LoadOp>(op);
|
||||
loadOp && succeeded(resolveIndexValue(loadOp.getResult(), knowledge)))
|
||||
continue;
|
||||
if (auto forOp = mlir::dyn_cast<mlir::scf::ForOp>(op)) {
|
||||
mlir::Block& loopBody = forOp.getRegion().front();
|
||||
auto lowerBound = resolveIndexValue(forOp.getLowerBound(), knowledge);
|
||||
auto upperBound = resolveIndexValue(forOp.getUpperBound(), knowledge);
|
||||
auto lower = resolveIndexValue(forOp.getLowerBound(), knowledge);
|
||||
auto upper = resolveIndexValue(forOp.getUpperBound(), knowledge);
|
||||
auto step = resolveIndexValue(forOp.getStep(), knowledge);
|
||||
if (failed(lowerBound) || failed(upperBound) || failed(step) || *step <= 0) {
|
||||
forOp.emitOpError("requires statically evaluable scf.for bounds for PIM codegen");
|
||||
if (failed(lower) || failed(upper) || failed(step)
|
||||
|| (mode == CoreWalkMode::ExecuteCommunication && *step <= 0)) {
|
||||
forOp.emitOpError() << "requires statically evaluable scf.for bounds for PIM " << purpose;
|
||||
hasFailure = true;
|
||||
continue;
|
||||
}
|
||||
if (*step <= 0) {
|
||||
forOp.emitOpError("requires positive scf.for step for PIM verification");
|
||||
hasFailure = true;
|
||||
continue;
|
||||
}
|
||||
|
||||
mlir::Block& body = forOp.getRegion().front();
|
||||
llvm::SmallVector<mlir::Value> iterValues(forOp.getInitArgs().begin(), forOp.getInitArgs().end());
|
||||
for (int64_t inductionValue = *lowerBound; inductionValue < *upperBound; inductionValue += *step) {
|
||||
auto visitIteration = [&](int64_t induction, bool carryValues) {
|
||||
StaticValueKnowledge loopKnowledge = knowledge;
|
||||
loopKnowledge.indexValues[forOp.getInductionVar()] = inductionValue;
|
||||
for (auto [iterArg, iterValue] : llvm::zip_equal(forOp.getRegionIterArgs(), iterValues))
|
||||
loopKnowledge.aliases[iterArg] = iterValue;
|
||||
loopKnowledge.indexValues[forOp.getInductionVar()] = induction;
|
||||
for (auto [index, iterArg] : llvm::enumerate(forOp.getRegionIterArgs()))
|
||||
loopKnowledge.aliases[iterArg] = carryValues ? iterValues[index] : forOp.getInitArgs()[index];
|
||||
hasFailure |= failed(walkPimCoreBlockImpl(body, loopKnowledge, mode, communicationPlan, callback));
|
||||
auto yield = mlir::cast<mlir::scf::YieldOp>(body.getTerminator());
|
||||
for (auto [index, yielded] : llvm::enumerate(yield.getOperands()))
|
||||
iterValues[index] = resolveLoopCarriedAlias(yielded, loopKnowledge);
|
||||
};
|
||||
|
||||
if (failed(walkPimCoreBlock(loopBody, loopKnowledge, callback)))
|
||||
hasFailure = true;
|
||||
|
||||
auto yieldOp = mlir::cast<mlir::scf::YieldOp>(loopBody.getTerminator());
|
||||
for (auto [index, yieldedValue] : llvm::enumerate(yieldOp.getOperands()))
|
||||
iterValues[index] = resolveLoopCarriedAlias(yieldedValue, loopKnowledge);
|
||||
if (mode == CoreWalkMode::ExecuteCommunication) {
|
||||
for (int64_t induction = *lower; induction < *upper; induction += *step)
|
||||
visitIteration(induction, true);
|
||||
}
|
||||
else if (*lower < *upper) {
|
||||
visitIteration(*lower, false);
|
||||
int64_t last = *lower + ((*upper - 1 - *lower) / *step) * *step;
|
||||
if (last != *lower)
|
||||
visitIteration(last, false);
|
||||
}
|
||||
for (auto [result, value] : llvm::zip(forOp.getResults(), iterValues))
|
||||
knowledge.aliases[result] = value;
|
||||
continue;
|
||||
}
|
||||
|
||||
if (failed(callback(op, knowledge)))
|
||||
hasFailure = true;
|
||||
if (auto ifOp = mlir::dyn_cast<mlir::scf::IfOp>(op)) {
|
||||
auto condition = resolveIndexValue(ifOp.getCondition(), knowledge);
|
||||
if (failed(condition)) {
|
||||
ifOp.emitOpError() << "requires statically evaluable scf.if condition for PIM " << purpose;
|
||||
hasFailure = true;
|
||||
continue;
|
||||
}
|
||||
mlir::Region& selected = *condition != 0 ? ifOp.getThenRegion() : ifOp.getElseRegion();
|
||||
if (mode == CoreWalkMode::ExecuteCommunication) {
|
||||
hasFailure |= !selected.empty()
|
||||
&& failed(walkPimCoreBlockImpl(selected.front(), knowledge, mode, communicationPlan, callback));
|
||||
}
|
||||
else {
|
||||
for (mlir::Region* region : {&ifOp.getThenRegion(), &ifOp.getElseRegion()})
|
||||
hasFailure |= !region->empty()
|
||||
&& failed(walkPimCoreBlockImpl(region->front(), knowledge, mode, communicationPlan, callback));
|
||||
}
|
||||
propagateRegionResults(ifOp.getResults(), selected, knowledge);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto switchOp = mlir::dyn_cast<mlir::scf::IndexSwitchOp>(op)) {
|
||||
auto selector = resolveIndexValue(switchOp.getArg(), knowledge);
|
||||
if (failed(selector)) {
|
||||
switchOp.emitOpError() << "requires a statically evaluable scf.index_switch selector for PIM " << purpose;
|
||||
hasFailure = true;
|
||||
continue;
|
||||
}
|
||||
mlir::Region* selected = &switchOp.getDefaultRegion();
|
||||
for (auto [caseValue, caseRegion] : llvm::zip(switchOp.getCases(), switchOp.getCaseRegions()))
|
||||
if (caseValue == *selector) {
|
||||
selected = &caseRegion;
|
||||
break;
|
||||
}
|
||||
for (mlir::Region& region : switchOp->getRegions()) {
|
||||
if (mode == CoreWalkMode::ExecuteCommunication && ®ion != selected)
|
||||
continue;
|
||||
hasFailure |= failed(walkPimCoreBlockImpl(region.front(), knowledge, mode, communicationPlan, callback));
|
||||
}
|
||||
propagateRegionResults(switchOp.getResults(), *selected, knowledge);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (mode != CoreWalkMode::ExecuteCommunication || mlir::isa<pim::PimSendOp, pim::PimReceiveOp>(op))
|
||||
hasFailure |= failed(callback(op, knowledge));
|
||||
}
|
||||
return mlir::success(!hasFailure);
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
PimCoreCommunicationPlan buildPimCoreCommunicationPlan(mlir::Block& block) {
|
||||
llvm::DenseSet<mlir::Operation*> communicationAncestors;
|
||||
block.walk([&](mlir::Operation* op) {
|
||||
if (!mlir::isa<pim::PimSendOp, pim::PimReceiveOp>(op))
|
||||
return;
|
||||
for (mlir::Operation* parent = op->getParentOp(); parent; parent = parent->getParentOp())
|
||||
communicationAncestors.insert(parent);
|
||||
});
|
||||
|
||||
PimCoreCommunicationPlan plan;
|
||||
block.walk([&](mlir::Operation* op) {
|
||||
bool isCommunication = mlir::isa<pim::PimSendOp, pim::PimReceiveOp>(op);
|
||||
bool isControlFlow = mlir::isa<mlir::scf::ForOp, mlir::scf::IfOp, mlir::scf::IndexSwitchOp>(op);
|
||||
if (isCommunication || (isControlFlow && (op->getNumResults() != 0 || communicationAncestors.contains(op))))
|
||||
plan[op->getBlock()].push_back(op);
|
||||
});
|
||||
return plan;
|
||||
}
|
||||
|
||||
mlir::LogicalResult walkPimCoreCommunicationBlock(
|
||||
mlir::Block& block,
|
||||
const PimCoreCommunicationPlan& plan,
|
||||
const StaticValueKnowledge& knowledge,
|
||||
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)> callback) {
|
||||
return walkPimCoreBlockImpl(block, knowledge, CoreWalkMode::ExecuteCommunication, &plan, callback);
|
||||
}
|
||||
|
||||
mlir::LogicalResult walkPimCoreBlockStructurally(
|
||||
mlir::Block& block,
|
||||
const StaticValueKnowledge& knowledge,
|
||||
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)> callback) {
|
||||
return walkPimCoreBlockImpl(block, knowledge, CoreWalkMode::StructuralExtremes, nullptr, callback);
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -3,22 +3,37 @@
|
||||
#include "mlir/IR/Block.h"
|
||||
#include "mlir/Support/LogicalResult.h"
|
||||
|
||||
#include "llvm/ADT/DenseMap.h"
|
||||
#include "llvm/ADT/STLFunctionalExtras.h"
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/AddressAnalysis.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
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
|
||||
/// address or index computation and therefore do not emit PIM instructions.
|
||||
bool isCoreStaticAddressOp(mlir::Operation* op);
|
||||
|
||||
/// Walks a `pim.core` body, statically unrolling nested `scf.for` loops when
|
||||
/// their bounds are known and invoking `callback` only on instruction-emitting
|
||||
/// operations.
|
||||
mlir::LogicalResult
|
||||
walkPimCoreBlock(mlir::Block& block,
|
||||
const StaticValueKnowledge& knowledge,
|
||||
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)> callback);
|
||||
/// Walks a `pim.core` body's communication stream, statically unrolling
|
||||
/// control flow that contains send/receive operations and invoking `callback`
|
||||
/// on those operations in execution order.
|
||||
PimCoreCommunicationPlan buildPimCoreCommunicationPlan(mlir::Block& block);
|
||||
|
||||
mlir::LogicalResult walkPimCoreCommunicationBlock(
|
||||
mlir::Block& block,
|
||||
const PimCoreCommunicationPlan& plan,
|
||||
const StaticValueKnowledge& knowledge,
|
||||
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)> callback);
|
||||
|
||||
/// Walks a `pim.core`-like body structurally for verification without
|
||||
/// enumerating full loop trip counts. Loop bounds must still be statically
|
||||
/// evaluable so address resolution remains well-defined.
|
||||
mlir::LogicalResult walkPimCoreBlockStructurally(
|
||||
mlir::Block& block,
|
||||
const StaticValueKnowledge& knowledge,
|
||||
llvm::function_ref<mlir::LogicalResult(mlir::Operation&, const StaticValueKnowledge&)> callback);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
#include <algorithm>
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/IndexingUtils.hpp"
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
int64_t normalizeAxis(int64_t axis, int64_t rank) { return axis >= 0 ? axis : rank + axis; }
|
||||
|
||||
FailureOr<int64_t> normalizeAxisChecked(int64_t axis, int64_t rank) {
|
||||
int64_t normalizedAxis = normalizeAxis(axis, rank);
|
||||
if (normalizedAxis < 0 || normalizedAxis >= rank)
|
||||
return failure();
|
||||
return normalizedAxis;
|
||||
}
|
||||
|
||||
int64_t normalizeIndex(int64_t index, int64_t dimSize) { return index >= 0 ? index : dimSize + index; }
|
||||
|
||||
static SmallVector<int64_t> normalizeAxesImpl(std::optional<ArrayAttr> axesAttr, int64_t rank) {
|
||||
SmallVector<int64_t> normalizedAxes;
|
||||
if (!axesAttr) {
|
||||
normalizedAxes.reserve(rank);
|
||||
for (int64_t axis = 0; axis < rank; ++axis)
|
||||
normalizedAxes.push_back(axis);
|
||||
}
|
||||
else {
|
||||
normalizedAxes.reserve(axesAttr->size());
|
||||
for (Attribute attr : *axesAttr)
|
||||
normalizedAxes.push_back(normalizeAxis(cast<IntegerAttr>(attr).getInt(), rank));
|
||||
llvm::sort(normalizedAxes);
|
||||
normalizedAxes.erase(std::unique(normalizedAxes.begin(), normalizedAxes.end()), normalizedAxes.end());
|
||||
}
|
||||
return normalizedAxes;
|
||||
}
|
||||
|
||||
FailureOr<SmallVector<int64_t>> normalizeAxesChecked(std::optional<ArrayAttr> axesAttr, int64_t rank) {
|
||||
SmallVector<int64_t> normalizedAxes = normalizeAxesImpl(axesAttr, rank);
|
||||
for (int64_t axis : normalizedAxes)
|
||||
if (axis < 0 || axis >= rank)
|
||||
return failure();
|
||||
return normalizedAxes;
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,20 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/IR/BuiltinAttributes.h"
|
||||
#include "mlir/Support/LogicalResult.h"
|
||||
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
|
||||
#include <optional>
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
int64_t normalizeAxis(int64_t axis, int64_t rank);
|
||||
|
||||
mlir::FailureOr<int64_t> normalizeAxisChecked(int64_t axis, int64_t rank);
|
||||
|
||||
int64_t normalizeIndex(int64_t index, int64_t dimSize);
|
||||
|
||||
mlir::FailureOr<llvm::SmallVector<int64_t>> normalizeAxesChecked(std::optional<mlir::ArrayAttr> axesAttr, int64_t rank);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,96 @@
|
||||
#include "mlir/Dialect/SCF/IR/SCF.h"
|
||||
|
||||
#include "llvm/Support/MathExtras.h"
|
||||
|
||||
#include <optional>
|
||||
|
||||
#include "ConstantUtils.hpp"
|
||||
#include "LoopUtils.hpp"
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
namespace {
|
||||
|
||||
static std::optional<int64_t> getStaticTripCount(Value lowerBound, Value upperBound, Value step) {
|
||||
auto lower = matchConstantIndexValue(lowerBound);
|
||||
auto upper = matchConstantIndexValue(upperBound);
|
||||
auto stepValue = matchConstantIndexValue(step);
|
||||
if (!lower || !upper || !stepValue)
|
||||
return std::nullopt;
|
||||
if (*stepValue <= 0)
|
||||
return std::nullopt;
|
||||
if (*upper <= *lower)
|
||||
return int64_t {0};
|
||||
return llvm::divideCeil(*upper - *lower, *stepValue);
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
static LogicalResult validateNormalizedLoopYields(Location loc, ValueRange initArgs, ArrayRef<Value> yieldedValues) {
|
||||
if (yieldedValues.size() == initArgs.size())
|
||||
return success();
|
||||
|
||||
emitError(loc) << "normalized loop body yielded " << yieldedValues.size() << " values for " << initArgs.size()
|
||||
<< " iter args";
|
||||
return failure();
|
||||
}
|
||||
|
||||
FailureOr<NormalizedLoopResult> buildNormalizedScfFor(OpBuilder& builder,
|
||||
Location loc,
|
||||
Value lowerBound,
|
||||
Value upperBound,
|
||||
Value step,
|
||||
ValueRange initArgs,
|
||||
NormalizedLoopBodyBuilder bodyBuilder) {
|
||||
NormalizedLoopResult result;
|
||||
|
||||
if (auto stepValue = matchConstantIndexValue(step); stepValue && *stepValue <= 0) {
|
||||
emitError(loc) << "normalized scf.for requires a positive step, got " << *stepValue;
|
||||
return failure();
|
||||
}
|
||||
|
||||
if (auto tripCount = getStaticTripCount(lowerBound, upperBound, step)) {
|
||||
if (*tripCount == 0) {
|
||||
llvm::append_range(result.results, initArgs);
|
||||
return result;
|
||||
}
|
||||
|
||||
if (*tripCount == 1) {
|
||||
result.inductionVar = lowerBound;
|
||||
if (failed(bodyBuilder(builder, loc, lowerBound, initArgs, result.results)))
|
||||
return failure();
|
||||
if (failed(validateNormalizedLoopYields(loc, initArgs, result.results)))
|
||||
return failure();
|
||||
return result;
|
||||
}
|
||||
}
|
||||
|
||||
result.loop = scf::ForOp::create(builder, loc, lowerBound, upperBound, step, initArgs);
|
||||
result.inductionVar = result.loop.getInductionVar();
|
||||
|
||||
{
|
||||
OpBuilder::InsertionGuard guard(builder);
|
||||
Block* body = result.loop.getBody();
|
||||
if (!body->empty())
|
||||
if (auto yieldOp = dyn_cast<scf::YieldOp>(body->back()))
|
||||
yieldOp->erase();
|
||||
builder.setInsertionPointToEnd(body);
|
||||
ValueRange iterArgs = result.loop.getRegionIterArgs();
|
||||
if (failed(bodyBuilder(builder, loc, result.inductionVar, iterArgs, result.results))) {
|
||||
result.loop.erase();
|
||||
return failure();
|
||||
}
|
||||
if (failed(validateNormalizedLoopYields(loc, initArgs, result.results))) {
|
||||
result.loop.erase();
|
||||
return failure();
|
||||
}
|
||||
scf::YieldOp::create(builder, loc, result.results);
|
||||
}
|
||||
builder.setInsertionPointAfter(result.loop);
|
||||
result.results.assign(result.loop.getResults().begin(), result.loop.getResults().end());
|
||||
return result;
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,30 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/Dialect/SCF/IR/SCF.h"
|
||||
#include "mlir/IR/Builders.h"
|
||||
|
||||
#include "llvm/ADT/STLFunctionalExtras.h"
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
struct NormalizedLoopResult {
|
||||
mlir::Value inductionVar;
|
||||
llvm::SmallVector<mlir::Value, 4> results;
|
||||
mlir::scf::ForOp loop;
|
||||
|
||||
bool wasInlined() const { return !loop; }
|
||||
};
|
||||
|
||||
using NormalizedLoopBodyBuilder = llvm::function_ref<mlir::LogicalResult(
|
||||
mlir::OpBuilder&, mlir::Location, mlir::Value, mlir::ValueRange, llvm::SmallVectorImpl<mlir::Value>&)>;
|
||||
|
||||
mlir::FailureOr<NormalizedLoopResult> buildNormalizedScfFor(mlir::OpBuilder& builder,
|
||||
mlir::Location loc,
|
||||
mlir::Value lowerBound,
|
||||
mlir::Value upperBound,
|
||||
mlir::Value step,
|
||||
mlir::ValueRange initArgs,
|
||||
NormalizedLoopBodyBuilder bodyBuilder);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -1,4 +1,8 @@
|
||||
#include "llvm/ADT/STLExtras.h"
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
#include "llvm/Support/ErrorHandling.h"
|
||||
|
||||
#include <functional>
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp"
|
||||
|
||||
@@ -35,6 +39,30 @@ int64_t getNumElements(llvm::ArrayRef<int64_t> shape) {
|
||||
return numElements;
|
||||
}
|
||||
|
||||
bool hasByteSizedElementType(mlir::Type elementType) {
|
||||
if (mlir::isa<mlir::IndexType>(elementType))
|
||||
return true;
|
||||
if (auto intType = mlir::dyn_cast<mlir::IntegerType>(elementType))
|
||||
return intType.getWidth() > 0 && intType.getWidth() % 8 == 0;
|
||||
if (auto floatType = mlir::dyn_cast<mlir::FloatType>(elementType))
|
||||
return floatType.getWidth() > 0 && floatType.getWidth() % 8 == 0;
|
||||
return false;
|
||||
}
|
||||
|
||||
size_t getElementTypeSizeInBytes(mlir::Type elementType) {
|
||||
if (mlir::isa<mlir::IndexType>(elementType))
|
||||
return mlir::IndexType::kInternalStorageBitWidth / 8;
|
||||
if (auto intType = mlir::dyn_cast<mlir::IntegerType>(elementType))
|
||||
return static_cast<size_t>(intType.getWidth() / 8);
|
||||
if (auto floatType = mlir::dyn_cast<mlir::FloatType>(elementType))
|
||||
return static_cast<size_t>(floatType.getWidth() / 8);
|
||||
llvm_unreachable("expected byte-sized integer, float, or index element type");
|
||||
}
|
||||
|
||||
size_t getShapedTypeSizeInBytes(mlir::ShapedType shapedType) {
|
||||
return static_cast<size_t>(shapedType.getNumElements()) * getElementTypeSizeInBytes(shapedType.getElementType());
|
||||
}
|
||||
|
||||
bool isMemoryContiguous(llvm::ArrayRef<int64_t> srcShape,
|
||||
llvm::ArrayRef<int64_t> offsets,
|
||||
llvm::ArrayRef<int64_t> sizes,
|
||||
@@ -86,4 +114,132 @@ bool isMemoryContiguous(llvm::ArrayRef<int64_t> srcShape,
|
||||
return true;
|
||||
}
|
||||
|
||||
bool isContiguousSubviewWithDynamicOffsets(llvm::ArrayRef<int64_t> sourceShape,
|
||||
llvm::ArrayRef<mlir::OpFoldResult> mixedOffsets,
|
||||
llvm::ArrayRef<int64_t> staticSizes,
|
||||
llvm::ArrayRef<int64_t> staticStrides) {
|
||||
if (sourceShape.size() != mixedOffsets.size() || sourceShape.size() != staticSizes.size()
|
||||
|| sourceShape.size() != staticStrides.size()) {
|
||||
return false;
|
||||
}
|
||||
|
||||
if (llvm::any_of(staticStrides, [](int64_t stride) { return stride != 1; }))
|
||||
return false;
|
||||
|
||||
auto reversedTriples =
|
||||
llvm::zip_equal(llvm::reverse(sourceShape), llvm::reverse(mixedOffsets), llvm::reverse(staticSizes));
|
||||
|
||||
auto firstNonZeroOrDynamicOffset = llvm::find_if(reversedTriples, [](auto triple) {
|
||||
auto [_sourceDim, offset, _size] = triple;
|
||||
if (auto attr = mlir::dyn_cast<mlir::Attribute>(offset))
|
||||
return mlir::cast<mlir::IntegerAttr>(attr).getInt() != 0;
|
||||
return true;
|
||||
});
|
||||
|
||||
if (firstNonZeroOrDynamicOffset != reversedTriples.end()) {
|
||||
auto [sourceDim, offset, size] = *firstNonZeroOrDynamicOffset;
|
||||
if (auto attr = mlir::dyn_cast<mlir::Attribute>(offset)) {
|
||||
int64_t staticOffset = mlir::cast<mlir::IntegerAttr>(attr).getInt();
|
||||
if (size > sourceDim - staticOffset)
|
||||
return false;
|
||||
}
|
||||
|
||||
++firstNonZeroOrDynamicOffset;
|
||||
for (auto it = firstNonZeroOrDynamicOffset; it != reversedTriples.end(); ++it)
|
||||
if (std::get<2>(*it) != 1)
|
||||
return false;
|
||||
}
|
||||
|
||||
auto reversedSizes = llvm::zip_equal(llvm::reverse(sourceShape), llvm::reverse(staticSizes));
|
||||
auto firstDifferentSize = llvm::find_if(reversedSizes, [](auto pair) {
|
||||
auto [sourceDim, size] = pair;
|
||||
return size != sourceDim;
|
||||
});
|
||||
|
||||
if (firstDifferentSize != reversedSizes.end()) {
|
||||
++firstDifferentSize;
|
||||
for (auto it = firstDifferentSize; it != reversedSizes.end(); ++it)
|
||||
if (std::get<1>(*it) != 1)
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
bool hasStaticPositiveShape(llvm::ArrayRef<int64_t> shape) {
|
||||
return llvm::all_of(shape, [](int64_t dim) { return dim > 0; });
|
||||
}
|
||||
|
||||
bool hasStaticPositiveShape(mlir::RankedTensorType type) {
|
||||
return type.hasStaticShape() && hasStaticPositiveShape(type.getShape());
|
||||
}
|
||||
|
||||
int64_t getStaticShapeElementCount(llvm::ArrayRef<int64_t> shape) {
|
||||
return std::accumulate(shape.begin(), shape.end(), int64_t {1}, std::multiplies<int64_t> {});
|
||||
}
|
||||
|
||||
llvm::SmallVector<int64_t> permuteShape(llvm::ArrayRef<int64_t> shape, llvm::ArrayRef<int64_t> permutation) {
|
||||
llvm::SmallVector<int64_t> permutedShape;
|
||||
permutedShape.reserve(permutation.size());
|
||||
for (int64_t axis : permutation)
|
||||
permutedShape.push_back(shape[axis]);
|
||||
return permutedShape;
|
||||
}
|
||||
|
||||
llvm::SmallVector<int64_t> invertPermutation(llvm::ArrayRef<int64_t> permutation) {
|
||||
llvm::SmallVector<int64_t> inversePermutation(permutation.size());
|
||||
for (auto [newIndex, oldIndex] : llvm::enumerate(permutation))
|
||||
inversePermutation[oldIndex] = static_cast<int64_t>(newIndex);
|
||||
return inversePermutation;
|
||||
}
|
||||
|
||||
mlir::FailureOr<llvm::SmallVector<int64_t>>
|
||||
getTransposePermutationChecked(std::optional<mlir::ArrayAttr> permAttr, int64_t rank) {
|
||||
llvm::SmallVector<int64_t> permutation;
|
||||
if (!permAttr) {
|
||||
permutation.reserve(rank);
|
||||
for (int64_t dim = rank - 1; dim >= 0; --dim)
|
||||
permutation.push_back(dim);
|
||||
return permutation;
|
||||
}
|
||||
|
||||
if (static_cast<int64_t>(permAttr->size()) != rank)
|
||||
return mlir::failure();
|
||||
|
||||
permutation.reserve(permAttr->size());
|
||||
llvm::SmallVector<bool> seen(rank, false);
|
||||
for (mlir::IntegerAttr attr : permAttr->getAsRange<mlir::IntegerAttr>()) {
|
||||
int64_t axis = attr.getInt();
|
||||
if (axis < 0 || axis >= rank || seen[axis])
|
||||
return mlir::failure();
|
||||
seen[axis] = true;
|
||||
permutation.push_back(axis);
|
||||
}
|
||||
return permutation;
|
||||
}
|
||||
|
||||
llvm::SmallVector<mlir::OpFoldResult> getStaticIndexAttrs(mlir::Builder& builder, llvm::ArrayRef<int64_t> values) {
|
||||
llvm::SmallVector<mlir::OpFoldResult> attrs;
|
||||
attrs.reserve(values.size());
|
||||
for (int64_t value : values)
|
||||
attrs.push_back(builder.getIndexAttr(value));
|
||||
return attrs;
|
||||
}
|
||||
|
||||
llvm::SmallVector<mlir::OpFoldResult> getUnitStrides(mlir::PatternRewriter& rewriter, int64_t rank) {
|
||||
return llvm::SmallVector<mlir::OpFoldResult>(rank, rewriter.getIndexAttr(1));
|
||||
}
|
||||
|
||||
llvm::SmallVector<mlir::OpFoldResult> getZeroOffsets(mlir::PatternRewriter& rewriter, int64_t rank) {
|
||||
return llvm::SmallVector<mlir::OpFoldResult>(rank, rewriter.getIndexAttr(0));
|
||||
}
|
||||
|
||||
llvm::SmallVector<mlir::OpFoldResult> getStaticSizes(mlir::PatternRewriter& rewriter, llvm::ArrayRef<int64_t> shape) {
|
||||
llvm::SmallVector<mlir::OpFoldResult> sizes;
|
||||
sizes.reserve(shape.size());
|
||||
for (int64_t dim : shape)
|
||||
sizes.push_back(rewriter.getIndexAttr(dim));
|
||||
return sizes;
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -1,10 +1,24 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/IR/BuiltinTypes.h"
|
||||
#include "mlir/IR/OpDefinition.h"
|
||||
#include "mlir/IR/PatternMatch.h"
|
||||
#include "mlir/IR/Value.h"
|
||||
|
||||
#include "llvm/ADT/ArrayRef.h"
|
||||
#include "llvm/ADT/DenseMap.h"
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
|
||||
#include <cstddef>
|
||||
#include <optional>
|
||||
#include <type_traits>
|
||||
#include <utility>
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
using HSliceId = size_t;
|
||||
using CoreId = size_t;
|
||||
|
||||
llvm::SmallVector<int64_t> computeRowMajorStrides(llvm::ArrayRef<int64_t> shape);
|
||||
|
||||
llvm::SmallVector<int64_t>
|
||||
@@ -14,9 +28,85 @@ int64_t linearizeIndex(llvm::ArrayRef<int64_t> indices, llvm::ArrayRef<int64_t>
|
||||
|
||||
int64_t getNumElements(llvm::ArrayRef<int64_t> shape);
|
||||
|
||||
bool hasByteSizedElementType(mlir::Type elementType);
|
||||
|
||||
size_t getElementTypeSizeInBytes(mlir::Type elementType);
|
||||
|
||||
size_t getShapedTypeSizeInBytes(mlir::ShapedType shapedType);
|
||||
|
||||
bool isMemoryContiguous(llvm::ArrayRef<int64_t> srcShape,
|
||||
llvm::ArrayRef<int64_t> offsets,
|
||||
llvm::ArrayRef<int64_t> sizes,
|
||||
llvm::ArrayRef<int64_t> strides);
|
||||
|
||||
bool isContiguousSubviewWithDynamicOffsets(llvm::ArrayRef<int64_t> sourceShape,
|
||||
llvm::ArrayRef<mlir::OpFoldResult> mixedOffsets,
|
||||
llvm::ArrayRef<int64_t> staticSizes,
|
||||
llvm::ArrayRef<int64_t> staticStrides);
|
||||
|
||||
template <class A, class B, class C = std::common_type_t<A, B>>
|
||||
constexpr C ceilIntegerDivide(A a, B b) {
|
||||
static_assert(std::is_integral_v<A>, "A must be an integer type");
|
||||
static_assert(std::is_integral_v<B>, "B must be an integer type");
|
||||
C ac = static_cast<C>(a);
|
||||
C bc = static_cast<C>(b);
|
||||
return 1 + (ac - 1) / bc;
|
||||
}
|
||||
|
||||
template <class A, class B, class C = std::common_type_t<A, B>>
|
||||
constexpr std::pair<C, C> ceilIntegerDivideWithRemainder(A a, B b) {
|
||||
static_assert(std::is_integral_v<A>, "A must be an integer type");
|
||||
static_assert(std::is_integral_v<B>, "B must be an integer type");
|
||||
C ac = static_cast<C>(a);
|
||||
C bc = static_cast<C>(b);
|
||||
return {ceilIntegerDivide(ac, bc), ac % bc};
|
||||
}
|
||||
|
||||
template <class T>
|
||||
bool isVectorShape(mlir::ArrayRef<T> shape) {
|
||||
return shape.size() == 2 && (shape[0] == 1 || shape[1] == 1);
|
||||
}
|
||||
|
||||
template <class T>
|
||||
bool isMatrixShape(mlir::ArrayRef<T> shape) {
|
||||
return shape.size() == 2;
|
||||
}
|
||||
|
||||
template <class T>
|
||||
bool isHVectorShape(mlir::ArrayRef<T> shape) {
|
||||
return shape.size() == 2 && shape[0] == 1;
|
||||
}
|
||||
|
||||
inline auto getTensorShape(mlir::Value tensor) {
|
||||
return mlir::cast<mlir::RankedTensorType>(tensor.getType()).getShape();
|
||||
}
|
||||
|
||||
inline bool haveSameStaticShape(mlir::Value lhs, mlir::Value rhs) {
|
||||
auto lhsType = mlir::dyn_cast<mlir::RankedTensorType>(lhs.getType());
|
||||
auto rhsType = mlir::dyn_cast<mlir::RankedTensorType>(rhs.getType());
|
||||
return lhsType && rhsType && lhsType.hasStaticShape() && rhsType.hasStaticShape()
|
||||
&& lhsType.getShape() == rhsType.getShape();
|
||||
}
|
||||
|
||||
bool hasStaticPositiveShape(mlir::ArrayRef<int64_t> shape);
|
||||
|
||||
bool hasStaticPositiveShape(mlir::RankedTensorType type);
|
||||
|
||||
int64_t getStaticShapeElementCount(mlir::ArrayRef<int64_t> shape);
|
||||
|
||||
llvm::SmallVector<int64_t> permuteShape(mlir::ArrayRef<int64_t> shape, mlir::ArrayRef<int64_t> permutation);
|
||||
|
||||
llvm::SmallVector<int64_t> invertPermutation(mlir::ArrayRef<int64_t> permutation);
|
||||
|
||||
mlir::FailureOr<llvm::SmallVector<int64_t>> getTransposePermutationChecked(std::optional<mlir::ArrayAttr> permAttr,
|
||||
int64_t rank);
|
||||
|
||||
llvm::SmallVector<mlir::OpFoldResult> getStaticIndexAttrs(mlir::Builder& builder, llvm::ArrayRef<int64_t> values);
|
||||
|
||||
llvm::SmallVector<mlir::OpFoldResult> getUnitStrides(mlir::PatternRewriter& rewriter, int64_t rank);
|
||||
|
||||
llvm::SmallVector<mlir::OpFoldResult> getZeroOffsets(mlir::PatternRewriter& rewriter, int64_t rank);
|
||||
|
||||
llvm::SmallVector<mlir::OpFoldResult> getStaticSizes(mlir::PatternRewriter& rewriter, llvm::ArrayRef<int64_t> shape);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
#include "mlir/Dialect/Linalg/IR/Linalg.h"
|
||||
#include "mlir/Dialect/Tensor/IR/Tensor.h"
|
||||
#include "mlir/Interfaces/SideEffectInterfaces.h"
|
||||
|
||||
#include "ShapingUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
|
||||
#include "src/Dialect/ONNX/ONNXOps.hpp"
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
bool isShapingOnlyOp(Operation *op) {
|
||||
return isa<tensor::CastOp,
|
||||
tensor::CollapseShapeOp,
|
||||
tensor::ExpandShapeOp,
|
||||
tensor::ExtractSliceOp,
|
||||
tensor::InsertSliceOp,
|
||||
tensor::ConcatOp,
|
||||
tensor::EmptyOp,
|
||||
tensor::ExtractOp,
|
||||
tensor::InsertOp,
|
||||
tensor::SplatOp,
|
||||
linalg::TransposeOp,
|
||||
ONNXTransposeOp,
|
||||
spatial::SpatConcatOp,
|
||||
spatial::SpatExtractRowsOp>(op);
|
||||
}
|
||||
|
||||
bool isPureIndexComputationOp(Operation *op) {
|
||||
if (op->getNumRegions() != 0 || op->getNumResults() == 0 || op->hasTrait<OpTrait::IsTerminator>()
|
||||
|| !isMemoryEffectFree(op))
|
||||
return false;
|
||||
auto isIndexOrInteger = [](Type type) { return type.isIndex() || isa<IntegerType>(type); };
|
||||
return llvm::all_of(op->getOperandTypes(), isIndexOrInteger)
|
||||
&& llvm::all_of(op->getResultTypes(), isIndexOrInteger);
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,13 @@
|
||||
#pragma once
|
||||
|
||||
namespace mlir {
|
||||
class Operation;
|
||||
}
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
bool isShapingOnlyOp(mlir::Operation *op);
|
||||
|
||||
bool isPureIndexComputationOp(mlir::Operation *op);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,223 @@
|
||||
#include "StaticIntGrid.hpp"
|
||||
|
||||
#include "mlir/Dialect/Arith/IR/Arith.h"
|
||||
|
||||
#include "AffineUtils.hpp"
|
||||
#include "ConstantUtils.hpp"
|
||||
#include <algorithm>
|
||||
#include <limits>
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
namespace {
|
||||
|
||||
static std::optional<size_t> cellCount(size_t rows, size_t columns) {
|
||||
if (!rows || !columns ||
|
||||
rows > static_cast<size_t>(std::numeric_limits<int64_t>::max()) / columns)
|
||||
return std::nullopt;
|
||||
return rows * columns;
|
||||
}
|
||||
|
||||
static bool checkedFlatIndex(
|
||||
size_t row, size_t columns, size_t column, size_t &flat) {
|
||||
bool overflow;
|
||||
flat = llvm::SaturatingMultiplyAdd(row, columns, column, &overflow);
|
||||
return !overflow &&
|
||||
flat <= static_cast<size_t>(std::numeric_limits<int64_t>::max());
|
||||
}
|
||||
|
||||
static bool affineValue(int64_t base, int64_t rowStep, int64_t columnStep,
|
||||
size_t row, size_t column, int64_t &result) {
|
||||
if (row > static_cast<size_t>(std::numeric_limits<int64_t>::max()) ||
|
||||
column > static_cast<size_t>(std::numeric_limits<int64_t>::max()))
|
||||
return false;
|
||||
int64_t rowValue, columnValue;
|
||||
return !llvm::MulOverflow(rowStep, static_cast<int64_t>(row), rowValue)
|
||||
&& !llvm::MulOverflow(columnStep, static_cast<int64_t>(column),
|
||||
columnValue)
|
||||
&& !llvm::AddOverflow(base, rowValue, result)
|
||||
&& !llvm::AddOverflow(result, columnValue, result);
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
FailureOr<StaticIntGrid> StaticIntGrid::fromSequences(
|
||||
ArrayRef<StaticIntSequence> input, bool columnsInput,
|
||||
int64_t sparseBase) {
|
||||
if (input.empty() || !input.front().size())
|
||||
return failure();
|
||||
size_t rowCount = columnsInput ? input.front().size() : input.size();
|
||||
size_t columnCount = columnsInput ? input.size() : input.front().size();
|
||||
auto cells = cellCount(rowCount, columnCount);
|
||||
if (!cells ||
|
||||
llvm::any_of(input, [&](const StaticIntSequence &sequence) {
|
||||
return sequence.size() != input.front().size();
|
||||
}))
|
||||
return failure();
|
||||
StaticIntGrid result(rowCount, columnCount, input.front().valueAt(0));
|
||||
if (llvm::all_equal(input)) {
|
||||
if (input.front().getKind() == StaticIntSequenceKind::Uniform)
|
||||
return result;
|
||||
result.kind = columnsInput ? Kind::ActionOnly : Kind::LaneOnly;
|
||||
result.values = input.front();
|
||||
return result;
|
||||
}
|
||||
SmallVector<int64_t> outerBases;
|
||||
for (const StaticIntSequence &sequence : input)
|
||||
outerBases.push_back(sequence.valueAt(0));
|
||||
if (llvm::all_of(input, [](const StaticIntSequence &sequence) {
|
||||
return sequence.getKind() == StaticIntSequenceKind::Uniform;
|
||||
})) {
|
||||
result.kind = columnsInput ? Kind::LaneOnly : Kind::ActionOnly;
|
||||
result.values = StaticIntSequence::fromValues(outerBases);
|
||||
return result;
|
||||
}
|
||||
auto innerStep = input.front().getAffineStep();
|
||||
StaticIntSequence bases = StaticIntSequence::fromValues(outerBases);
|
||||
auto outerStep = bases.getAffineStep();
|
||||
if (innerStep && outerStep &&
|
||||
llvm::all_of(input, [&](const StaticIntSequence &sequence) {
|
||||
return sequence.getAffineStep() == innerStep;
|
||||
}))
|
||||
return affine2D(result.base,
|
||||
columnsInput ? *innerStep : *outerStep,
|
||||
columnsInput ? *outerStep : *innerStep, rowCount, columnCount);
|
||||
SmallVector<int64_t> values;
|
||||
values.reserve(*cells);
|
||||
for (size_t row = 0; row < rowCount; ++row)
|
||||
for (size_t column = 0; column < columnCount; ++column)
|
||||
values.push_back(columnsInput ? input[column].valueAt(row)
|
||||
: input[row].valueAt(column));
|
||||
result.values = StaticIntSequence::fromValues(values);
|
||||
for (size_t index = 0; index < *cells; ++index)
|
||||
if (values[index] != sparseBase)
|
||||
result.overrideKeys.push_back(static_cast<int64_t>(index));
|
||||
if (result.overrideKeys.size() <= *cells / 4) {
|
||||
result.kind = Kind::SparseLaneOverrides;
|
||||
result.base = sparseBase;
|
||||
} else {
|
||||
result.kind = Kind::Dense;
|
||||
result.overrideKeys.clear();
|
||||
}
|
||||
return result;
|
||||
}
|
||||
|
||||
FailureOr<StaticIntGrid> StaticIntGrid::fromRows(
|
||||
ArrayRef<StaticIntSequence> rows) {
|
||||
if (rows.empty() || !rows.front().size())
|
||||
return failure();
|
||||
return fromSequences(rows, false, rows.front().valueAt(0));
|
||||
}
|
||||
|
||||
FailureOr<StaticIntGrid> StaticIntGrid::fromColumns(
|
||||
size_t rowCount, ArrayRef<StaticIntSequence> columnSequences,
|
||||
int64_t defaultValue) {
|
||||
if (!cellCount(rowCount, columnSequences.size()))
|
||||
return failure();
|
||||
SmallVector<StaticIntSequence> padded;
|
||||
padded.reserve(columnSequences.size());
|
||||
for (const StaticIntSequence &sequence : columnSequences) {
|
||||
if (sequence.size() > rowCount)
|
||||
return failure();
|
||||
if (sequence.size() == rowCount) {
|
||||
padded.push_back(sequence);
|
||||
continue;
|
||||
}
|
||||
SmallVector<int64_t> values(rowCount, defaultValue);
|
||||
for (size_t row = 0; row < sequence.size(); ++row)
|
||||
values[row] = sequence.valueAt(row);
|
||||
padded.push_back(StaticIntSequence::fromValues(values));
|
||||
}
|
||||
return fromSequences(padded, true, defaultValue);
|
||||
}
|
||||
|
||||
FailureOr<StaticIntGrid> StaticIntGrid::affine2D(
|
||||
int64_t base, int64_t rowStep, int64_t columnStep,
|
||||
size_t rows, size_t columns) {
|
||||
int64_t last;
|
||||
if (!cellCount(rows, columns) ||
|
||||
!affineValue(base, rowStep, columnStep, rows - 1, columns - 1, last))
|
||||
return failure();
|
||||
StaticIntGrid result(rows, columns, base);
|
||||
result.kind = rowStep || columnStep ? Kind::Affine : Kind::Uniform;
|
||||
result.rowStep = rowStep;
|
||||
result.columnStep = columnStep;
|
||||
return result;
|
||||
}
|
||||
|
||||
FailureOr<StaticIntGrid> StaticIntGrid::laneIntervals(
|
||||
size_t columns, ArrayRef<std::pair<size_t, size_t>> intervals,
|
||||
int64_t insideValue, int64_t outsideValue) {
|
||||
if (!columns)
|
||||
return failure();
|
||||
SmallVector<int64_t> values(columns, outsideValue);
|
||||
for (auto [begin, end] : intervals) {
|
||||
if (begin > end || end > columns)
|
||||
return failure();
|
||||
std::fill(values.begin() + begin, values.begin() + end, insideValue);
|
||||
}
|
||||
StaticIntSequence row = StaticIntSequence::fromValues(values);
|
||||
return fromRows(ArrayRef<StaticIntSequence>(row));
|
||||
}
|
||||
|
||||
int64_t StaticIntGrid::valueAt(size_t row, size_t column) const {
|
||||
assert(row < rows && column < columns);
|
||||
if (kind == Kind::Uniform)
|
||||
return base;
|
||||
if (kind == Kind::ActionOnly)
|
||||
return values->valueAt(row);
|
||||
if (kind == Kind::LaneOnly)
|
||||
return values->valueAt(column);
|
||||
if (kind == Kind::Affine) {
|
||||
int64_t result;
|
||||
bool valid = affineValue(base, rowStep, columnStep, row, column, result);
|
||||
assert(valid);
|
||||
return result;
|
||||
}
|
||||
size_t flat;
|
||||
bool valid = checkedFlatIndex(row, columns, column, flat);
|
||||
assert(valid);
|
||||
if (kind == Kind::SparseLaneOverrides) {
|
||||
auto found = llvm::lower_bound(overrideKeys, static_cast<int64_t>(flat));
|
||||
if (found == overrideKeys.end() || *found != static_cast<int64_t>(flat))
|
||||
return base;
|
||||
}
|
||||
assert(kind == Kind::Dense || kind == Kind::SparseLaneOverrides);
|
||||
return values->valueAt(flat);
|
||||
}
|
||||
|
||||
Value StaticIntGrid::emitLookup(Value row, Value column,
|
||||
Operation *constantAnchor,
|
||||
ConstantPool &constants, OpBuilder &builder,
|
||||
Location loc) const {
|
||||
if (kind == Kind::Uniform)
|
||||
return constants.getIndex(base);
|
||||
if (kind == Kind::ActionOnly || kind == Kind::LaneOnly)
|
||||
return emitStaticIntLookup(
|
||||
*values, kind == Kind::ActionOnly ? row : column,
|
||||
constantAnchor, constants, builder, loc);
|
||||
Value flat = affineMulConst(builder, loc, row, columns, constantAnchor);
|
||||
flat = arith::AddIOp::create(builder, loc, flat, column);
|
||||
if (kind == Kind::Affine) {
|
||||
Value rowValue = affineMulConst(
|
||||
builder, loc, row, rowStep, constantAnchor);
|
||||
Value columnValue = affineMulConst(
|
||||
builder, loc, column, columnStep, constantAnchor);
|
||||
Value result = arith::AddIOp::create(builder, loc, rowValue, columnValue);
|
||||
return affineAddConst(builder, loc, result, base, constantAnchor);
|
||||
}
|
||||
return emitStaticIntLookup(
|
||||
*values, flat, constantAnchor, constants, builder, loc);
|
||||
}
|
||||
|
||||
OpFoldResult StaticIntGrid::emitFoldedLookup(
|
||||
Value row, Value column, Operation *constantAnchor,
|
||||
ConstantPool &constants, OpBuilder &builder, Location loc) const {
|
||||
return kind == Kind::Uniform
|
||||
? OpFoldResult(builder.getIndexAttr(base))
|
||||
: OpFoldResult(emitLookup(
|
||||
row, column, constantAnchor, constants, builder, loc));
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,61 @@
|
||||
#pragma once
|
||||
|
||||
#include "StaticIntSequence.hpp"
|
||||
|
||||
#include "llvm/ADT/ArrayRef.h"
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
|
||||
#include <utility>
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
class ConstantPool;
|
||||
|
||||
class StaticIntGrid {
|
||||
public:
|
||||
static mlir::FailureOr<StaticIntGrid> fromColumns(
|
||||
size_t rows, llvm::ArrayRef<StaticIntSequence> columns,
|
||||
int64_t defaultValue);
|
||||
static mlir::FailureOr<StaticIntGrid> fromRows(
|
||||
llvm::ArrayRef<StaticIntSequence> rows);
|
||||
static mlir::FailureOr<StaticIntGrid> affine2D(
|
||||
int64_t base, int64_t rowStep, int64_t columnStep,
|
||||
size_t rows, size_t columns);
|
||||
static mlir::FailureOr<StaticIntGrid> laneIntervals(
|
||||
size_t columns,
|
||||
llvm::ArrayRef<std::pair<size_t, size_t>> intervals,
|
||||
int64_t insideValue, int64_t outsideValue);
|
||||
|
||||
int64_t valueAt(size_t row, size_t column) const;
|
||||
|
||||
mlir::Value emitLookup(mlir::Value row, mlir::Value column,
|
||||
mlir::Operation *constantAnchor,
|
||||
ConstantPool &constants, mlir::OpBuilder &builder,
|
||||
mlir::Location loc) const;
|
||||
mlir::OpFoldResult emitFoldedLookup(
|
||||
mlir::Value row, mlir::Value column, mlir::Operation *constantAnchor,
|
||||
ConstantPool &constants, mlir::OpBuilder &builder,
|
||||
mlir::Location loc) const;
|
||||
|
||||
private:
|
||||
enum class Kind { Uniform, ActionOnly, LaneOnly, Affine,
|
||||
SparseLaneOverrides, Dense };
|
||||
|
||||
StaticIntGrid(size_t rows, size_t columns, int64_t base)
|
||||
: rows(rows), columns(columns), base(base) {}
|
||||
|
||||
static mlir::FailureOr<StaticIntGrid> fromSequences(
|
||||
llvm::ArrayRef<StaticIntSequence> sequences, bool columns,
|
||||
int64_t sparseBase);
|
||||
|
||||
Kind kind = Kind::Uniform;
|
||||
size_t rows = 0;
|
||||
size_t columns = 0;
|
||||
int64_t base = 0;
|
||||
int64_t rowStep = 0;
|
||||
int64_t columnStep = 0;
|
||||
llvm::SmallVector<int64_t> overrideKeys;
|
||||
std::optional<StaticIntSequence> values;
|
||||
};
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,539 @@
|
||||
#include "mlir/Dialect/Arith/IR/Arith.h"
|
||||
#include "mlir/Dialect/Tensor/IR/Tensor.h"
|
||||
|
||||
#include "llvm/ADT/STLExtras.h"
|
||||
#include "llvm/Support/ErrorHandling.h"
|
||||
#include "llvm/Support/MathExtras.h"
|
||||
|
||||
#include <limits>
|
||||
|
||||
#include "AffineUtils.hpp"
|
||||
#include "ConstantUtils.hpp"
|
||||
#include "StaticIntSequence.hpp"
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
namespace {
|
||||
|
||||
static bool getAffineValue(int64_t base, int64_t step, size_t index,
|
||||
int64_t &value) {
|
||||
if (index > static_cast<size_t>(std::numeric_limits<int64_t>::max()))
|
||||
return false;
|
||||
int64_t scaled;
|
||||
return !llvm::MulOverflow(static_cast<int64_t>(index), step, scaled)
|
||||
&& !llvm::AddOverflow(base, scaled, value);
|
||||
}
|
||||
|
||||
static FailureOr<SmallVector<int64_t>> getI64Values(Operation *op,
|
||||
StringRef name) {
|
||||
Attribute attr = op->getAttr(name);
|
||||
if (!attr)
|
||||
return op->emitOpError() << "is missing " << name << " metadata",
|
||||
failure();
|
||||
if (auto scalar = dyn_cast<IntegerAttr>(attr))
|
||||
return SmallVector<int64_t> {scalar.getInt()};
|
||||
if (auto array = dyn_cast<DenseI64ArrayAttr>(attr))
|
||||
return SmallVector<int64_t>(array.asArrayRef());
|
||||
auto elements = dyn_cast<DenseIntElementsAttr>(attr);
|
||||
auto type = elements ? dyn_cast<RankedTensorType>(elements.getType())
|
||||
: RankedTensorType();
|
||||
if (!elements || !type || type.getRank() != 1
|
||||
|| !type.getElementType().isInteger(64))
|
||||
return op->emitOpError() << "has invalid " << name << " metadata",
|
||||
failure();
|
||||
SmallVector<int64_t> values;
|
||||
values.reserve(elements.getNumElements());
|
||||
for (APInt value : elements.getValues<APInt>())
|
||||
values.push_back(value.getSExtValue());
|
||||
return values;
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
StaticIntSequence StaticIntSequence::uniform(int64_t value, size_t count) {
|
||||
assert(count != 0 && "empty static integer sequence");
|
||||
StaticIntSequence result;
|
||||
result.kind = StaticIntSequenceKind::Uniform;
|
||||
result.count = count;
|
||||
result.base = value;
|
||||
return result;
|
||||
}
|
||||
|
||||
StaticIntSequence StaticIntSequence::affine(int64_t base, int64_t step,
|
||||
size_t count) {
|
||||
assert(count != 0 && "empty static integer sequence");
|
||||
int64_t last;
|
||||
assert(getAffineValue(base, step, count - 1, last)
|
||||
&& "overflowing static affine sequence");
|
||||
if (count == 1 || step == 0)
|
||||
return uniform(base, count);
|
||||
StaticIntSequence result;
|
||||
result.kind = StaticIntSequenceKind::Affine;
|
||||
result.count = count;
|
||||
result.base = base;
|
||||
result.step = step;
|
||||
return result;
|
||||
}
|
||||
|
||||
StaticIntSequence StaticIntSequence::runLengthEncoded(
|
||||
ArrayRef<int64_t> runs, size_t count) {
|
||||
assert(count != 0 && runs.size() % 2 == 0
|
||||
&& "invalid run-length encoded sequence");
|
||||
StaticIntSequence result;
|
||||
result.kind = StaticIntSequenceKind::RunLengthEncoded;
|
||||
result.count = count;
|
||||
result.data.assign(runs);
|
||||
return result;
|
||||
}
|
||||
|
||||
StaticIntSequence StaticIntSequence::fromValues(ArrayRef<int64_t> values) {
|
||||
assert(!values.empty() && "empty static integer sequence");
|
||||
if (llvm::all_equal(values))
|
||||
return uniform(values.front(), values.size());
|
||||
int64_t step;
|
||||
bool isAffine = !llvm::SubOverflow(values[1], values[0], step);
|
||||
for (size_t index = 1; isAffine && index < values.size(); ++index) {
|
||||
int64_t difference;
|
||||
isAffine = !llvm::SubOverflow(values[index], values[index - 1],
|
||||
difference)
|
||||
&& difference == step;
|
||||
}
|
||||
if (isAffine)
|
||||
return affine(values.front(), step, values.size());
|
||||
|
||||
SmallVector<int64_t> runs;
|
||||
for (int64_t value : values) {
|
||||
if (!runs.empty() && runs[runs.size() - 2] == value) {
|
||||
++runs.back();
|
||||
continue;
|
||||
}
|
||||
runs.push_back(value);
|
||||
runs.push_back(1);
|
||||
}
|
||||
if (runs.size() < values.size())
|
||||
return runLengthEncoded(runs, values.size());
|
||||
StaticIntSequence result;
|
||||
result.kind = StaticIntSequenceKind::Dense;
|
||||
result.count = values.size();
|
||||
result.data.assign(values);
|
||||
return result;
|
||||
}
|
||||
|
||||
int64_t StaticIntSequence::valueAt(size_t index) const {
|
||||
assert(index < count && "static integer sequence index out of range");
|
||||
if (kind == StaticIntSequenceKind::Uniform)
|
||||
return base;
|
||||
if (kind == StaticIntSequenceKind::Affine) {
|
||||
int64_t value;
|
||||
bool valid = getAffineValue(base, step, index, value);
|
||||
assert(valid && "overflowing static affine sequence");
|
||||
return value;
|
||||
}
|
||||
if (kind == StaticIntSequenceKind::Dense)
|
||||
return data[index];
|
||||
for (size_t run = 0; run < data.size(); run += 2) {
|
||||
size_t length = static_cast<size_t>(data[run + 1]);
|
||||
if (index < length)
|
||||
return data[run];
|
||||
index -= length;
|
||||
}
|
||||
llvm_unreachable("malformed run-length encoded sequence");
|
||||
}
|
||||
|
||||
std::optional<size_t> StaticIntSequence::find(int64_t value, size_t begin,
|
||||
size_t length) const {
|
||||
assert(begin <= count && length <= count - begin
|
||||
&& "invalid static integer sequence search");
|
||||
if (length == 0)
|
||||
return std::nullopt;
|
||||
size_t end = begin + length;
|
||||
if (kind == StaticIntSequenceKind::Uniform)
|
||||
return value == base ? std::optional<size_t>(begin) : std::nullopt;
|
||||
if (kind == StaticIntSequenceKind::Affine) {
|
||||
int64_t delta;
|
||||
if (llvm::SubOverflow(value, base, delta) || delta % step != 0)
|
||||
return std::nullopt;
|
||||
int64_t index = delta / step;
|
||||
return index >= static_cast<int64_t>(begin)
|
||||
&& index < static_cast<int64_t>(end)
|
||||
? std::optional<size_t>(index)
|
||||
: std::nullopt;
|
||||
}
|
||||
if (kind == StaticIntSequenceKind::Dense) {
|
||||
ArrayRef<int64_t> selected = ArrayRef(data).slice(begin, length);
|
||||
auto found = llvm::find(selected, value);
|
||||
return found == selected.end()
|
||||
? std::nullopt
|
||||
: std::optional<size_t>(begin + (found - selected.begin()));
|
||||
}
|
||||
size_t runBegin = 0;
|
||||
for (size_t run = 0; run < data.size(); run += 2) {
|
||||
size_t runEnd = runBegin + static_cast<size_t>(data[run + 1]);
|
||||
if (runEnd > begin && runBegin < end && data[run] == value)
|
||||
return std::max(begin, runBegin);
|
||||
if (runBegin >= end)
|
||||
break;
|
||||
runBegin = runEnd;
|
||||
}
|
||||
return std::nullopt;
|
||||
}
|
||||
|
||||
StaticIntSequence StaticIntSequence::slice(size_t begin, size_t length) const {
|
||||
assert(length != 0 && begin <= count - length && "invalid sequence slice");
|
||||
if (kind == StaticIntSequenceKind::Uniform)
|
||||
return uniform(base, length);
|
||||
if (kind == StaticIntSequenceKind::Affine)
|
||||
return affine(valueAt(begin), step, length);
|
||||
if (kind == StaticIntSequenceKind::Dense)
|
||||
return fromValues(ArrayRef(data).slice(begin, length));
|
||||
SmallVector<int64_t> runs;
|
||||
size_t end = begin + length;
|
||||
forEachEqualRun([&](int64_t value, size_t runBegin, size_t runCount) {
|
||||
size_t selectedBegin = std::max(begin, runBegin);
|
||||
size_t selectedEnd = std::min(end, runBegin + runCount);
|
||||
if (selectedBegin >= selectedEnd)
|
||||
return;
|
||||
if (!runs.empty() && runs[runs.size() - 2] == value)
|
||||
runs.back() += selectedEnd - selectedBegin;
|
||||
else {
|
||||
runs.push_back(value);
|
||||
runs.push_back(selectedEnd - selectedBegin);
|
||||
}
|
||||
});
|
||||
if (runs.size() == 2)
|
||||
return uniform(runs.front(), length);
|
||||
if (runs.size() < length)
|
||||
return runLengthEncoded(runs, length);
|
||||
SmallVector<int64_t> values;
|
||||
for (size_t run = 0; run < runs.size(); run += 2)
|
||||
values.append(runs[run + 1], runs[run]);
|
||||
return fromValues(values);
|
||||
}
|
||||
|
||||
StaticIntSequence StaticIntSequence::remap(ArrayRef<unsigned> indices) const {
|
||||
assert(!indices.empty() && "empty static integer sequence remap");
|
||||
SmallVector<int64_t> values;
|
||||
values.reserve(indices.size());
|
||||
for (unsigned index : indices)
|
||||
values.push_back(valueAt(index));
|
||||
return fromValues(values);
|
||||
}
|
||||
|
||||
bool StaticIntSequence::operator==(const StaticIntSequence& other) const {
|
||||
return kind == other.kind && count == other.count && base == other.base
|
||||
&& step == other.step && data == other.data;
|
||||
}
|
||||
|
||||
llvm::hash_code StaticIntSequence::hash() const {
|
||||
return llvm::hash_combine(kind, count, base, step,
|
||||
llvm::hash_combine_range(data.begin(), data.end()));
|
||||
}
|
||||
|
||||
void StaticIntSequence::forEachEqualRun(
|
||||
llvm::function_ref<void(int64_t, size_t, size_t)> callback) const {
|
||||
if (kind == StaticIntSequenceKind::Uniform) {
|
||||
callback(base, 0, count);
|
||||
return;
|
||||
}
|
||||
if (kind == StaticIntSequenceKind::RunLengthEncoded) {
|
||||
size_t begin = 0;
|
||||
for (size_t run = 0; run < data.size(); run += 2) {
|
||||
size_t runCount = static_cast<size_t>(data[run + 1]);
|
||||
callback(data[run], begin, runCount);
|
||||
begin += runCount;
|
||||
}
|
||||
return;
|
||||
}
|
||||
size_t begin = 0;
|
||||
while (begin < count) {
|
||||
int64_t value = valueAt(begin);
|
||||
size_t end = begin + 1;
|
||||
while (end < count && valueAt(end) == value)
|
||||
++end;
|
||||
callback(value, begin, end - begin);
|
||||
begin = end;
|
||||
}
|
||||
}
|
||||
|
||||
void StaticIntSequenceChain::append(const StaticIntSequence &sequence,
|
||||
size_t begin, size_t length) {
|
||||
assert(length != 0 && begin <= sequence.size() - length
|
||||
&& "invalid static integer sequence chain slice");
|
||||
if (!slices.empty()) {
|
||||
StaticIntSequenceSlice &last = slices.back();
|
||||
if (last.sequence == &sequence && last.begin + last.count == begin) {
|
||||
last.count += length;
|
||||
count += length;
|
||||
return;
|
||||
}
|
||||
auto affinePart = [](const StaticIntSequenceSlice &slice,
|
||||
int64_t &base, int64_t &step) {
|
||||
base = slice.sequence->valueAt(slice.begin);
|
||||
if (slice.count == 1) {
|
||||
step = 0;
|
||||
return true;
|
||||
}
|
||||
return !llvm::SubOverflow(slice.sequence->valueAt(slice.begin + 1),
|
||||
base, step)
|
||||
&& (slice.sequence->kind == StaticIntSequenceKind::Uniform
|
||||
|| slice.sequence->kind == StaticIntSequenceKind::Affine);
|
||||
};
|
||||
StaticIntSequenceSlice next {&sequence, begin, length};
|
||||
int64_t leftBase, leftStep, rightBase, rightStep, expected;
|
||||
if (affinePart(last, leftBase, leftStep)
|
||||
&& affinePart(next, rightBase, rightStep)
|
||||
&& (last.count == 1 || length == 1 || leftStep == rightStep)) {
|
||||
int64_t step = last.count == 1 ? rightStep : leftStep;
|
||||
if (getAffineValue(leftBase, step, last.count, expected)
|
||||
&& expected == rightBase) {
|
||||
owned.push_back(std::make_unique<StaticIntSequence>(
|
||||
StaticIntSequence::affine(leftBase, step, last.count + length)));
|
||||
last = {owned.back().get(), 0, last.count + length};
|
||||
count += length;
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
slices.push_back({&sequence, begin, length});
|
||||
count += length;
|
||||
}
|
||||
|
||||
void StaticIntSequenceChain::append(StaticIntSequence sequence) {
|
||||
size_t length = sequence.size();
|
||||
owned.push_back(std::make_unique<StaticIntSequence>(std::move(sequence)));
|
||||
append(*owned.back(), 0, length);
|
||||
}
|
||||
|
||||
int64_t StaticIntSequenceChain::valueAt(size_t index) const {
|
||||
assert(index < count && "static integer sequence chain index out of range");
|
||||
for (const StaticIntSequenceSlice &slice : slices) {
|
||||
if (index < slice.count)
|
||||
return slice.sequence->valueAt(slice.begin + index);
|
||||
index -= slice.count;
|
||||
}
|
||||
llvm_unreachable("malformed static integer sequence chain");
|
||||
}
|
||||
|
||||
void StaticIntSequenceChain::forEachSegment(llvm::function_ref<void(
|
||||
const StaticIntSequence &, size_t, size_t)> callback) const {
|
||||
for (const StaticIntSequenceSlice &slice : slices)
|
||||
callback(*slice.sequence, slice.begin, slice.count);
|
||||
}
|
||||
|
||||
void StaticIntSequenceChain::forEachEqualRun(
|
||||
llvm::function_ref<void(int64_t, size_t, size_t)> callback) const {
|
||||
std::optional<int64_t> pendingValue;
|
||||
size_t pendingBegin = 0, pendingCount = 0, chainBegin = 0;
|
||||
auto flush = [&] {
|
||||
if (pendingValue)
|
||||
callback(*pendingValue, pendingBegin, pendingCount);
|
||||
};
|
||||
for (const StaticIntSequenceSlice &slice : slices) {
|
||||
size_t sliceEnd = slice.begin + slice.count;
|
||||
slice.sequence->forEachEqualRun(
|
||||
[&](int64_t value, size_t runBegin, size_t runCount) {
|
||||
size_t begin = std::max(slice.begin, runBegin);
|
||||
size_t end = std::min(sliceEnd, runBegin + runCount);
|
||||
if (begin >= end)
|
||||
return;
|
||||
size_t selectedCount = end - begin;
|
||||
size_t globalBegin = chainBegin + begin - slice.begin;
|
||||
if (pendingValue && *pendingValue == value
|
||||
&& pendingBegin + pendingCount == globalBegin) {
|
||||
pendingCount += selectedCount;
|
||||
return;
|
||||
}
|
||||
flush();
|
||||
pendingValue = value;
|
||||
pendingBegin = globalBegin;
|
||||
pendingCount = selectedCount;
|
||||
});
|
||||
chainBegin += slice.count;
|
||||
}
|
||||
flush();
|
||||
}
|
||||
|
||||
StaticIntSequence StaticIntSequenceChain::canonicalize() const {
|
||||
assert(count != 0 && "empty static integer sequence chain");
|
||||
int64_t first = valueAt(0);
|
||||
bool uniform = true;
|
||||
forEachEqualRun([&](int64_t value, size_t, size_t) {
|
||||
uniform &= value == first;
|
||||
});
|
||||
if (uniform)
|
||||
return StaticIntSequence::uniform(first, count);
|
||||
|
||||
int64_t step = 0, previous = first;
|
||||
bool affine = true, haveStep = false;
|
||||
size_t position = 0;
|
||||
forEachSegment([&](const StaticIntSequence &sequence, size_t begin,
|
||||
size_t length) {
|
||||
if (!affine)
|
||||
return;
|
||||
for (size_t index = 0; index < length; ++index) {
|
||||
int64_t value = sequence.valueAt(begin + index);
|
||||
if (position++ == 0) {
|
||||
previous = value;
|
||||
continue;
|
||||
}
|
||||
if (!haveStep) {
|
||||
affine = !llvm::SubOverflow(value, previous, step);
|
||||
haveStep = true;
|
||||
} else if (haveStep) {
|
||||
int64_t difference;
|
||||
affine = !llvm::SubOverflow(value, previous, difference)
|
||||
&& difference == step;
|
||||
}
|
||||
previous = value;
|
||||
if (!affine)
|
||||
break;
|
||||
}
|
||||
});
|
||||
if (affine && haveStep)
|
||||
return StaticIntSequence::affine(first, step, count);
|
||||
|
||||
SmallVector<int64_t> runs;
|
||||
forEachEqualRun([&](int64_t value, size_t, size_t runCount) {
|
||||
runs.push_back(value);
|
||||
runs.push_back(runCount);
|
||||
});
|
||||
if (runs.size() < count)
|
||||
return StaticIntSequence::runLengthEncoded(runs, count);
|
||||
SmallVector<int64_t> values;
|
||||
values.reserve(count);
|
||||
for (size_t run = 0; run < runs.size(); run += 2)
|
||||
values.append(runs[run + 1], runs[run]);
|
||||
return StaticIntSequence::fromValues(values);
|
||||
}
|
||||
|
||||
int64_t StaticIntSequenceChainCursor::value() const {
|
||||
assert(!done() && "static integer sequence chain cursor is done");
|
||||
const StaticIntSequenceSlice ¤t = chain.slices[slice];
|
||||
return current.sequence->valueAt(current.begin + offset);
|
||||
}
|
||||
|
||||
void StaticIntSequenceChainCursor::advance() {
|
||||
assert(!done() && "static integer sequence chain cursor is done");
|
||||
if (++offset != chain.slices[slice].count)
|
||||
return;
|
||||
offset = 0;
|
||||
++slice;
|
||||
}
|
||||
|
||||
void setStaticIntSequenceAttr(Operation *op, StringRef name,
|
||||
const StaticIntSequence &sequence,
|
||||
size_t logicalCount) {
|
||||
assert(sequence.size() == logicalCount && logicalCount != 0
|
||||
&& "invalid static integer metadata count");
|
||||
SmallVector<int64_t> values;
|
||||
StringRef encoding;
|
||||
switch (sequence.kind) {
|
||||
case StaticIntSequenceKind::Uniform:
|
||||
encoding = "uniform";
|
||||
values.push_back(sequence.base);
|
||||
break;
|
||||
case StaticIntSequenceKind::Affine:
|
||||
encoding = "affine";
|
||||
values = {sequence.base, sequence.step};
|
||||
break;
|
||||
case StaticIntSequenceKind::RunLengthEncoded:
|
||||
encoding = "rle";
|
||||
values = sequence.data;
|
||||
break;
|
||||
case StaticIntSequenceKind::Dense:
|
||||
encoding = "dense";
|
||||
values = sequence.data;
|
||||
break;
|
||||
}
|
||||
OpBuilder builder(op);
|
||||
auto type = RankedTensorType::get(
|
||||
{static_cast<int64_t>(values.size())}, builder.getI64Type());
|
||||
op->setAttr(name, DenseIntElementsAttr::get(type, values));
|
||||
if (sequence.kind != StaticIntSequenceKind::Dense)
|
||||
op->setAttr((name + "_encoding").str(), builder.getStringAttr(encoding));
|
||||
}
|
||||
|
||||
FailureOr<StaticIntSequence> getStaticIntSequenceAttr(
|
||||
Operation *op, StringRef name, size_t logicalCount) {
|
||||
if (logicalCount == 0)
|
||||
return op->emitOpError() << "has zero logical count for " << name,
|
||||
failure();
|
||||
auto values = getI64Values(op, name);
|
||||
if (failed(values))
|
||||
return failure();
|
||||
auto encoding = op->getAttrOfType<StringAttr>((name + "_encoding").str());
|
||||
if (!encoding) {
|
||||
if (values->size() != logicalCount)
|
||||
return op->emitOpError() << "has invalid dense " << name << " count",
|
||||
failure();
|
||||
return StaticIntSequence::fromValues(*values);
|
||||
}
|
||||
if (encoding.getValue() == "uniform") {
|
||||
if (values->size() != 1)
|
||||
return op->emitOpError() << "has invalid uniform " << name,
|
||||
failure();
|
||||
return StaticIntSequence::uniform(values->front(), logicalCount);
|
||||
}
|
||||
if (encoding.getValue() == "affine") {
|
||||
int64_t last;
|
||||
if (values->size() != 2
|
||||
|| !getAffineValue((*values)[0], (*values)[1], logicalCount - 1, last))
|
||||
return op->emitOpError() << "has invalid affine " << name,
|
||||
failure();
|
||||
return StaticIntSequence::affine((*values)[0], (*values)[1], logicalCount);
|
||||
}
|
||||
if (encoding.getValue() == "rle") {
|
||||
size_t count = 0;
|
||||
if (values->empty() || values->size() % 2 != 0)
|
||||
return op->emitOpError() << "has invalid RLE " << name, failure();
|
||||
for (size_t index = 1; index < values->size(); index += 2) {
|
||||
if ((*values)[index] <= 0
|
||||
|| static_cast<uint64_t>((*values)[index]) > logicalCount - count)
|
||||
return op->emitOpError() << "has invalid RLE " << name, failure();
|
||||
count += (*values)[index];
|
||||
}
|
||||
if (count != logicalCount)
|
||||
return op->emitOpError() << "has mismatched RLE " << name << " count",
|
||||
failure();
|
||||
return StaticIntSequence::runLengthEncoded(*values, count);
|
||||
}
|
||||
if (encoding.getValue() == "dense") {
|
||||
if (values->size() != logicalCount)
|
||||
return op->emitOpError() << "has invalid dense " << name << " count",
|
||||
failure();
|
||||
return StaticIntSequence::fromValues(*values);
|
||||
}
|
||||
return op->emitOpError() << "has unknown " << name << " encoding",
|
||||
failure();
|
||||
}
|
||||
|
||||
Value emitStaticIntLookup(const StaticIntSequence& sequence, Value position,
|
||||
Operation* constantAnchor,
|
||||
ConstantPool& constants, OpBuilder& builder,
|
||||
Location loc) {
|
||||
if (sequence.getKind() == StaticIntSequenceKind::Uniform)
|
||||
return constants.getIndex(sequence.valueAt(0));
|
||||
if (sequence.getKind() == StaticIntSequenceKind::Affine) {
|
||||
Value scaled = affineMulConst(builder, loc, position,
|
||||
sequence.valueAt(1) - sequence.valueAt(0),
|
||||
constantAnchor);
|
||||
return affineAddConst(builder, loc, scaled, sequence.valueAt(0),
|
||||
constantAnchor);
|
||||
}
|
||||
SmallVector<int64_t> values;
|
||||
values.reserve(sequence.size());
|
||||
sequence.forEachEqualRun([&](int64_t value, size_t, size_t count) {
|
||||
values.append(count, value);
|
||||
});
|
||||
auto type = RankedTensorType::get(
|
||||
{static_cast<int64_t>(values.size())}, builder.getI64Type());
|
||||
Value table = constants.get(type,
|
||||
DenseElementsAttr::get(type, ArrayRef<int64_t>(values)));
|
||||
Value selected = tensor::ExtractOp::create(
|
||||
builder, loc, table, ValueRange {position});
|
||||
return arith::IndexCastOp::create(
|
||||
builder, loc, builder.getIndexType(), selected);
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,123 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/IR/Builders.h"
|
||||
|
||||
#include "llvm/ADT/ArrayRef.h"
|
||||
#include "llvm/ADT/FunctionExtras.h"
|
||||
#include "llvm/ADT/Hashing.h"
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <memory>
|
||||
#include <optional>
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
class ConstantPool;
|
||||
|
||||
enum class StaticIntSequenceKind {
|
||||
Uniform,
|
||||
Affine,
|
||||
RunLengthEncoded,
|
||||
Dense
|
||||
};
|
||||
|
||||
class StaticIntSequence {
|
||||
public:
|
||||
static StaticIntSequence fromValues(llvm::ArrayRef<int64_t> values);
|
||||
static StaticIntSequence uniform(int64_t value, size_t count);
|
||||
static StaticIntSequence affine(int64_t base, int64_t step, size_t count);
|
||||
|
||||
size_t size() const { return count; }
|
||||
int64_t valueAt(size_t index) const;
|
||||
std::optional<size_t> find(int64_t value, size_t begin, size_t length) const;
|
||||
StaticIntSequence slice(size_t begin, size_t count) const;
|
||||
StaticIntSequence remap(llvm::ArrayRef<unsigned> indices) const;
|
||||
StaticIntSequenceKind getKind() const { return kind; }
|
||||
std::optional<int64_t> getAffineStep() const {
|
||||
if (kind == StaticIntSequenceKind::Uniform)
|
||||
return 0;
|
||||
return kind == StaticIntSequenceKind::Affine
|
||||
? std::optional<int64_t>(step) : std::nullopt;
|
||||
}
|
||||
|
||||
bool operator==(const StaticIntSequence& other) const;
|
||||
llvm::hash_code hash() const;
|
||||
|
||||
void forEachEqualRun(
|
||||
llvm::function_ref<void(int64_t, size_t, size_t)> callback) const;
|
||||
|
||||
private:
|
||||
friend class StaticIntSequenceChain;
|
||||
friend void setStaticIntSequenceAttr(mlir::Operation *, llvm::StringRef,
|
||||
const StaticIntSequence &, size_t);
|
||||
friend mlir::FailureOr<StaticIntSequence>
|
||||
getStaticIntSequenceAttr(mlir::Operation *, llvm::StringRef, size_t);
|
||||
|
||||
static StaticIntSequence runLengthEncoded(
|
||||
llvm::ArrayRef<int64_t> runs, size_t count);
|
||||
StaticIntSequenceKind kind = StaticIntSequenceKind::Dense;
|
||||
size_t count = 0;
|
||||
int64_t base = 0;
|
||||
int64_t step = 0;
|
||||
llvm::SmallVector<int64_t> data;
|
||||
};
|
||||
|
||||
struct StaticIntSequenceSlice {
|
||||
const StaticIntSequence *sequence = nullptr;
|
||||
size_t begin = 0;
|
||||
size_t count = 0;
|
||||
};
|
||||
|
||||
class StaticIntSequenceChain {
|
||||
public:
|
||||
void append(const StaticIntSequence &sequence, size_t begin, size_t count);
|
||||
void append(StaticIntSequence sequence);
|
||||
size_t size() const { return count; }
|
||||
int64_t valueAt(size_t index) const;
|
||||
void forEachSegment(llvm::function_ref<void(
|
||||
const StaticIntSequence &, size_t, size_t)> callback) const;
|
||||
void forEachEqualRun(
|
||||
llvm::function_ref<void(int64_t, size_t, size_t)> callback) const;
|
||||
StaticIntSequence canonicalize() const;
|
||||
|
||||
private:
|
||||
friend class StaticIntSequenceChainCursor;
|
||||
llvm::SmallVector<StaticIntSequenceSlice> slices;
|
||||
llvm::SmallVector<std::unique_ptr<StaticIntSequence>> owned;
|
||||
size_t count = 0;
|
||||
};
|
||||
|
||||
class StaticIntSequenceChainCursor {
|
||||
public:
|
||||
explicit StaticIntSequenceChainCursor(const StaticIntSequenceChain &chain)
|
||||
: chain(chain) {}
|
||||
|
||||
bool done() const { return slice == chain.slices.size(); }
|
||||
int64_t value() const;
|
||||
void advance();
|
||||
|
||||
private:
|
||||
const StaticIntSequenceChain &chain;
|
||||
size_t slice = 0;
|
||||
size_t offset = 0;
|
||||
};
|
||||
|
||||
void setStaticIntSequenceAttr(mlir::Operation *op, llvm::StringRef name,
|
||||
const StaticIntSequence &sequence,
|
||||
size_t logicalCount);
|
||||
|
||||
mlir::FailureOr<StaticIntSequence>
|
||||
getStaticIntSequenceAttr(mlir::Operation *op, llvm::StringRef name,
|
||||
size_t logicalCount);
|
||||
|
||||
mlir::Value emitStaticIntLookup(const StaticIntSequence& sequence,
|
||||
mlir::Value position,
|
||||
mlir::Operation* constantAnchor,
|
||||
ConstantPool& constants,
|
||||
mlir::OpBuilder& builder,
|
||||
mlir::Location loc);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,106 @@
|
||||
#include "mlir/IR/BuiltinTypeInterfaces.h"
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/SubviewUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/PimCommon.hpp"
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
Value stripMemRefCasts(Value value) {
|
||||
while (auto castOp = value.getDefiningOp<memref::CastOp>())
|
||||
value = castOp.getSource();
|
||||
return value;
|
||||
}
|
||||
|
||||
Value stripMemRefViewOps(Value value) {
|
||||
while (true) {
|
||||
if (auto castOp = value.getDefiningOp<memref::CastOp>()) {
|
||||
value = castOp.getSource();
|
||||
continue;
|
||||
}
|
||||
if (auto collapseOp = value.getDefiningOp<memref::CollapseShapeOp>()) {
|
||||
value = collapseOp.getSrc();
|
||||
continue;
|
||||
}
|
||||
if (auto expandOp = value.getDefiningOp<memref::ExpandShapeOp>()) {
|
||||
value = expandOp.getSrc();
|
||||
continue;
|
||||
}
|
||||
return value;
|
||||
}
|
||||
}
|
||||
|
||||
Value stripMemRefAddressingOps(Value value) {
|
||||
while (true) {
|
||||
if (auto subviewOp = value.getDefiningOp<memref::SubViewOp>()) {
|
||||
value = subviewOp.getSource();
|
||||
continue;
|
||||
}
|
||||
Value strippedValue = stripMemRefViewOps(value);
|
||||
if (strippedValue == value)
|
||||
return value;
|
||||
value = strippedValue;
|
||||
}
|
||||
}
|
||||
|
||||
bool hasAllStaticSubviewParts(memref::SubViewOp subview) {
|
||||
return llvm::all_of(subview.getStaticOffsets(), [](int64_t value) { return !ShapedType::isDynamic(value); })
|
||||
&& llvm::all_of(subview.getStaticSizes(), [](int64_t value) { return !ShapedType::isDynamic(value); })
|
||||
&& llvm::all_of(subview.getStaticStrides(), [](int64_t value) { return !ShapedType::isDynamic(value); });
|
||||
}
|
||||
|
||||
FailureOr<StaticSubviewInfo> getStaticSubviewInfo(Value value) {
|
||||
value = stripMemRefViewOps(value);
|
||||
auto subviewOp = value.getDefiningOp<memref::SubViewOp>();
|
||||
if (!subviewOp)
|
||||
return failure();
|
||||
|
||||
auto source = stripMemRefCasts(subviewOp.getSource());
|
||||
auto sourceType = dyn_cast<MemRefType>(source.getType());
|
||||
auto subviewType = dyn_cast<MemRefType>(subviewOp.getType());
|
||||
if (!sourceType || !subviewType || !sourceType.hasStaticShape() || !subviewType.hasStaticShape())
|
||||
return failure();
|
||||
|
||||
StaticSubviewInfo info;
|
||||
info.source = source;
|
||||
info.sourceShape.assign(sourceType.getShape().begin(), sourceType.getShape().end());
|
||||
SmallVector<OpFoldResult> mixedOffsets = subviewOp.getMixedOffsets();
|
||||
info.offsets.assign(mixedOffsets.begin(), mixedOffsets.end());
|
||||
for (OpFoldResult size : subviewOp.getMixedSizes()) {
|
||||
auto staticSize = getConstantIntValue(size);
|
||||
if (!staticSize)
|
||||
return failure();
|
||||
info.sizes.push_back(*staticSize);
|
||||
}
|
||||
for (OpFoldResult stride : subviewOp.getMixedStrides()) {
|
||||
auto staticStride = getConstantIntValue(stride);
|
||||
if (!staticStride)
|
||||
return failure();
|
||||
info.strides.push_back(*staticStride);
|
||||
}
|
||||
return info;
|
||||
}
|
||||
|
||||
FailureOr<SmallVector<int64_t>> getStaticSubviewOffsets(const StaticSubviewInfo& info) {
|
||||
SmallVector<int64_t> staticOffsets;
|
||||
staticOffsets.reserve(info.offsets.size());
|
||||
for (OpFoldResult offset : info.offsets) {
|
||||
auto staticOffset = getConstantIntValue(offset);
|
||||
if (!staticOffset)
|
||||
return failure();
|
||||
staticOffsets.push_back(*staticOffset);
|
||||
}
|
||||
return staticOffsets;
|
||||
}
|
||||
|
||||
bool isMemRefBaseAddressableValue(Value value) {
|
||||
value = stripMemRefAddressingOps(value);
|
||||
if (isa<BlockArgument>(value))
|
||||
return true;
|
||||
|
||||
Operation* defOp = value.getDefiningOp();
|
||||
return defOp && isa<memref::AllocOp, memref::GetGlobalOp>(defOp);
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,34 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/Dialect/MemRef/IR/MemRef.h"
|
||||
#include "mlir/IR/Value.h"
|
||||
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
#include "llvm/Support/LogicalResult.h"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
struct StaticSubviewInfo {
|
||||
mlir::Value source;
|
||||
llvm::SmallVector<int64_t> sourceShape;
|
||||
llvm::SmallVector<mlir::OpFoldResult> offsets;
|
||||
llvm::SmallVector<int64_t> sizes;
|
||||
llvm::SmallVector<int64_t> strides;
|
||||
};
|
||||
|
||||
mlir::Value stripMemRefCasts(mlir::Value value);
|
||||
|
||||
mlir::Value stripMemRefViewOps(mlir::Value value);
|
||||
|
||||
mlir::Value stripMemRefAddressingOps(mlir::Value value);
|
||||
|
||||
bool hasAllStaticSubviewParts(mlir::memref::SubViewOp subview);
|
||||
|
||||
llvm::FailureOr<StaticSubviewInfo> getStaticSubviewInfo(mlir::Value value);
|
||||
|
||||
/// Returns the offsets in `info` as int64_t, failing if any offset is dynamic.
|
||||
llvm::FailureOr<llvm::SmallVector<int64_t>> getStaticSubviewOffsets(const StaticSubviewInfo& info);
|
||||
|
||||
bool isMemRefBaseAddressableValue(mlir::Value value);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,133 @@
|
||||
#include "mlir/Dialect/Tensor/IR/Tensor.h"
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/TensorSliceUtils.hpp"
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
Value extractAxisSlice(
|
||||
PatternRewriter& rewriter, Location loc, Value source, int64_t axis, int64_t offset, int64_t size) {
|
||||
auto sourceType = cast<RankedTensorType>(source.getType());
|
||||
SmallVector<int64_t> resultShape(sourceType.getShape());
|
||||
resultShape[axis] = size;
|
||||
auto resultType = RankedTensorType::get(resultShape, sourceType.getElementType(), sourceType.getEncoding());
|
||||
|
||||
SmallVector<OpFoldResult> offsets = getZeroOffsets(rewriter, sourceType.getRank());
|
||||
SmallVector<OpFoldResult> sizes = getStaticSizes(rewriter, sourceType.getShape());
|
||||
offsets[axis] = rewriter.getIndexAttr(offset);
|
||||
sizes[axis] = rewriter.getIndexAttr(size);
|
||||
return tensor::ExtractSliceOp::create(
|
||||
rewriter, loc, resultType, source, offsets, sizes, getUnitStrides(rewriter, sourceType.getRank()))
|
||||
.getResult();
|
||||
}
|
||||
|
||||
Value extractStaticSliceOrIdentity(OpBuilder& rewriter,
|
||||
Location loc,
|
||||
Value source,
|
||||
RankedTensorType resultType,
|
||||
ArrayRef<OpFoldResult> offsets,
|
||||
ArrayRef<OpFoldResult> sizes,
|
||||
ArrayRef<OpFoldResult> strides) {
|
||||
auto sourceType = cast<RankedTensorType>(source.getType());
|
||||
size_t rank = static_cast<size_t>(sourceType.getRank());
|
||||
|
||||
bool isIdentitySlice =
|
||||
sourceType == resultType && sourceType.hasStaticShape() && offsets.size() == rank && sizes.size() == rank
|
||||
&& strides.size() == rank;
|
||||
if (isIdentitySlice) {
|
||||
ArrayRef<int64_t> sourceShape = sourceType.getShape();
|
||||
for (auto [dim, offset, size, stride] : llvm::zip_equal(sourceShape, offsets, sizes, strides)) {
|
||||
std::optional<int64_t> staticOffset = mlir::getConstantIntValue(offset);
|
||||
std::optional<int64_t> staticSize = mlir::getConstantIntValue(size);
|
||||
std::optional<int64_t> staticStride = mlir::getConstantIntValue(stride);
|
||||
if (!staticOffset || !staticSize || !staticStride || *staticOffset != 0 || *staticSize != dim
|
||||
|| *staticStride != 1) {
|
||||
isIdentitySlice = false;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (isIdentitySlice)
|
||||
return source;
|
||||
|
||||
return rewriter.createOrFold<tensor::ExtractSliceOp>(
|
||||
loc, resultType, source, offsets, sizes, strides);
|
||||
}
|
||||
|
||||
Value insertStaticSlice(
|
||||
PatternRewriter& rewriter, Location loc, Value source, Value dest, ArrayRef<OpFoldResult> offsets) {
|
||||
auto sourceType = cast<RankedTensorType>(source.getType());
|
||||
return tensor::InsertSliceOp::create(rewriter,
|
||||
loc,
|
||||
source,
|
||||
dest,
|
||||
offsets,
|
||||
getStaticSizes(rewriter, sourceType.getShape()),
|
||||
getUnitStrides(rewriter, sourceType.getRank()))
|
||||
.getResult();
|
||||
}
|
||||
|
||||
Value extractMixedSliceOrIdentity(OpBuilder &rewriter,
|
||||
Location loc,
|
||||
Value source,
|
||||
RankedTensorType resultType,
|
||||
const MixedSliceGeometry &geometry) {
|
||||
return extractStaticSliceOrIdentity(rewriter, loc, source, resultType,
|
||||
geometry.offsets, geometry.sizes,
|
||||
geometry.strides);
|
||||
}
|
||||
|
||||
Value insertMixedSlice(OpBuilder &builder, Location loc, Value source,
|
||||
Value dest, const MixedSliceGeometry &geometry) {
|
||||
SmallVector<OpFoldResult> sizes(geometry.sizes);
|
||||
auto sourceType = dyn_cast<RankedTensorType>(source.getType());
|
||||
auto destType = dyn_cast<RankedTensorType>(dest.getType());
|
||||
if (sourceType && destType && sourceType.hasStaticShape()
|
||||
&& sourceType.getRank() == destType.getRank()) {
|
||||
sizes.clear();
|
||||
for (int64_t dimension : sourceType.getShape())
|
||||
sizes.push_back(builder.getIndexAttr(dimension));
|
||||
}
|
||||
return tensor::InsertSliceOp::create(builder, loc, source, dest,
|
||||
geometry.offsets, sizes,
|
||||
geometry.strides);
|
||||
}
|
||||
|
||||
FailureOr<Value> addLeadingUnitTensorDimension(OpBuilder& builder, Location loc, Value value) {
|
||||
auto type = dyn_cast<RankedTensorType>(value.getType());
|
||||
if (!type || !type.hasStaticShape())
|
||||
return failure();
|
||||
SmallVector<int64_t> shape {1};
|
||||
llvm::append_range(shape, type.getShape());
|
||||
auto resultType = RankedTensorType::get(shape, type.getElementType(), type.getEncoding());
|
||||
SmallVector<ReassociationIndices> reassociation;
|
||||
if (type.getRank() != 0) {
|
||||
reassociation.push_back({0, 1});
|
||||
for (int64_t dim = 1; dim < type.getRank(); ++dim)
|
||||
reassociation.push_back({dim + 1});
|
||||
}
|
||||
return tensor::ExpandShapeOp::create(builder, loc, resultType, value, reassociation).getResult();
|
||||
}
|
||||
|
||||
FailureOr<Value> removeLeadingUnitTensorDimension(
|
||||
OpBuilder& builder, Location loc, Value value, RankedTensorType resultType) {
|
||||
if (value.getType() == resultType)
|
||||
return value;
|
||||
auto type = dyn_cast<RankedTensorType>(value.getType());
|
||||
if (!type || !resultType || !type.hasStaticShape() || !resultType.hasStaticShape()
|
||||
|| type.getRank() != resultType.getRank() + 1 || type.getDimSize(0) != 1
|
||||
|| type.getElementType() != resultType.getElementType()
|
||||
|| !llvm::equal(type.getShape().drop_front(), resultType.getShape()))
|
||||
return failure();
|
||||
SmallVector<ReassociationIndices> reassociation;
|
||||
if (resultType.getRank() != 0) {
|
||||
reassociation.push_back({0, 1});
|
||||
for (int64_t dim = 1; dim < resultType.getRank(); ++dim)
|
||||
reassociation.push_back({dim + 1});
|
||||
}
|
||||
return tensor::CollapseShapeOp::create(builder, loc, resultType, value, reassociation).getResult();
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,52 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/Dialect/Tensor/IR/Tensor.h"
|
||||
#include "mlir/IR/ValueRange.h"
|
||||
#include "mlir/Transforms/DialectConversion.h"
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
struct MixedSliceGeometry {
|
||||
llvm::SmallVector<mlir::OpFoldResult> offsets;
|
||||
llvm::SmallVector<mlir::OpFoldResult> sizes;
|
||||
llvm::SmallVector<mlir::OpFoldResult> strides;
|
||||
};
|
||||
|
||||
mlir::Value extractAxisSlice(
|
||||
mlir::PatternRewriter& rewriter, mlir::Location loc, mlir::Value source, int64_t axis, int64_t offset, int64_t size);
|
||||
|
||||
mlir::Value extractStaticSliceOrIdentity(mlir::OpBuilder& rewriter,
|
||||
mlir::Location loc,
|
||||
mlir::Value source,
|
||||
mlir::RankedTensorType resultType,
|
||||
llvm::ArrayRef<mlir::OpFoldResult> offsets,
|
||||
llvm::ArrayRef<mlir::OpFoldResult> sizes,
|
||||
llvm::ArrayRef<mlir::OpFoldResult> strides);
|
||||
|
||||
mlir::Value insertStaticSlice(mlir::PatternRewriter& rewriter,
|
||||
mlir::Location loc,
|
||||
mlir::Value source,
|
||||
mlir::Value dest,
|
||||
llvm::ArrayRef<mlir::OpFoldResult> offsets);
|
||||
|
||||
mlir::Value extractMixedSliceOrIdentity(mlir::OpBuilder &rewriter,
|
||||
mlir::Location loc,
|
||||
mlir::Value source,
|
||||
mlir::RankedTensorType resultType,
|
||||
const MixedSliceGeometry &geometry);
|
||||
|
||||
mlir::Value insertMixedSlice(mlir::OpBuilder &builder,
|
||||
mlir::Location loc,
|
||||
mlir::Value source,
|
||||
mlir::Value dest,
|
||||
const MixedSliceGeometry &geometry);
|
||||
|
||||
mlir::FailureOr<mlir::Value>
|
||||
addLeadingUnitTensorDimension(mlir::OpBuilder& builder, mlir::Location loc, mlir::Value value);
|
||||
|
||||
mlir::FailureOr<mlir::Value> removeLeadingUnitTensorDimension(
|
||||
mlir::OpBuilder& builder, mlir::Location loc, mlir::Value value, mlir::RankedTensorType resultType);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -1,8 +1,14 @@
|
||||
#include "mlir/Dialect/Linalg/IR/Linalg.h"
|
||||
#include "mlir/Dialect/Tensor/IR/Tensor.h"
|
||||
#include "mlir/IR/BuiltinAttributes.h"
|
||||
#include "mlir/IR/BuiltinTypes.h"
|
||||
|
||||
#include "llvm/ADT/SmallPtrSet.h"
|
||||
#include "llvm/ADT/SmallSet.h"
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/AddressAnalysis.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/SubviewUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/WeightUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialOps.hpp"
|
||||
@@ -19,29 +25,67 @@ void markWeightAlways(mlir::Operation* op) {
|
||||
|
||||
namespace {
|
||||
|
||||
template <typename MVMOpTy, typename VMMOpTy, typename ParentOpTy>
|
||||
bool hasMvmVmmWeightUse(ParentOpTy parentOp, unsigned weightIndex) {
|
||||
CompiledIndexExpr makeConstantExpr(int64_t constant) {
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = CompiledIndexExprNode::Kind::Constant;
|
||||
expr.constant = constant;
|
||||
return CompiledIndexExpr(std::make_shared<CompiledIndexExprNode>(std::move(expr)));
|
||||
}
|
||||
|
||||
CompiledIndexExpr makeBinaryExpr(CompiledIndexExprNode::Kind kind, CompiledIndexExpr lhs, CompiledIndexExpr rhs) {
|
||||
CompiledIndexExprNode expr;
|
||||
expr.kind = kind;
|
||||
expr.operands = {std::move(lhs), std::move(rhs)};
|
||||
return CompiledIndexExpr(std::make_shared<CompiledIndexExprNode>(std::move(expr)));
|
||||
}
|
||||
|
||||
CompiledIndexExpr addExpr(CompiledIndexExpr lhs, CompiledIndexExpr rhs) {
|
||||
return makeBinaryExpr(CompiledIndexExprNode::Kind::Add, std::move(lhs), std::move(rhs));
|
||||
}
|
||||
|
||||
CompiledIndexExpr mulExpr(CompiledIndexExpr lhs, int64_t rhs) {
|
||||
return makeBinaryExpr(CompiledIndexExprNode::Kind::Mul, std::move(lhs), makeConstantExpr(rhs));
|
||||
}
|
||||
|
||||
llvm::FailureOr<llvm::SmallVector<int64_t>> getStaticMemRefTypeStrides(mlir::MemRefType type) {
|
||||
llvm::SmallVector<int64_t> strides;
|
||||
int64_t offset = 0;
|
||||
if (failed(type.getStridesAndOffset(strides, offset)))
|
||||
return mlir::failure();
|
||||
if (llvm::is_contained(strides, mlir::ShapedType::kDynamic))
|
||||
return mlir::failure();
|
||||
return strides;
|
||||
}
|
||||
|
||||
template <typename VMMOpTy, typename ParentOpTy>
|
||||
bool hasVmmWeightUse(ParentOpTy parentOp, unsigned weightIndex) {
|
||||
auto weightArg = parentOp.getWeightArgument(weightIndex);
|
||||
if (!weightArg)
|
||||
return false;
|
||||
bool found = false;
|
||||
parentOp.walk([&](mlir::Operation* op) {
|
||||
if (auto mvmOp = mlir::dyn_cast<MVMOpTy>(op))
|
||||
found |= mvmOp.getWeightIndex() == weightIndex;
|
||||
else if (auto vmmOp = mlir::dyn_cast<VMMOpTy>(op))
|
||||
found |= vmmOp.getWeightIndex() == weightIndex;
|
||||
if (auto vmmOp = mlir::dyn_cast<VMMOpTy>(op))
|
||||
found |= vmmOp.getWeight() == *weightArg;
|
||||
});
|
||||
return found;
|
||||
}
|
||||
|
||||
template <typename MVMOpTy, typename VMMOpTy, typename ParentOpTy>
|
||||
void walkMvmVmmWeightUses(ParentOpTy parentOp, llvm::function_ref<void(mlir::OpOperand&)> callback) {
|
||||
template <typename VMMOpTy, typename ParentOpTy>
|
||||
void walkVmmWeightUses(ParentOpTy parentOp, llvm::function_ref<void(mlir::OpOperand&)> callback) {
|
||||
auto weights = parentOp.getWeights();
|
||||
llvm::SmallSet<unsigned, 8> visited;
|
||||
auto walkWeightIndex = [&](unsigned weightIndex) {
|
||||
if (weightIndex < weights.size() && visited.insert(weightIndex).second)
|
||||
callback(parentOp->getOpOperand(weightIndex));
|
||||
auto walkWeight = [&](mlir::Value weight) {
|
||||
for (unsigned weightIndex = 0; weightIndex < weights.size(); ++weightIndex) {
|
||||
auto weightArg = parentOp.getWeightArgument(weightIndex);
|
||||
if (!weightArg || *weightArg != weight)
|
||||
continue;
|
||||
if (visited.insert(weightIndex).second)
|
||||
callback(parentOp->getOpOperand(weightIndex));
|
||||
break;
|
||||
}
|
||||
};
|
||||
|
||||
parentOp.walk([&](MVMOpTy op) { walkWeightIndex(op.getWeightIndex()); });
|
||||
parentOp.walk([&](VMMOpTy op) { walkWeightIndex(op.getWeightIndex()); });
|
||||
parentOp.walk([&](VMMOpTy op) { walkWeight(op.getWeight()); });
|
||||
}
|
||||
|
||||
} // namespace
|
||||
@@ -54,7 +98,7 @@ bool isSpatialMvmVmmWeightUse(mlir::OpOperand& use) {
|
||||
if (!computeOp || operandIndex >= computeOp.getWeights().size())
|
||||
return false;
|
||||
|
||||
return hasMvmVmmWeightUse<spatial::SpatWeightedMVMOp, spatial::SpatWeightedVMMOp>(computeOp, operandIndex);
|
||||
return hasVmmWeightUse<spatial::SpatVMMOp>(computeOp, operandIndex);
|
||||
}
|
||||
|
||||
bool hasOnlySpatialMvmVmmWeightUses(mlir::Value value) {
|
||||
@@ -76,8 +120,8 @@ bool hasOnlySpatialMvmVmmWeightUses(mlir::Value value) {
|
||||
return expandShapeOp.getSrc() == currentValue && self(expandShapeOp.getResult(), self);
|
||||
if (auto collapseShapeOp = mlir::dyn_cast<mlir::tensor::CollapseShapeOp>(user))
|
||||
return collapseShapeOp.getSrc() == currentValue && self(collapseShapeOp.getResult(), self);
|
||||
if (auto transposeOp = mlir::dyn_cast<mlir::ONNXTransposeOp>(user))
|
||||
return transposeOp.getData() == currentValue && self(transposeOp.getResult(), self);
|
||||
if (auto transposeOp = mlir::dyn_cast<mlir::linalg::TransposeOp>(user))
|
||||
return transposeOp.getInput() == currentValue && self(transposeOp.getResult()[0], self);
|
||||
|
||||
return false;
|
||||
});
|
||||
@@ -88,14 +132,185 @@ bool hasOnlySpatialMvmVmmWeightUses(mlir::Value value) {
|
||||
|
||||
void walkPimMvmVmmWeightUses(mlir::Operation* root, llvm::function_ref<void(mlir::OpOperand&)> callback) {
|
||||
assert(root && "expected valid root op");
|
||||
root->walk([&](pim::PimCoreOp coreOp) { walkMvmVmmWeightUses<pim::PimMVMOp, pim::PimVMMOp>(coreOp, callback); });
|
||||
root->walk([&](pim::PimCoreOp coreOp) {
|
||||
coreOp.walk([&](pim::PimVMMOp vmmOp) {
|
||||
if (auto weightIndex = resolveWeightIndex(coreOp.getOperation(), vmmOp.getWeight()))
|
||||
callback(coreOp->getOpOperand(*weightIndex));
|
||||
});
|
||||
});
|
||||
root->walk([&](pim::PimCoreBatchOp coreBatchOp) {
|
||||
auto weights = coreBatchOp.getWeights();
|
||||
for (auto weight : weights)
|
||||
for (mlir::OpOperand& use : weight.getUses())
|
||||
if (use.getOwner() == coreBatchOp.getOperation())
|
||||
callback(use);
|
||||
coreBatchOp.walk([&](pim::PimVMMOp vmmOp) {
|
||||
if (auto weightIndex = resolveWeightIndex(coreBatchOp.getOperation(), vmmOp.getWeight()))
|
||||
callback(coreBatchOp->getOpOperand(*weightIndex));
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
std::optional<unsigned> resolveWeightIndex(mlir::Operation* coreLikeOp, mlir::Value weight) {
|
||||
weight = stripMemRefAddressingOps(weight);
|
||||
|
||||
if (auto coreOp = mlir::dyn_cast_or_null<pim::PimCoreOp>(coreLikeOp)) {
|
||||
for (unsigned weightIndex = 0; weightIndex < coreOp.getWeights().size(); ++weightIndex)
|
||||
if (coreOp.getWeightArgument(weightIndex) == weight)
|
||||
return weightIndex;
|
||||
return std::nullopt;
|
||||
}
|
||||
|
||||
if (auto coreBatchOp = mlir::dyn_cast_or_null<pim::PimCoreBatchOp>(coreLikeOp)) {
|
||||
for (unsigned weightIndex = 0; weightIndex < coreBatchOp.getWeights().size(); ++weightIndex)
|
||||
if (coreBatchOp.getWeightArgument(weightIndex) == weight)
|
||||
return weightIndex;
|
||||
return std::nullopt;
|
||||
}
|
||||
|
||||
return std::nullopt;
|
||||
}
|
||||
|
||||
llvm::FailureOr<ResolvedWeightView>
|
||||
resolveWeightView(mlir::Operation* coreLikeOp, mlir::Value weight, const StaticValueKnowledge& knowledge) {
|
||||
llvm::SmallVector<mlir::Operation*> viewOps;
|
||||
mlir::Value current = weight;
|
||||
|
||||
while (true) {
|
||||
if (mlir::Value directAlias = knowledge.aliases.lookup(current); directAlias && directAlias != current) {
|
||||
current = directAlias;
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto defOp = current.getDefiningOp()) {
|
||||
if (auto getGlobalOp = mlir::dyn_cast<mlir::memref::GetGlobalOp>(defOp)) {
|
||||
auto moduleOp = coreLikeOp ? coreLikeOp->getParentOfType<mlir::ModuleOp>() : mlir::ModuleOp {};
|
||||
auto globalOp = lookupGlobalForGetGlobal(moduleOp, getGlobalOp);
|
||||
if (!globalOp || !globalOp.getInitialValue())
|
||||
return mlir::failure();
|
||||
|
||||
auto denseAttr = mlir::dyn_cast<mlir::DenseElementsAttr>(*globalOp.getInitialValue());
|
||||
if (!denseAttr)
|
||||
return mlir::failure();
|
||||
|
||||
ResolvedWeightView view;
|
||||
view.globalOp = globalOp;
|
||||
view.shape.assign(denseAttr.getType().getShape().begin(), denseAttr.getType().getShape().end());
|
||||
view.strides = computeRowMajorStrides(view.shape);
|
||||
|
||||
CompiledIndexExpr offsetExpr = makeConstantExpr(0);
|
||||
for (mlir::Operation* viewOp : llvm::reverse(viewOps)) {
|
||||
if (auto subview = mlir::dyn_cast<mlir::memref::SubViewOp>(viewOp)) {
|
||||
for (auto [offset, stride, sourceStride] :
|
||||
llvm::zip_equal(subview.getMixedOffsets(), subview.getStaticStrides(), view.strides)) {
|
||||
CompiledIndexExpr offsetValue = makeConstantExpr(0);
|
||||
if (auto attr = mlir::dyn_cast<mlir::Attribute>(offset)) {
|
||||
auto intAttr = mlir::dyn_cast<mlir::IntegerAttr>(attr);
|
||||
if (!intAttr)
|
||||
return mlir::failure();
|
||||
offsetValue = makeConstantExpr(intAttr.getInt());
|
||||
}
|
||||
else if (auto value = mlir::dyn_cast<mlir::Value>(offset)) {
|
||||
auto compiledOffset = compileIndexExpr(value);
|
||||
if (failed(compiledOffset))
|
||||
return mlir::failure();
|
||||
offsetValue = *compiledOffset;
|
||||
}
|
||||
else {
|
||||
return mlir::failure();
|
||||
}
|
||||
offsetExpr = addExpr(std::move(offsetExpr), mulExpr(std::move(offsetValue), sourceStride));
|
||||
}
|
||||
auto resultType = mlir::cast<mlir::MemRefType>(subview.getResult().getType());
|
||||
auto resultStrides = getStaticMemRefTypeStrides(resultType);
|
||||
if (failed(resultStrides))
|
||||
return mlir::failure();
|
||||
view.shape.assign(resultType.getShape().begin(), resultType.getShape().end());
|
||||
view.strides = std::move(*resultStrides);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto collapse = mlir::dyn_cast<mlir::memref::CollapseShapeOp>(viewOp)) {
|
||||
auto resultType = mlir::cast<mlir::MemRefType>(collapse.getResult().getType());
|
||||
auto resultStrides = getStaticMemRefTypeStrides(resultType);
|
||||
if (failed(resultStrides))
|
||||
return mlir::failure();
|
||||
view.shape.assign(resultType.getShape().begin(), resultType.getShape().end());
|
||||
view.strides = std::move(*resultStrides);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto expand = mlir::dyn_cast<mlir::memref::ExpandShapeOp>(viewOp)) {
|
||||
auto resultType = mlir::cast<mlir::MemRefType>(expand.getResult().getType());
|
||||
auto resultStrides = getStaticMemRefTypeStrides(resultType);
|
||||
if (failed(resultStrides))
|
||||
return mlir::failure();
|
||||
view.shape.assign(resultType.getShape().begin(), resultType.getShape().end());
|
||||
view.strides = std::move(*resultStrides);
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto castOp = mlir::dyn_cast<mlir::memref::CastOp>(viewOp)) {
|
||||
auto resultType = mlir::cast<mlir::MemRefType>(castOp.getResult().getType());
|
||||
auto resultStrides = getStaticMemRefTypeStrides(resultType);
|
||||
if (failed(resultStrides))
|
||||
return mlir::failure();
|
||||
view.shape.assign(resultType.getShape().begin(), resultType.getShape().end());
|
||||
view.strides = std::move(*resultStrides);
|
||||
continue;
|
||||
}
|
||||
|
||||
return mlir::failure();
|
||||
}
|
||||
|
||||
auto resolvedOffset = offsetExpr.evaluate(knowledge);
|
||||
if (failed(resolvedOffset))
|
||||
return mlir::failure();
|
||||
view.offset = *resolvedOffset;
|
||||
return view;
|
||||
}
|
||||
|
||||
if (auto subview = mlir::dyn_cast<mlir::memref::SubViewOp>(defOp)) {
|
||||
viewOps.push_back(defOp);
|
||||
current = subview.getSource();
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto collapse = mlir::dyn_cast<mlir::memref::CollapseShapeOp>(defOp)) {
|
||||
viewOps.push_back(defOp);
|
||||
current = collapse.getSrc();
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto expand = mlir::dyn_cast<mlir::memref::ExpandShapeOp>(defOp)) {
|
||||
viewOps.push_back(defOp);
|
||||
current = expand.getSrc();
|
||||
continue;
|
||||
}
|
||||
|
||||
if (auto castOp = mlir::dyn_cast<mlir::memref::CastOp>(defOp)) {
|
||||
viewOps.push_back(defOp);
|
||||
current = castOp.getSource();
|
||||
continue;
|
||||
}
|
||||
|
||||
return mlir::failure();
|
||||
}
|
||||
|
||||
if (mlir::Value loopAlias = resolveLoopCarriedAlias(current, knowledge); loopAlias && loopAlias != current) {
|
||||
current = loopAlias;
|
||||
continue;
|
||||
}
|
||||
|
||||
auto weightIndex = resolveWeightIndex(coreLikeOp, current);
|
||||
if (!weightIndex)
|
||||
return mlir::failure();
|
||||
|
||||
if (auto coreOp = mlir::dyn_cast_or_null<pim::PimCoreOp>(coreLikeOp)) {
|
||||
current = coreOp.getWeights()[*weightIndex];
|
||||
continue;
|
||||
}
|
||||
if (auto coreBatchOp = mlir::dyn_cast_or_null<pim::PimCoreBatchOp>(coreLikeOp)) {
|
||||
current = coreBatchOp.getWeights()[*weightIndex];
|
||||
continue;
|
||||
}
|
||||
return mlir::failure();
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -1,15 +1,34 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/IR/Operation.h"
|
||||
#include "mlir/Dialect/MemRef/IR/MemRef.h"
|
||||
#include "mlir/IR/BuiltinOps.h"
|
||||
#include "mlir/IR/Value.h"
|
||||
|
||||
#include "llvm/ADT/STLExtras.h"
|
||||
#include "llvm/ADT/STLFunctionalExtras.h"
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
|
||||
#include <optional>
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/AddressAnalysis.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp"
|
||||
|
||||
inline constexpr llvm::StringRef PimWeightAlwaysAttrName = "weightAlways";
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
struct ResolvedWeightView {
|
||||
mlir::memref::GlobalOp globalOp;
|
||||
llvm::SmallVector<int64_t> shape;
|
||||
llvm::SmallVector<int64_t> strides;
|
||||
int64_t offset = 0;
|
||||
|
||||
bool operator==(const ResolvedWeightView& other) const {
|
||||
return globalOp == other.globalOp && shape == other.shape && strides == other.strides && offset == other.offset;
|
||||
}
|
||||
};
|
||||
|
||||
bool hasWeightAlways(mlir::Operation* op);
|
||||
|
||||
/// Tags an op as producing a value that should stay materialized as a reusable
|
||||
@@ -26,4 +45,20 @@ bool hasOnlySpatialMvmVmmWeightUses(mlir::Value value);
|
||||
/// passes can identify globals that must remain weight-backed.
|
||||
void walkPimMvmVmmWeightUses(mlir::Operation* root, llvm::function_ref<void(mlir::OpOperand&)> callback);
|
||||
|
||||
std::optional<unsigned> resolveWeightIndex(mlir::Operation* coreLikeOp, mlir::Value weight);
|
||||
llvm::FailureOr<ResolvedWeightView>
|
||||
resolveWeightView(mlir::Operation* coreLikeOp, mlir::Value weight, const StaticValueKnowledge& knowledge = {});
|
||||
|
||||
template <typename CoreLikeOpTy>
|
||||
llvm::SmallVector<unsigned, 8> getUsedWeightIndices(CoreLikeOpTy coreLikeOp) {
|
||||
llvm::SmallVector<unsigned, 8> indices;
|
||||
coreLikeOp.walk([&](pim::PimVMMOp vmmOp) {
|
||||
auto weightIndex = resolveWeightIndex(coreLikeOp.getOperation(), vmmOp.getWeight());
|
||||
if (weightIndex && !llvm::is_contained(indices, *weightIndex))
|
||||
indices.push_back(*weightIndex);
|
||||
});
|
||||
llvm::sort(indices);
|
||||
return indices;
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -1,315 +0,0 @@
|
||||
#pragma once
|
||||
|
||||
#include "llvm/ADT/STLExtras.h"
|
||||
#include "llvm/ADT/ilist_node.h"
|
||||
#include "llvm/ADT/simple_ilist.h"
|
||||
|
||||
#include <cassert>
|
||||
#include <iterator>
|
||||
#include <limits>
|
||||
#include <type_traits>
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
template <typename NodeT>
|
||||
class LabeledList;
|
||||
|
||||
template <typename NodeT>
|
||||
class LabeledListNode : public llvm::ilist_node<NodeT> {
|
||||
friend class LabeledList<NodeT>;
|
||||
|
||||
public:
|
||||
using Label = uint64_t;
|
||||
|
||||
LabeledListNode() = default;
|
||||
LabeledListNode(const LabeledListNode&) = delete;
|
||||
LabeledListNode(LabeledListNode&&) = default;
|
||||
LabeledListNode& operator=(LabeledListNode&&) = delete;
|
||||
|
||||
~LabeledListNode() { assert(owner_ == nullptr && "destroying a linked LabeledListNode"); }
|
||||
|
||||
bool isLinked() const { return owner_ != nullptr; }
|
||||
Label getOrderLabel() const { return label; }
|
||||
|
||||
friend bool operator<(const LabeledListNode& lft, const LabeledListNode& rgt) { return lft.label < rgt.label; }
|
||||
|
||||
private:
|
||||
const void* owner_ = nullptr;
|
||||
Label label = 0;
|
||||
};
|
||||
|
||||
template <typename NodeT>
|
||||
class LabeledList {
|
||||
|
||||
using Label = typename NodeT::Label;
|
||||
|
||||
static constexpr Label kLowerSentinel = 0;
|
||||
static constexpr Label kUpperSentinel = std::numeric_limits<Label>::max();
|
||||
static constexpr Label kRelabelGap = 2;
|
||||
|
||||
public:
|
||||
using List = llvm::simple_ilist<NodeT>;
|
||||
using Iterator = typename List::iterator;
|
||||
using RIterator = typename List::reverse_iterator;
|
||||
using ConstIterator = typename List::const_iterator;
|
||||
|
||||
LabeledList() = default;
|
||||
LabeledList(const LabeledList&) = delete;
|
||||
LabeledList& operator=(const LabeledList&) = delete;
|
||||
LabeledList(LabeledList&&) = delete;
|
||||
LabeledList& operator=(LabeledList&&) = delete;
|
||||
|
||||
~LabeledList() { clear(); }
|
||||
|
||||
bool empty() const { return size_ == 0; }
|
||||
size_t size() const { return size_; }
|
||||
|
||||
NodeT* front() { return empty() ? nullptr : &nodes_.front(); }
|
||||
const NodeT* front() const { return empty() ? nullptr : &nodes_.front(); }
|
||||
|
||||
NodeT* back() { return empty() ? nullptr : &nodes_.back(); }
|
||||
const NodeT* back() const { return empty() ? nullptr : &nodes_.back(); }
|
||||
|
||||
static NodeT* previous(NodeT* node) {
|
||||
if (!node || !owner(node))
|
||||
return nullptr;
|
||||
auto* list = owner(node);
|
||||
auto it = node->getIterator();
|
||||
if (it == list->nodes_.begin())
|
||||
return nullptr;
|
||||
return &*std::prev(it);
|
||||
}
|
||||
|
||||
static const NodeT* previous(const NodeT* node) {
|
||||
if (!node || !owner(node))
|
||||
return nullptr;
|
||||
const auto* list = owner(node);
|
||||
auto it = const_cast<NodeT*>(node)->getIterator();
|
||||
if (it == list->nodes_.begin())
|
||||
return nullptr;
|
||||
return &*std::prev(it);
|
||||
}
|
||||
|
||||
static NodeT* next(NodeT* node) {
|
||||
if (!node || !owner(node))
|
||||
return nullptr;
|
||||
auto* list = owner(node);
|
||||
auto it = std::next(node->getIterator());
|
||||
if (it == list->nodes_.end())
|
||||
return nullptr;
|
||||
return &*it;
|
||||
}
|
||||
|
||||
static const NodeT* next(const NodeT* node) {
|
||||
if (!node || !owner(node))
|
||||
return nullptr;
|
||||
const auto* list = owner(node);
|
||||
auto it = std::next(const_cast<NodeT*>(node)->getIterator());
|
||||
if (it == list->nodes_.end())
|
||||
return nullptr;
|
||||
return &*it;
|
||||
}
|
||||
|
||||
bool contains(const NodeT* node) const { return node && node->owner_ == this; }
|
||||
|
||||
Label getOrderLabel(const NodeT* node) const {
|
||||
assert(contains(node) && "node must belong to this list");
|
||||
return node->label;
|
||||
}
|
||||
|
||||
bool comesBefore(const NodeT* lhs, const NodeT* rhs) const {
|
||||
assert(contains(lhs) && contains(rhs) && "nodes must belong to this list");
|
||||
return lhs->label < rhs->label;
|
||||
}
|
||||
|
||||
void pushFront(NodeT* node) { insertBefore(front(), node); }
|
||||
|
||||
void pushBack(NodeT* node) { insertBefore(nullptr, node); }
|
||||
|
||||
void insertBefore(NodeT* nextNode, NodeT* node) {
|
||||
assert(node && "cannot insert a null node");
|
||||
assert(!node->owner_ && "node is already linked");
|
||||
assert(nextNode == nullptr || contains(nextNode));
|
||||
|
||||
Iterator nextIt = nextNode ? getIteratorFor(nextNode) : nodes_.end();
|
||||
nodes_.insert(nextIt, *node);
|
||||
node->owner_ = this;
|
||||
++size_;
|
||||
assignLabel(getIteratorFor(node));
|
||||
}
|
||||
|
||||
void insertAfter(NodeT* prevNode, NodeT* node) {
|
||||
assert(prevNode == nullptr || contains(prevNode));
|
||||
if (prevNode == nullptr)
|
||||
insertBefore(front(), node);
|
||||
else
|
||||
insertBefore(next(prevNode), node);
|
||||
}
|
||||
|
||||
void remove(NodeT* node) {
|
||||
assert(contains(node) && "node must belong to this list");
|
||||
nodes_.remove(*node);
|
||||
node->owner_ = nullptr;
|
||||
node->label = 0;
|
||||
--size_;
|
||||
}
|
||||
|
||||
void moveBefore(NodeT* node, NodeT* nextNode) {
|
||||
assert(contains(node) && "node must belong to this list");
|
||||
assert(nextNode == nullptr || contains(nextNode));
|
||||
|
||||
Iterator nodeIt = getIteratorFor(node);
|
||||
Iterator nextIt = nextNode ? getIteratorFor(nextNode) : nodes_.end();
|
||||
if (nodeIt == nextIt || std::next(nodeIt) == nextIt)
|
||||
return;
|
||||
|
||||
nodes_.splice(nextIt, nodes_, nodeIt);
|
||||
assignLabel(getIteratorFor(node));
|
||||
}
|
||||
|
||||
void moveAfter(NodeT* node, NodeT* prevNode) {
|
||||
assert(contains(node) && "node must belong to this list");
|
||||
assert(prevNode == nullptr || contains(prevNode));
|
||||
|
||||
Iterator nextIt = prevNode ? std::next(getIteratorFor(prevNode)) : nodes_.begin();
|
||||
if (getIteratorFor(node) == nextIt)
|
||||
return;
|
||||
moveBefore(node, nextIt == nodes_.end() ? nullptr : &*nextIt);
|
||||
}
|
||||
|
||||
void clear() {
|
||||
while (!nodes_.empty()) {
|
||||
NodeT* node = &nodes_.front();
|
||||
node->owner_ = nullptr;
|
||||
node->label = 0;
|
||||
nodes_.remove(*node);
|
||||
}
|
||||
size_ = 0;
|
||||
}
|
||||
|
||||
Iterator begin() { return nodes_.begin(); }
|
||||
Iterator end() { return nodes_.end(); }
|
||||
|
||||
RIterator rbegin() { return nodes_.rbegin(); }
|
||||
RIterator rend() { return nodes_.rend(); }
|
||||
|
||||
private:
|
||||
static const LabeledList* owner(const NodeT* node) { return static_cast<const LabeledList*>(node->owner_); }
|
||||
static LabeledList* owner(NodeT* node) { return static_cast<LabeledList*>(const_cast<void*>(node->owner_)); }
|
||||
|
||||
static Label lowerLabel(const NodeT* node) { return node ? node->label : kLowerSentinel; }
|
||||
static Label upperLabel(const NodeT* node) { return node ? node->label : kUpperSentinel; }
|
||||
|
||||
static Label labelGap(Label lower, Label upper) {
|
||||
assert(lower < upper && "labels must be strictly ordered");
|
||||
return upper - lower;
|
||||
}
|
||||
|
||||
static bool hasMidpoint(Label lower, Label upper) { return labelGap(lower, upper) > 1; }
|
||||
|
||||
static bool hasRelabelSlack(Label lower, Label upper, size_t nodeCount) {
|
||||
Label gap = labelGap(lower, upper);
|
||||
return gap / static_cast<Label>(nodeCount + 1) >= kRelabelGap;
|
||||
}
|
||||
|
||||
Iterator getIteratorFor(NodeT* node) { return node->getIterator(); }
|
||||
ConstIterator getiteratorFor(const NodeT* node) const { return node->getIterator(); }
|
||||
|
||||
NodeT* previousNode(Iterator it) {
|
||||
if (it == nodes_.begin())
|
||||
return nullptr;
|
||||
return &*std::prev(it);
|
||||
}
|
||||
|
||||
const NodeT* previousNode(ConstIterator it) const {
|
||||
if (it == nodes_.begin())
|
||||
return nullptr;
|
||||
return &*std::prev(it);
|
||||
}
|
||||
|
||||
NodeT* nextNode(Iterator it) {
|
||||
++it;
|
||||
if (it == nodes_.end())
|
||||
return nullptr;
|
||||
return &*it;
|
||||
}
|
||||
|
||||
const NodeT* nextNode(ConstIterator it) const {
|
||||
++it;
|
||||
if (it == nodes_.end())
|
||||
return nullptr;
|
||||
return &*it;
|
||||
}
|
||||
|
||||
void assignLabel(Iterator it) {
|
||||
Label lower = lowerLabel(previousNode(it));
|
||||
Label upper = upperLabel(nextNode(it));
|
||||
if (hasMidpoint(lower, upper)) {
|
||||
(*it).label = lower + static_cast<Label>(labelGap(lower, upper) / 2);
|
||||
return;
|
||||
}
|
||||
|
||||
relabelAround(it);
|
||||
}
|
||||
|
||||
void relabelAround(Iterator center) {
|
||||
size_t targetCount = 1;
|
||||
while (true) {
|
||||
Iterator left = center;
|
||||
Iterator right = center;
|
||||
size_t actualCount = 1;
|
||||
expandWindow(center, targetCount, left, right, actualCount);
|
||||
|
||||
Label lower = lowerLabel(previousNode(left));
|
||||
Label upper = upperLabel(nextNode(right));
|
||||
if (hasRelabelSlack(lower, upper, actualCount)) {
|
||||
relabelWindow(left, actualCount, lower, upper);
|
||||
return;
|
||||
}
|
||||
|
||||
if (left == nodes_.begin() && nextNode(right) == nullptr) {
|
||||
assert(hasRelabelSlack(lower, upper, actualCount) && "label space exhausted");
|
||||
relabelWindow(left, actualCount, lower, upper);
|
||||
return;
|
||||
}
|
||||
|
||||
targetCount *= 2;
|
||||
}
|
||||
}
|
||||
|
||||
void expandWindow(Iterator center, size_t targetCount, Iterator& left, Iterator& right, size_t& actualCount) {
|
||||
left = center;
|
||||
right = center;
|
||||
actualCount = 1;
|
||||
|
||||
while (actualCount < targetCount && (left != nodes_.begin() || nextNode(right) != nullptr)) {
|
||||
if (left != nodes_.begin()) {
|
||||
--left;
|
||||
++actualCount;
|
||||
if (actualCount == targetCount)
|
||||
break;
|
||||
}
|
||||
if (nextNode(right) != nullptr) {
|
||||
++right;
|
||||
++actualCount;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void relabelWindow(Iterator left, size_t nodeCount, Label lower, Label upper) {
|
||||
assert(nodeCount > 0 && "relabel window must not be empty");
|
||||
Label step = labelGap(lower, upper) / static_cast<Label>(nodeCount + 1);
|
||||
assert(step >= 1 && "relabel step must be positive");
|
||||
|
||||
Iterator it = left;
|
||||
for (size_t index = 1; index <= nodeCount; ++index) {
|
||||
(*it).label = lower + step * index;
|
||||
++it;
|
||||
}
|
||||
}
|
||||
|
||||
List nodes_;
|
||||
size_t size_ = 0;
|
||||
};
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -12,17 +12,35 @@
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/AddressAnalysis.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/ConstantUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/CoreBlockUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/EntryPointUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/IndexingUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/WeightUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/Support/DebugDump.hpp"
|
||||
#include "src/Accelerators/PIM/Common/Support/FileSystemUtils.hpp"
|
||||
#include "src/Compiler/CompilerOptions.hpp"
|
||||
|
||||
#include <cstdint>
|
||||
#include <array>
|
||||
#include <limits>
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
inline constexpr llvm::StringLiteral kCoreIdAttrName = "coreId";
|
||||
inline constexpr llvm::StringLiteral kCoreIdsAttrName = "coreIds";
|
||||
inline constexpr llvm::StringLiteral kLocalMemoryAddressAttrName = "pim.local_memory_address";
|
||||
inline constexpr llvm::StringLiteral kLocalMemorySizeAttrName = "pim.local_memory_size";
|
||||
inline constexpr llvm::StringLiteral kPipelineHostBufferBytesAttrName = "pim.pipeline_host_buffer_bytes";
|
||||
inline constexpr llvm::StringLiteral kPipelineHostBufferName = "pim_pipeline_channels";
|
||||
inline constexpr size_t kPimEventRegisterCount = 32;
|
||||
inline constexpr std::array<llvm::StringLiteral, 4> kRemovedLocalMemoryPlanAttrNames = {
|
||||
"pim.local_memory_slot",
|
||||
"pim.local_memory_slot_size",
|
||||
"pim.local_memory_fallback_count",
|
||||
"pim.local_memory_nested_single_use_count",
|
||||
};
|
||||
inline constexpr size_t kPimLocalMemoryAddressLimit = static_cast<size_t>(std::numeric_limits<int32_t>::max());
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -0,0 +1,222 @@
|
||||
#include "llvm/Support/ErrorHandling.h"
|
||||
#include "llvm/Support/raw_ostream.h"
|
||||
|
||||
#include "CheckedArithmetic.hpp"
|
||||
#include "src/Accelerators/PIM/Common/PimCommon.hpp"
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir::pim {
|
||||
|
||||
namespace {
|
||||
|
||||
static void emitCrashMessage(llvm::StringRef fieldName, llvm::StringRef message) {
|
||||
llvm::errs() << "PIM " << fieldName << " " << message << "\n";
|
||||
}
|
||||
|
||||
template <typename To, typename From>
|
||||
static FailureOr<To> checkedCastAtLocation(From value, Location loc, llvm::StringRef fieldName) {
|
||||
static_assert(std::is_integral_v<To> && std::is_integral_v<From>, "checkedCastAtLocation requires integral types");
|
||||
|
||||
using ToLimits = std::numeric_limits<To>;
|
||||
|
||||
if constexpr (std::is_signed_v<From> == std::is_signed_v<To>) {
|
||||
if (value < static_cast<From>(ToLimits::min()) || value > static_cast<From>(ToLimits::max())) {
|
||||
emitCheckedArithmeticError(loc, fieldName, "is outside representable range");
|
||||
return failure();
|
||||
}
|
||||
}
|
||||
else if constexpr (std::is_signed_v<From>) {
|
||||
using UnsignedFrom = std::make_unsigned_t<From>;
|
||||
using UnsignedTo = std::make_unsigned_t<To>;
|
||||
if (value < 0 || static_cast<UnsignedFrom>(value) > static_cast<UnsignedTo>(ToLimits::max())) {
|
||||
emitCheckedArithmeticError(loc, fieldName, "is outside representable range");
|
||||
return failure();
|
||||
}
|
||||
}
|
||||
else {
|
||||
using UnsignedFrom = std::make_unsigned_t<From>;
|
||||
using UnsignedTo = std::conditional_t<std::is_signed_v<To>, std::make_unsigned_t<To>, To>;
|
||||
if (static_cast<UnsignedFrom>(value) > static_cast<UnsignedTo>(ToLimits::max())) {
|
||||
emitCheckedArithmeticError(loc, fieldName, "is outside representable range");
|
||||
return failure();
|
||||
}
|
||||
}
|
||||
|
||||
return static_cast<To>(value);
|
||||
}
|
||||
|
||||
template <typename UInt>
|
||||
FailureOr<UInt> checkedMulAtLocation(UInt lhs, UInt rhs, Location loc, llvm::StringRef fieldName) {
|
||||
static_assert(std::is_integral_v<UInt> && std::is_unsigned_v<UInt>,
|
||||
"checkedMulAtLocation requires unsigned integral types");
|
||||
if (lhs != 0 && rhs > std::numeric_limits<UInt>::max() / lhs) {
|
||||
emitCheckedArithmeticError(loc, fieldName, "multiplication overflow");
|
||||
return failure();
|
||||
}
|
||||
return lhs * rhs;
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
InFlightDiagnostic emitCheckedArithmeticError(Operation* anchor, llvm::StringRef fieldName, llvm::StringRef message) {
|
||||
assert(anchor && "expected arithmetic diagnostics to have an anchor op");
|
||||
return anchor->emitOpError() << fieldName << " " << message;
|
||||
}
|
||||
|
||||
InFlightDiagnostic emitCheckedArithmeticError(Location loc, llvm::StringRef fieldName, llvm::StringRef message) {
|
||||
return emitError(loc) << "PIM " << fieldName << " " << message;
|
||||
}
|
||||
|
||||
FailureOr<int32_t> checkedI32(int64_t value, Operation* anchor, llvm::StringRef fieldName) {
|
||||
return checkedCast<int32_t>(value, anchor, fieldName);
|
||||
}
|
||||
|
||||
FailureOr<int32_t> checkedI32(uint64_t value, Operation* anchor, llvm::StringRef fieldName) {
|
||||
return checkedCast<int32_t>(value, anchor, fieldName);
|
||||
}
|
||||
|
||||
FailureOr<uint8_t> checkedU8(uint64_t value, Operation* anchor, llvm::StringRef fieldName) {
|
||||
return checkedCast<uint8_t>(value, anchor, fieldName);
|
||||
}
|
||||
|
||||
FailureOr<size_t> checkedSize(int64_t value, Operation* anchor, llvm::StringRef fieldName) {
|
||||
return checkedCast<size_t>(value, anchor, fieldName);
|
||||
}
|
||||
|
||||
FailureOr<IntegerAttr>
|
||||
getCheckedI32Attr(Builder& builder, Operation* anchor, int64_t value, llvm::StringRef fieldName) {
|
||||
assert(anchor && "checked op-based attrs require a non-null diagnostic anchor");
|
||||
auto checkedValue = checkedI32(value, anchor, fieldName);
|
||||
if (failed(checkedValue))
|
||||
return failure();
|
||||
return builder.getI32IntegerAttr(*checkedValue);
|
||||
}
|
||||
|
||||
FailureOr<IntegerAttr>
|
||||
getCheckedI32Attr(Builder& builder, Operation* anchor, uint64_t value, llvm::StringRef fieldName) {
|
||||
assert(anchor && "checked op-based attrs require a non-null diagnostic anchor");
|
||||
auto checkedValue = checkedI32(value, anchor, fieldName);
|
||||
if (failed(checkedValue))
|
||||
return failure();
|
||||
return builder.getI32IntegerAttr(*checkedValue);
|
||||
}
|
||||
|
||||
FailureOr<IntegerAttr> getCheckedI32Attr(Builder& builder, Location loc, int64_t value, llvm::StringRef fieldName) {
|
||||
auto checkedValue = checkedCastAtLocation<int32_t>(value, loc, fieldName);
|
||||
if (failed(checkedValue))
|
||||
return failure();
|
||||
return builder.getI32IntegerAttr(*checkedValue);
|
||||
}
|
||||
|
||||
FailureOr<IntegerAttr> getCheckedI32Attr(Builder& builder, Location loc, uint64_t value, llvm::StringRef fieldName) {
|
||||
auto checkedValue = checkedCastAtLocation<int32_t>(value, loc, fieldName);
|
||||
if (failed(checkedValue))
|
||||
return failure();
|
||||
return builder.getI32IntegerAttr(*checkedValue);
|
||||
}
|
||||
|
||||
FailureOr<uint64_t> getCheckedShapedTypeSizeInBytes(ShapedType type, Operation* anchor, llvm::StringRef fieldName) {
|
||||
assert(anchor && "checked op-based size helpers require a non-null diagnostic anchor");
|
||||
if (!type.hasStaticShape()) {
|
||||
emitCheckedArithmeticError(anchor, fieldName, "requires static shaped type");
|
||||
return failure();
|
||||
}
|
||||
if (!hasByteSizedElementType(type.getElementType())) {
|
||||
emitCheckedArithmeticError(anchor, fieldName, "requires byte-sized element type");
|
||||
return failure();
|
||||
}
|
||||
|
||||
uint64_t elements = 1;
|
||||
for (int64_t dim : type.getShape()) {
|
||||
if (dim < 0) {
|
||||
emitCheckedArithmeticError(anchor, fieldName, "requires nonnegative dimensions");
|
||||
return failure();
|
||||
}
|
||||
|
||||
auto nextElements = checkedMul(elements, static_cast<uint64_t>(dim), anchor, fieldName);
|
||||
if (failed(nextElements))
|
||||
return failure();
|
||||
elements = *nextElements;
|
||||
}
|
||||
|
||||
return checkedMul(
|
||||
elements, static_cast<uint64_t>(getElementTypeSizeInBytes(type.getElementType())), anchor, fieldName);
|
||||
}
|
||||
|
||||
FailureOr<uint64_t> getCheckedShapedTypeSizeInBytes(ShapedType type, Location loc, llvm::StringRef fieldName) {
|
||||
if (!type.hasStaticShape()) {
|
||||
emitCheckedArithmeticError(loc, fieldName, "requires static shaped type");
|
||||
return failure();
|
||||
}
|
||||
if (!hasByteSizedElementType(type.getElementType())) {
|
||||
emitCheckedArithmeticError(loc, fieldName, "requires byte-sized element type");
|
||||
return failure();
|
||||
}
|
||||
|
||||
uint64_t elements = 1;
|
||||
for (int64_t dim : type.getShape()) {
|
||||
if (dim < 0) {
|
||||
emitCheckedArithmeticError(loc, fieldName, "requires nonnegative dimensions");
|
||||
return failure();
|
||||
}
|
||||
|
||||
auto nextElements = checkedMulAtLocation(elements, static_cast<uint64_t>(dim), loc, fieldName);
|
||||
if (failed(nextElements))
|
||||
return failure();
|
||||
elements = *nextElements;
|
||||
}
|
||||
|
||||
return checkedMulAtLocation(
|
||||
elements, static_cast<uint64_t>(getElementTypeSizeInBytes(type.getElementType())), loc, 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()) {
|
||||
emitCrashMessage(fieldName, "is outside representable range");
|
||||
llvm_unreachable("PIM checked arithmetic failure");
|
||||
}
|
||||
return static_cast<int32_t>(value);
|
||||
}
|
||||
|
||||
int32_t checkedI32OrCrash(uint64_t value, llvm::StringRef fieldName) {
|
||||
if (value > static_cast<uint64_t>(std::numeric_limits<int32_t>::max())) {
|
||||
emitCrashMessage(fieldName, "is outside representable range");
|
||||
llvm_unreachable("PIM checked arithmetic failure");
|
||||
}
|
||||
return static_cast<int32_t>(value);
|
||||
}
|
||||
|
||||
uint8_t checkedU8OrCrash(uint64_t value, llvm::StringRef fieldName) {
|
||||
if (value > static_cast<uint64_t>(std::numeric_limits<uint8_t>::max())) {
|
||||
emitCrashMessage(fieldName, "is outside representable range");
|
||||
llvm_unreachable("PIM checked arithmetic failure");
|
||||
}
|
||||
return static_cast<uint8_t>(value);
|
||||
}
|
||||
|
||||
size_t checkedSizeOrCrash(int64_t value, llvm::StringRef fieldName) {
|
||||
if (value < 0) {
|
||||
emitCrashMessage(fieldName, "is outside representable range");
|
||||
llvm_unreachable("PIM checked arithmetic failure");
|
||||
}
|
||||
return static_cast<size_t>(value);
|
||||
}
|
||||
|
||||
size_t checkedAddOrCrash(size_t lhs, size_t rhs, llvm::StringRef fieldName) {
|
||||
if (rhs > std::numeric_limits<size_t>::max() - lhs) {
|
||||
emitCrashMessage(fieldName, "addition overflow");
|
||||
llvm_unreachable("PIM checked arithmetic failure");
|
||||
}
|
||||
return lhs + rhs;
|
||||
}
|
||||
|
||||
size_t checkedMulOrCrash(size_t lhs, size_t rhs, llvm::StringRef fieldName) {
|
||||
if (lhs != 0 && rhs > std::numeric_limits<size_t>::max() / lhs) {
|
||||
emitCrashMessage(fieldName, "multiplication overflow");
|
||||
llvm_unreachable("PIM checked arithmetic failure");
|
||||
}
|
||||
return lhs * rhs;
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir::pim
|
||||
@@ -0,0 +1,107 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/IR/Builders.h"
|
||||
#include "mlir/IR/BuiltinTypeInterfaces.h"
|
||||
#include "mlir/IR/Operation.h"
|
||||
#include "mlir/Support/LogicalResult.h"
|
||||
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <limits>
|
||||
#include <type_traits>
|
||||
|
||||
namespace onnx_mlir::pim {
|
||||
|
||||
mlir::InFlightDiagnostic
|
||||
emitCheckedArithmeticError(mlir::Operation* anchor, llvm::StringRef fieldName, llvm::StringRef message);
|
||||
|
||||
mlir::InFlightDiagnostic
|
||||
emitCheckedArithmeticError(mlir::Location loc, llvm::StringRef fieldName, llvm::StringRef message);
|
||||
|
||||
template <typename To, typename From>
|
||||
mlir::FailureOr<To> checkedCast(From value, mlir::Operation* anchor, llvm::StringRef fieldName) {
|
||||
static_assert(std::is_integral_v<To> && std::is_integral_v<From>, "checkedCast requires integral types");
|
||||
|
||||
using ToLimits = std::numeric_limits<To>;
|
||||
|
||||
if constexpr (std::is_signed_v<From> == std::is_signed_v<To>) {
|
||||
if (value < static_cast<From>(ToLimits::min()) || value > static_cast<From>(ToLimits::max())) {
|
||||
emitCheckedArithmeticError(anchor, fieldName, "is outside representable range");
|
||||
return mlir::failure();
|
||||
}
|
||||
}
|
||||
else if constexpr (std::is_signed_v<From>) {
|
||||
using UnsignedFrom = std::make_unsigned_t<From>;
|
||||
using UnsignedTo = std::make_unsigned_t<To>;
|
||||
if (value < 0 || static_cast<UnsignedFrom>(value) > static_cast<UnsignedTo>(ToLimits::max())) {
|
||||
emitCheckedArithmeticError(anchor, fieldName, "is outside representable range");
|
||||
return mlir::failure();
|
||||
}
|
||||
}
|
||||
else {
|
||||
using UnsignedFrom = std::make_unsigned_t<From>;
|
||||
using UnsignedTo = std::conditional_t<std::is_signed_v<To>, std::make_unsigned_t<To>, To>;
|
||||
if (static_cast<UnsignedFrom>(value) > static_cast<UnsignedTo>(ToLimits::max())) {
|
||||
emitCheckedArithmeticError(anchor, fieldName, "is outside representable range");
|
||||
return mlir::failure();
|
||||
}
|
||||
}
|
||||
|
||||
return static_cast<To>(value);
|
||||
}
|
||||
|
||||
template <typename UInt>
|
||||
mlir::FailureOr<UInt> checkedAdd(UInt lhs, UInt rhs, mlir::Operation* anchor, llvm::StringRef fieldName) {
|
||||
static_assert(std::is_integral_v<UInt> && std::is_unsigned_v<UInt>, "checkedAdd requires unsigned integral types");
|
||||
if (rhs > std::numeric_limits<UInt>::max() - lhs) {
|
||||
emitCheckedArithmeticError(anchor, fieldName, "addition overflow");
|
||||
return mlir::failure();
|
||||
}
|
||||
return lhs + rhs;
|
||||
}
|
||||
|
||||
template <typename UInt>
|
||||
mlir::FailureOr<UInt> checkedMul(UInt lhs, UInt rhs, mlir::Operation* anchor, llvm::StringRef fieldName) {
|
||||
static_assert(std::is_integral_v<UInt> && std::is_unsigned_v<UInt>, "checkedMul requires unsigned integral types");
|
||||
if (lhs != 0 && rhs > std::numeric_limits<UInt>::max() / lhs) {
|
||||
emitCheckedArithmeticError(anchor, fieldName, "multiplication overflow");
|
||||
return mlir::failure();
|
||||
}
|
||||
return lhs * rhs;
|
||||
}
|
||||
|
||||
mlir::FailureOr<int32_t> checkedI32(int64_t value, mlir::Operation* anchor, llvm::StringRef fieldName);
|
||||
mlir::FailureOr<int32_t> checkedI32(uint64_t value, mlir::Operation* anchor, llvm::StringRef fieldName);
|
||||
|
||||
mlir::FailureOr<uint8_t> checkedU8(uint64_t value, mlir::Operation* anchor, llvm::StringRef fieldName);
|
||||
|
||||
mlir::FailureOr<size_t> checkedSize(int64_t value, mlir::Operation* anchor, llvm::StringRef fieldName);
|
||||
|
||||
mlir::FailureOr<mlir::IntegerAttr>
|
||||
getCheckedI32Attr(mlir::Builder& builder, mlir::Operation* anchor, int64_t value, llvm::StringRef fieldName);
|
||||
|
||||
mlir::FailureOr<mlir::IntegerAttr>
|
||||
getCheckedI32Attr(mlir::Builder& builder, mlir::Operation* anchor, uint64_t value, llvm::StringRef fieldName);
|
||||
|
||||
mlir::FailureOr<mlir::IntegerAttr>
|
||||
getCheckedI32Attr(mlir::Builder& builder, mlir::Location loc, int64_t value, llvm::StringRef fieldName);
|
||||
|
||||
mlir::FailureOr<mlir::IntegerAttr>
|
||||
getCheckedI32Attr(mlir::Builder& builder, mlir::Location loc, uint64_t value, llvm::StringRef fieldName);
|
||||
|
||||
mlir::FailureOr<uint64_t>
|
||||
getCheckedShapedTypeSizeInBytes(mlir::ShapedType type, mlir::Operation* anchor, llvm::StringRef fieldName);
|
||||
|
||||
mlir::FailureOr<uint64_t>
|
||||
getCheckedShapedTypeSizeInBytes(mlir::ShapedType type, mlir::Location loc, llvm::StringRef fieldName);
|
||||
|
||||
int32_t checkedI32OrCrash(int64_t value, llvm::StringRef fieldName);
|
||||
int32_t checkedI32OrCrash(uint64_t value, llvm::StringRef fieldName);
|
||||
uint8_t checkedU8OrCrash(uint64_t value, llvm::StringRef fieldName);
|
||||
size_t checkedSizeOrCrash(int64_t value, llvm::StringRef fieldName);
|
||||
size_t checkedAddOrCrash(size_t lhs, size_t rhs, llvm::StringRef fieldName);
|
||||
size_t checkedMulOrCrash(size_t lhs, size_t rhs, llvm::StringRef fieldName);
|
||||
|
||||
} // namespace onnx_mlir::pim
|
||||
@@ -7,18 +7,26 @@
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
void dumpModule(mlir::ModuleOp moduleOp, const std::string& name) {
|
||||
std::fstream openDialectDumpFileWithExtension(const std::string& name, llvm::StringRef destination, llvm::StringRef extension) {
|
||||
std::string outputDir = getOutputDir();
|
||||
if (outputDir.empty())
|
||||
return {};
|
||||
|
||||
std::string dialectsDir = (outputDir + destination).str();
|
||||
createDirectory(dialectsDir);
|
||||
return std::fstream(dialectsDir + "/" + name + "." + extension.str(), std::ios::out);
|
||||
}
|
||||
|
||||
void dumpModule(mlir::ModuleOp moduleOp, const std::string& name, bool assumeVerified) {
|
||||
std::fstream file = openDialectDumpFileWithExtension(name, "/dialects", "mlir");
|
||||
if (!file.is_open())
|
||||
return;
|
||||
|
||||
std::string dialectsDir = outputDir + "/dialects";
|
||||
createDirectory(dialectsDir);
|
||||
|
||||
std::fstream file(dialectsDir + "/" + name + ".mlir", std::ios::out);
|
||||
llvm::raw_os_ostream os(file);
|
||||
mlir::OpPrintingFlags flags;
|
||||
flags.elideLargeElementsAttrs();
|
||||
flags.elideLargeElementsAttrs().enableDebugInfo(true, false);
|
||||
if (assumeVerified)
|
||||
flags.assumeVerified();
|
||||
moduleOp.print(os, flags);
|
||||
os.flush();
|
||||
file.close();
|
||||
|
||||
@@ -1,13 +1,18 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/IR/BuiltinOps.h"
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
|
||||
#include <fstream>
|
||||
#include <string>
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
/// Emits a MLIR snapshot under the current compiler output
|
||||
/// directory for pass-level debugging.
|
||||
void dumpModule(mlir::ModuleOp moduleOp, const std::string& name);
|
||||
void dumpModule(mlir::ModuleOp moduleOp, const std::string& name, bool assumeVerified = false);
|
||||
|
||||
/// Opens a file under the same dialect dump directory used by dumpModule.
|
||||
std::fstream openDialectDumpFileWithExtension(const std::string& name,llvm::StringRef destination = "/dialects", llvm::StringRef extension = "mlir");
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -7,10 +7,36 @@
|
||||
#include "llvm/ADT/ArrayRef.h"
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
|
||||
#include <cstdint>
|
||||
#include <system_error>
|
||||
|
||||
namespace onnx_mlir::pim {
|
||||
|
||||
struct CappedDiagnosticReporter {
|
||||
explicit CappedDiagnosticReporter(int64_t maxReportedFailures = 8)
|
||||
: maxReportedFailures(maxReportedFailures) {}
|
||||
|
||||
template <typename EmitFn>
|
||||
void report(mlir::Operation* op, EmitFn&& emit) {
|
||||
numFailures++;
|
||||
if (numFailures <= maxReportedFailures)
|
||||
emit(op);
|
||||
}
|
||||
|
||||
void emitSuppressedSummary(mlir::Operation* op, llvm::StringRef failureDescription) const {
|
||||
if (numFailures > maxReportedFailures)
|
||||
op->emitError() << "suppressed " << (numFailures - maxReportedFailures) << " additional " << failureDescription;
|
||||
}
|
||||
|
||||
void noteFailures(int64_t count) { numFailures += count; }
|
||||
|
||||
bool hasFailure() const { return numFailures != 0; }
|
||||
|
||||
private:
|
||||
int64_t maxReportedFailures;
|
||||
int64_t numFailures = 0;
|
||||
};
|
||||
|
||||
/// Emits a consistent diagnostic for target paths that require static shapes.
|
||||
mlir::InFlightDiagnostic emitUnsupportedStaticShapeDiagnostic(mlir::Operation* op, llvm::StringRef valueDescription);
|
||||
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
#include "llvm/Support/Format.h"
|
||||
|
||||
#include "src/Accelerators/PIM/Common/Support/FileSystemUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/Support/ReportUtils.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
std::fstream openReportFileWithExtension(const std::string& name, llvm::StringRef extension) {
|
||||
std::string outputDir = getOutputDir();
|
||||
if (outputDir.empty())
|
||||
return {};
|
||||
|
||||
std::string reportsDir = outputDir + "/reports";
|
||||
createDirectory(reportsDir);
|
||||
return std::fstream(reportsDir + "/" + name + "." + extension.str(), std::ios::out);
|
||||
}
|
||||
|
||||
std::fstream openReportFile(const std::string& name) { return openReportFileWithExtension(name, "txt"); }
|
||||
|
||||
std::fstream openAppendedReportFileWithExtension(const std::string& name, llvm::StringRef extension) {
|
||||
std::string outputDir = getOutputDir();
|
||||
if (outputDir.empty())
|
||||
return {};
|
||||
|
||||
std::string reportsDir = outputDir + "/reports";
|
||||
createDirectory(reportsDir);
|
||||
return std::fstream(reportsDir + "/" + name + "." + extension.str(), std::ios::out | std::ios::app);
|
||||
}
|
||||
|
||||
std::fstream openAppendedReportFile(const std::string& name) {
|
||||
return openAppendedReportFileWithExtension(name, "txt");
|
||||
}
|
||||
|
||||
std::string formatReportMemory(uint64_t bytes) {
|
||||
const char* units[] = {"B", "KB", "MB", "GB", "TB", "PB", "EB"};
|
||||
int i = 0;
|
||||
double size = static_cast<double>(bytes);
|
||||
while (size >= 1024 && i < 6) {
|
||||
size /= 1024;
|
||||
i++;
|
||||
}
|
||||
|
||||
std::string out;
|
||||
llvm::raw_string_ostream rss(out);
|
||||
rss << llvm::format("%.2f ", size) << units[i];
|
||||
return rss.str();
|
||||
}
|
||||
|
||||
void printReportFlatFields(llvm::raw_ostream& os, llvm::ArrayRef<ReportField> fields) {
|
||||
for (const ReportField& field : fields)
|
||||
os << "\t" << field.label << ": " << field.value << "\n";
|
||||
}
|
||||
|
||||
void printReportFieldBlock(llvm::raw_ostream& os, llvm::StringRef title, llvm::ArrayRef<ReportField> fields) {
|
||||
os << "\t" << title << ":\n";
|
||||
for (const ReportField& field : fields)
|
||||
os << "\t " << field.label << ": " << field.value << "\n";
|
||||
}
|
||||
|
||||
void printReportTotalsBlock(llvm::raw_ostream& os, llvm::ArrayRef<ReportField> fields) {
|
||||
os << "Totals:\n";
|
||||
for (const ReportField& field : fields)
|
||||
os << "\t" << field.label << ": " << field.value << "\n";
|
||||
}
|
||||
|
||||
void printReportPerCoreAndTotalFields(llvm::raw_ostream& os,
|
||||
llvm::ArrayRef<ReportField> perCoreFields,
|
||||
llvm::ArrayRef<ReportField> totalFields) {
|
||||
printReportFieldBlock(os, "Per core", perCoreFields);
|
||||
printReportFieldBlock(os, "Total", totalFields);
|
||||
}
|
||||
|
||||
void printReportEntrySeparator(llvm::raw_ostream& os, bool hasNextEntry) {
|
||||
if (hasNextEntry)
|
||||
os << "\n";
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,50 @@
|
||||
#pragma once
|
||||
|
||||
#include "llvm/ADT/ArrayRef.h"
|
||||
#include "llvm/ADT/STLExtras.h"
|
||||
#include "llvm/Support/raw_ostream.h"
|
||||
|
||||
#include <fstream>
|
||||
#include <limits>
|
||||
#include <string>
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
std::fstream openReportFile(const std::string& name);
|
||||
std::fstream openReportFileWithExtension(const std::string& name, llvm::StringRef extension);
|
||||
std::fstream openAppendedReportFile(const std::string& name);
|
||||
std::fstream openAppendedReportFileWithExtension(const std::string& name, llvm::StringRef extension);
|
||||
std::string formatReportMemory(uint64_t bytes);
|
||||
|
||||
struct ReportField {
|
||||
std::string label;
|
||||
std::string value;
|
||||
};
|
||||
|
||||
void printReportFlatFields(llvm::raw_ostream& os, llvm::ArrayRef<ReportField> fields);
|
||||
void printReportFieldBlock(llvm::raw_ostream& os, llvm::StringRef title, llvm::ArrayRef<ReportField> fields);
|
||||
void printReportTotalsBlock(llvm::raw_ostream& os, llvm::ArrayRef<ReportField> fields);
|
||||
void printReportPerCoreAndTotalFields(llvm::raw_ostream& os,
|
||||
llvm::ArrayRef<ReportField> perCoreFields,
|
||||
llvm::ArrayRef<ReportField> totalFields);
|
||||
void printReportEntrySeparator(llvm::raw_ostream& os, bool hasNextEntry);
|
||||
|
||||
template <typename EntryTy>
|
||||
int32_t getFirstReportCoreId(const EntryTy& entry) {
|
||||
if (entry.coreIds.empty())
|
||||
return std::numeric_limits<int32_t>::max();
|
||||
return entry.coreIds.front();
|
||||
}
|
||||
|
||||
template <typename EntryRange>
|
||||
void sortReportEntriesByFirstCore(EntryRange& entries) {
|
||||
llvm::stable_sort(entries, [](const auto& lhs, const auto& rhs) {
|
||||
int32_t lhsFirstCore = getFirstReportCoreId(lhs);
|
||||
int32_t rhsFirstCore = getFirstReportCoreId(rhs);
|
||||
if (lhsFirstCore != rhsFirstCore)
|
||||
return lhsFirstCore < rhsFirstCore;
|
||||
return lhs.id < rhs.id;
|
||||
});
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -15,7 +15,10 @@ add_pim_library(OMPimCompilerOptions
|
||||
|
||||
add_pim_library(OMPimCompilerUtils
|
||||
PimCompilerUtils.cpp
|
||||
PimArtifactWriter.cpp
|
||||
PimCodeGen.cpp
|
||||
PimCoreProgram.cpp
|
||||
PimWeightEmitter.cpp
|
||||
|
||||
EXCLUDE_FROM_OM_LIBS
|
||||
|
||||
@@ -23,9 +26,14 @@ add_pim_library(OMPimCompilerUtils
|
||||
${PIM_COMPILER_INCLUDE_DIRS}
|
||||
|
||||
LINK_LIBS PUBLIC
|
||||
MLIRAffineToStandard
|
||||
OMPimCompilerOptions
|
||||
OMPimCommon
|
||||
OMPimBufferization
|
||||
OMPimHostConstantFolding
|
||||
OMPimInstructionSelection
|
||||
OMPimLocalMemoryPlanning
|
||||
OMPimVerification
|
||||
OMPimPasses
|
||||
OMONNXToSpatial
|
||||
OMSpatialToPim
|
||||
|
||||
@@ -0,0 +1,165 @@
|
||||
#include "mlir/Dialect/MemRef/IR/MemRef.h"
|
||||
|
||||
#include "llvm/ADT/SmallPtrSet.h"
|
||||
#include "llvm/Support/FileSystem.h"
|
||||
#include "llvm/Support/raw_ostream.h"
|
||||
|
||||
#include <algorithm>
|
||||
#include <cassert>
|
||||
#include <cstring>
|
||||
#include <limits>
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/WeightUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimArtifactWriter.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimBinaryFormat.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimCodeGen.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
|
||||
|
||||
using namespace llvm;
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
PimInstructionWriter::PimInstructionWriter(raw_pwrite_stream& stream)
|
||||
: stream(stream), buffer(kBufferSize) {
|
||||
pim_binary::writeHeader(stream);
|
||||
}
|
||||
|
||||
PimInstructionWriter::PimInstructionWriter(raw_fd_ostream& stream)
|
||||
: PimInstructionWriter(static_cast<raw_pwrite_stream&>(stream)) {
|
||||
fileStream = &stream;
|
||||
}
|
||||
|
||||
bool PimInstructionWriter::consumeStreamError() {
|
||||
if (!fileStream || !fileStream->has_error())
|
||||
return false;
|
||||
fileStream->clear_error();
|
||||
hasFailure = true;
|
||||
return true;
|
||||
}
|
||||
|
||||
LogicalResult PimInstructionWriter::flushBuffer() {
|
||||
if (hasFailure)
|
||||
return failure();
|
||||
if (bufferedBytes != 0) {
|
||||
stream.write(buffer.data(), bufferedBytes);
|
||||
bufferedBytes = 0;
|
||||
}
|
||||
return consumeStreamError() ? failure() : success();
|
||||
}
|
||||
|
||||
LogicalResult PimInstructionWriter::append(const pim_binary::InstructionRecord& record) {
|
||||
if (hasFailure || finalized || count == std::numeric_limits<uint32_t>::max()) {
|
||||
hasFailure = true;
|
||||
return failure();
|
||||
}
|
||||
|
||||
if (bufferedBytes + pim_binary::kRecordSize > buffer.size() && failed(flushBuffer()))
|
||||
return failure();
|
||||
|
||||
pim_binary::EncodedInstruction encoded = pim_binary::encodeInstructionRecord(record);
|
||||
std::memcpy(buffer.data() + bufferedBytes, encoded.data(), encoded.size());
|
||||
bufferedBytes += encoded.size();
|
||||
++count;
|
||||
return success();
|
||||
}
|
||||
|
||||
LogicalResult PimInstructionWriter::finalize() {
|
||||
if (finalized)
|
||||
return failure();
|
||||
finalized = true;
|
||||
if (hasFailure || failed(flushBuffer()))
|
||||
return failure();
|
||||
|
||||
stream.flush();
|
||||
if (consumeStreamError())
|
||||
return failure();
|
||||
pim_binary::patchInstructionCount(stream, count);
|
||||
stream.flush();
|
||||
return consumeStreamError() ? failure() : success();
|
||||
}
|
||||
|
||||
OnnxMlirCompilerErrorCodes
|
||||
writeMemoryBinary(ModuleOp moduleOp, func::FuncOp funcOp, PimAcceleratorMemory& memory, StringRef outputDirPath) {
|
||||
auto memoryFilePath = (outputDirPath + "/memory.bin").str();
|
||||
std::error_code errorCode;
|
||||
raw_fd_ostream memoryFileStream(memoryFilePath, errorCode, sys::fs::OF_None);
|
||||
if (errorCode) {
|
||||
errs() << "Error while opening memory file " << memoryFilePath << ": " << errorCode.message() << '\n';
|
||||
return InvalidOutputFileAccess;
|
||||
}
|
||||
|
||||
std::vector<char> memoryBuffer(memory.hostMem.getFirstAvailableAddress(), 0);
|
||||
|
||||
SmallPtrSet<Operation*, 16> writtenGlobals;
|
||||
funcOp.walk([&](memref::GetGlobalOp getGlobalOp) {
|
||||
if (hasWeightAlways(getGlobalOp))
|
||||
return;
|
||||
auto globalOp = lookupGlobalForGetGlobal(moduleOp, getGlobalOp);
|
||||
if (!globalOp)
|
||||
return;
|
||||
if (!writtenGlobals.insert(globalOp.getOperation()).second)
|
||||
return;
|
||||
auto initialValue = globalOp.getInitialValue();
|
||||
if (!initialValue)
|
||||
return;
|
||||
auto denseAttr = dyn_cast<DenseElementsAttr>(*initialValue);
|
||||
if (!denseAttr)
|
||||
return;
|
||||
|
||||
MemEntry memEntry = memory.hostMem.getMemEntry({getGlobalOp.getResult(), std::nullopt});
|
||||
ArrayRef<char> rawData = denseAttr.getRawData();
|
||||
char* dst = memoryBuffer.data() + memEntry.address;
|
||||
|
||||
if (denseAttr.isSplat()) {
|
||||
size_t elementSize = rawData.size();
|
||||
assert(elementSize * getGlobalOp.getType().getNumElements() == memEntry.size && "Data size mismatch");
|
||||
for (size_t offset = 0; offset < memEntry.size; offset += elementSize)
|
||||
std::memcpy(dst + offset, rawData.data(), std::min(elementSize, memEntry.size - offset));
|
||||
}
|
||||
else {
|
||||
assert(rawData.size() == memEntry.size && "Data size mismatch");
|
||||
std::memcpy(dst, rawData.data(), rawData.size());
|
||||
}
|
||||
});
|
||||
|
||||
memoryFileStream.write(memoryBuffer.data(), memoryBuffer.size());
|
||||
memoryFileStream.close();
|
||||
return CompilerSuccess;
|
||||
}
|
||||
|
||||
OnnxMlirCompilerErrorCodes writeConfigJson(func::FuncOp funcOp,
|
||||
PimAcceleratorMemory& memory,
|
||||
size_t maxCoreId,
|
||||
json::Object xbarsPerArrayGroup,
|
||||
StringRef outputDirPath) {
|
||||
json::Object configJson;
|
||||
|
||||
configJson["core_cnt"] = maxCoreId + 1;
|
||||
configJson["xbar_size"] = {crossbarSize.getValue(), crossbarSize.getValue()};
|
||||
configJson["array_group_map"] = std::move(xbarsPerArrayGroup);
|
||||
|
||||
json::Array inputsAddresses;
|
||||
for (BlockArgument input : funcOp.getArguments())
|
||||
inputsAddresses.push_back(memory.getValueAddress(input));
|
||||
configJson["inputs_addresses"] = std::move(inputsAddresses);
|
||||
|
||||
json::Array outputsAddresses;
|
||||
for (func::ReturnOp returnOp : funcOp.getOps<func::ReturnOp>())
|
||||
for (mlir::Value output : returnOp.getOperands())
|
||||
outputsAddresses.push_back(memory.getValueAddress(output));
|
||||
configJson["outputs_addresses"] = std::move(outputsAddresses);
|
||||
|
||||
auto configPath = (outputDirPath + "/config.json").str();
|
||||
std::error_code errorCode;
|
||||
raw_fd_ostream jsonOS(configPath, errorCode);
|
||||
if (errorCode) {
|
||||
errs() << "Error while opening config file: " << errorCode.message() << '\n';
|
||||
return InvalidOutputFileAccess;
|
||||
}
|
||||
jsonOS << json::Value(std::move(configJson)) << '\n';
|
||||
jsonOS.close();
|
||||
return CompilerSuccess;
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,54 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/Dialect/Func/IR/FuncOps.h"
|
||||
#include "mlir/IR/BuiltinOps.h"
|
||||
#include "mlir/Support/LogicalResult.h"
|
||||
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
#include "llvm/Support/JSON.h"
|
||||
#include "llvm/Support/raw_ostream.h"
|
||||
|
||||
#include <cstddef>
|
||||
#include <cstdint>
|
||||
#include <vector>
|
||||
|
||||
#include "onnx-mlir/Compiler/OMCompilerTypes.h"
|
||||
#include "src/Accelerators/PIM/Compiler/PimBinaryFormat.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
class PimAcceleratorMemory;
|
||||
|
||||
class PimInstructionWriter {
|
||||
public:
|
||||
explicit PimInstructionWriter(llvm::raw_pwrite_stream& stream);
|
||||
explicit PimInstructionWriter(llvm::raw_fd_ostream& stream);
|
||||
|
||||
mlir::LogicalResult append(const pim_binary::InstructionRecord& record);
|
||||
mlir::LogicalResult finalize();
|
||||
|
||||
private:
|
||||
static constexpr size_t kBufferSize = 256 * 1024;
|
||||
mlir::LogicalResult flushBuffer();
|
||||
bool consumeStreamError();
|
||||
|
||||
llvm::raw_pwrite_stream& stream;
|
||||
llvm::raw_fd_ostream* fileStream = nullptr;
|
||||
std::vector<char> buffer;
|
||||
size_t bufferedBytes = 0;
|
||||
uint32_t count = 0;
|
||||
bool finalized = false;
|
||||
bool hasFailure = false;
|
||||
};
|
||||
|
||||
OnnxMlirCompilerErrorCodes writeMemoryBinary(mlir::ModuleOp moduleOp,
|
||||
mlir::func::FuncOp funcOp,
|
||||
PimAcceleratorMemory& memory,
|
||||
llvm::StringRef outputDirPath);
|
||||
OnnxMlirCompilerErrorCodes writeConfigJson(mlir::func::FuncOp funcOp,
|
||||
PimAcceleratorMemory& memory,
|
||||
size_t maxCoreId,
|
||||
llvm::json::Object xbarsPerArrayGroup,
|
||||
llvm::StringRef outputDirPath);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,238 @@
|
||||
#pragma once
|
||||
|
||||
#include "llvm/ADT/STLExtras.h"
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
#include "llvm/Support/Endian.h"
|
||||
#include "llvm/Support/JSON.h"
|
||||
#include "llvm/Support/raw_ostream.h"
|
||||
|
||||
#include <array>
|
||||
|
||||
#include "src/Accelerators/PIM/Common/Support/CheckedArithmetic.hpp"
|
||||
|
||||
namespace onnx_mlir::pim_binary {
|
||||
|
||||
inline constexpr char kMagic[4] = {'P', 'I', 'M', 'B'};
|
||||
inline constexpr uint32_t kVersion = 1;
|
||||
inline constexpr uint64_t kCountOffset = 8;
|
||||
inline constexpr size_t kHeaderSize = 12;
|
||||
inline constexpr size_t kRecordSize = 20;
|
||||
|
||||
enum class Opcode : uint32_t {
|
||||
nop = 0,
|
||||
sldi = 1,
|
||||
sld = 2,
|
||||
sadd = 3,
|
||||
ssub = 4,
|
||||
smul = 5,
|
||||
saddi = 6,
|
||||
smuli = 7,
|
||||
setbw = 8,
|
||||
mvmul = 9,
|
||||
vvadd = 10,
|
||||
vvsub = 11,
|
||||
vvmul = 12,
|
||||
vvdmul = 13,
|
||||
vvmax = 14,
|
||||
vvsll = 15,
|
||||
vvsra = 16,
|
||||
vavg = 17,
|
||||
vrelu = 18,
|
||||
vtanh = 19,
|
||||
vsigm = 20,
|
||||
vsoftmax = 21,
|
||||
vmv = 22,
|
||||
vrsu = 23,
|
||||
vrsl = 24,
|
||||
ld = 25,
|
||||
st = 26,
|
||||
lldi = 27,
|
||||
lmv = 28,
|
||||
send = 29,
|
||||
recv = 30,
|
||||
wait = 31,
|
||||
sync = 32,
|
||||
};
|
||||
|
||||
inline constexpr size_t kOpcodeCount = static_cast<size_t>(Opcode::sync) + 1;
|
||||
inline constexpr std::array<llvm::StringLiteral, kOpcodeCount> kOpcodeNames = {
|
||||
"nop", "sldi", "sld", "sadd", "ssub", "smul", "saddi", "smuli", "setbw", "mvmul", "vvadd",
|
||||
"vvsub", "vvmul", "vvdmul", "vvmax", "vvsll", "vvsra", "vavg", "vrelu", "vtanh", "vsigm", "vsoftmax",
|
||||
"vmv", "vrsu", "vrsl", "ld", "st", "lldi", "lmv", "send", "recv", "wait", "sync",
|
||||
};
|
||||
static_assert(kOpcodeNames.size() == kOpcodeCount);
|
||||
|
||||
struct InstructionRecord {
|
||||
Opcode opcode = Opcode::nop;
|
||||
uint8_t rd = 0;
|
||||
uint8_t r1 = 0;
|
||||
int32_t r2OrImm = 0;
|
||||
int32_t generic1 = 0;
|
||||
int32_t generic2 = 0;
|
||||
int32_t generic3 = 0;
|
||||
uint8_t flags = 0;
|
||||
};
|
||||
|
||||
using EncodedInstruction = std::array<char, kRecordSize>;
|
||||
static_assert(kRecordSize == 20);
|
||||
static_assert(sizeof(EncodedInstruction) == kRecordSize);
|
||||
|
||||
inline EncodedInstruction encodeInstructionRecord(const InstructionRecord& record) {
|
||||
EncodedInstruction encoded = {};
|
||||
encoded[0] = static_cast<char>(static_cast<uint8_t>(record.opcode));
|
||||
encoded[1] = static_cast<char>(record.rd);
|
||||
encoded[2] = static_cast<char>(record.r1);
|
||||
encoded[3] = static_cast<char>(record.flags);
|
||||
llvm::support::endian::write32le(encoded.data() + 4, static_cast<uint32_t>(record.r2OrImm));
|
||||
llvm::support::endian::write32le(encoded.data() + 8, static_cast<uint32_t>(record.generic1));
|
||||
llvm::support::endian::write32le(encoded.data() + 12, static_cast<uint32_t>(record.generic2));
|
||||
llvm::support::endian::write32le(encoded.data() + 16, static_cast<uint32_t>(record.generic3));
|
||||
return encoded;
|
||||
}
|
||||
|
||||
inline void writeUint32LE(llvm::raw_ostream& os, uint32_t value) {
|
||||
std::array<char, sizeof(uint32_t)> bytes;
|
||||
llvm::support::endian::write32le(bytes.data(), value);
|
||||
os.write(bytes.data(), bytes.size());
|
||||
}
|
||||
|
||||
inline void writeHeader(llvm::raw_ostream& os) {
|
||||
os.write(kMagic, sizeof(kMagic));
|
||||
writeUint32LE(os, kVersion);
|
||||
writeUint32LE(os, 0);
|
||||
}
|
||||
|
||||
inline void patchInstructionCount(llvm::raw_pwrite_stream& os, uint32_t instructionCount) {
|
||||
std::array<char, sizeof(uint32_t)> bytes;
|
||||
llvm::support::endian::write32le(bytes.data(), instructionCount);
|
||||
os.pwrite(bytes.data(), bytes.size(), kCountOffset);
|
||||
}
|
||||
|
||||
inline int32_t toI32(int64_t value) { return onnx_mlir::pim::checkedI32OrCrash(value, "binary field"); }
|
||||
|
||||
inline uint8_t toU8(int64_t value) {
|
||||
return onnx_mlir::pim::checkedU8OrCrash(static_cast<uint64_t>(value), "binary field");
|
||||
}
|
||||
|
||||
inline int32_t getOptionalInt(const llvm::json::Object& object, llvm::StringRef key, int32_t defaultValue = 0) {
|
||||
if (std::optional<int64_t> value = object.getInteger(key))
|
||||
return toI32(*value);
|
||||
return defaultValue;
|
||||
}
|
||||
|
||||
struct InstructionJsonFormat {
|
||||
bool rd;
|
||||
bool r1;
|
||||
bool offset;
|
||||
llvm::StringLiteral r2;
|
||||
llvm::StringLiteral generic1;
|
||||
llvm::StringLiteral generic2;
|
||||
llvm::StringLiteral generic3;
|
||||
};
|
||||
|
||||
inline constexpr std::array<InstructionJsonFormat, kOpcodeCount> kInstructionJsonFormats = {{
|
||||
{false, false, false, "", "", "", "" }, // nop
|
||||
{true, false, false, "imm", "", "", "" }, // sldi
|
||||
{true, true, true, "", "", "", "" }, // sld
|
||||
{true, true, false, "rs2", "", "", "" }, // sadd
|
||||
{true, true, false, "rs2", "", "", "" }, // ssub
|
||||
{true, true, false, "rs2", "", "", "" }, // smul
|
||||
{true, true, false, "imm", "", "", "" }, // saddi
|
||||
{true, true, false, "imm", "", "", "" }, // smuli
|
||||
{false, false, false, "", "ibiw", "obiw", "" }, // setbw
|
||||
{true, true, false, "mbiw", "relu", "group", "" }, // mvmul
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvadd
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvsub
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvmul
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvdmul
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvmax
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvsll
|
||||
{true, true, true, "rs2", "", "", "len" }, // vvsra
|
||||
{true, true, true, "rs2", "", "", "len" }, // vavg
|
||||
{true, true, true, "", "", "", "len" }, // vrelu
|
||||
{true, true, true, "", "", "", "len" }, // vtanh
|
||||
{true, true, true, "", "", "", "len" }, // vsigm
|
||||
{true, true, true, "", "", "", "len" }, // vsoftmax
|
||||
{true, true, true, "rs2", "", "", "len" }, // vmv
|
||||
{true, true, true, "rs2", "", "", "len" }, // vrsu
|
||||
{true, true, true, "rs2", "", "", "len" }, // vrsl
|
||||
{true, true, true, "", "", "", "size"}, // ld
|
||||
{true, true, true, "", "", "", "size"}, // st
|
||||
{true, false, true, "imm", "", "", "len" }, // lldi
|
||||
{true, true, true, "", "", "", "len" }, // lmv
|
||||
{true, false, true, "core", "", "", "size"}, // send
|
||||
{true, false, true, "core", "", "", "size"}, // recv
|
||||
{false, false, false, "", "event_register", "wait_value", ""}, // wait
|
||||
{false, false, false, "core", "event_register", "", ""}, // sync
|
||||
}};
|
||||
static_assert(kInstructionJsonFormats.size() == kOpcodeCount);
|
||||
|
||||
inline Opcode opcodeFromString(llvm::StringRef opName) {
|
||||
for (auto [index, name] : llvm::enumerate(kOpcodeNames))
|
||||
if (opName == name)
|
||||
return static_cast<Opcode>(index);
|
||||
llvm_unreachable("Unsupported PIM binary opcode");
|
||||
}
|
||||
|
||||
inline llvm::StringRef opcodeToString(Opcode opcode) {
|
||||
size_t index = static_cast<size_t>(opcode);
|
||||
assert(index < kOpcodeNames.size() && "Unsupported PIM binary opcode");
|
||||
return kOpcodeNames[index];
|
||||
}
|
||||
|
||||
inline InstructionRecord makeInstructionRecord(const llvm::json::Object& instruction) {
|
||||
InstructionRecord record;
|
||||
std::optional<llvm::StringRef> opName = instruction.getString("op");
|
||||
assert(opName && "Missing op field in PIM instruction");
|
||||
record.opcode = opcodeFromString(*opName);
|
||||
const auto& format = kInstructionJsonFormats[static_cast<size_t>(record.opcode)];
|
||||
if (format.rd)
|
||||
record.rd = toU8(getOptionalInt(instruction, "rd"));
|
||||
if (format.r1)
|
||||
record.r1 = toU8(getOptionalInt(instruction, "rs1"));
|
||||
if (!format.r2.empty())
|
||||
record.r2OrImm = getOptionalInt(instruction, format.r2);
|
||||
if (!format.generic1.empty())
|
||||
record.generic1 = getOptionalInt(instruction, format.generic1);
|
||||
if (!format.generic2.empty())
|
||||
record.generic2 = getOptionalInt(instruction, format.generic2);
|
||||
if (format.offset) {
|
||||
if (auto* offsetValue = instruction.getObject("offset")) {
|
||||
record.generic1 = getOptionalInt(*offsetValue, "offset_select");
|
||||
record.generic2 = getOptionalInt(*offsetValue, "offset_value");
|
||||
}
|
||||
}
|
||||
if (!format.generic3.empty())
|
||||
record.generic3 = getOptionalInt(instruction, format.generic3);
|
||||
return record;
|
||||
}
|
||||
|
||||
inline llvm::json::Object makeInstructionJson(const InstructionRecord& record) {
|
||||
llvm::json::Object instruction;
|
||||
instruction["op"] = opcodeToString(record.opcode).str();
|
||||
|
||||
auto addOffset = [&](int32_t offsetSelect, int32_t offsetValue) {
|
||||
llvm::json::Object offset;
|
||||
offset["offset_select"] = offsetSelect;
|
||||
offset["offset_value"] = offsetValue;
|
||||
instruction["offset"] = std::move(offset);
|
||||
};
|
||||
const auto& format = kInstructionJsonFormats[static_cast<size_t>(record.opcode)];
|
||||
if (format.rd)
|
||||
instruction["rd"] = static_cast<int64_t>(record.rd);
|
||||
if (format.r1)
|
||||
instruction["rs1"] = static_cast<int64_t>(record.r1);
|
||||
if (!format.r2.empty())
|
||||
instruction[format.r2] = record.r2OrImm;
|
||||
if (!format.generic1.empty())
|
||||
instruction[format.generic1] = record.generic1;
|
||||
if (!format.generic2.empty())
|
||||
instruction[format.generic2] = record.generic2;
|
||||
if (format.offset)
|
||||
addOffset(record.generic1, record.generic2);
|
||||
if (!format.generic3.empty())
|
||||
instruction[format.generic3] = record.generic3;
|
||||
return instruction;
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir::pim_binary
|
||||
+1177
-958
File diff suppressed because it is too large
Load Diff
+162
-54
@@ -1,140 +1,248 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/IR/Operation.h"
|
||||
|
||||
#include "llvm-project/clang/include/clang/Basic/LLVM.h"
|
||||
#include "llvm/ADT/DenseMap.h"
|
||||
#include "llvm/ADT/Hashing.h"
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
#include "llvm/Support/JSON.h"
|
||||
#include "llvm/Support/raw_os_ostream.h"
|
||||
|
||||
#include <array>
|
||||
#include <fstream>
|
||||
#include <limits>
|
||||
#include <optional>
|
||||
#include <string>
|
||||
|
||||
#include "onnx-mlir/Compiler/OMCompilerTypes.h"
|
||||
#include "src/Accelerators/PIM/Common/IR/AddressAnalysis.hpp"
|
||||
#include "src/Accelerators/PIM/Common/PimCommon.hpp"
|
||||
#include "src/Accelerators/PIM/Common/Support/ReportUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimBinaryFormat.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
struct CompiledCoreProgram;
|
||||
class PimInstructionWriter;
|
||||
|
||||
struct MemEntry {
|
||||
size_t address;
|
||||
size_t size;
|
||||
};
|
||||
|
||||
class PimMemory {
|
||||
llvm::SmallVector<std::pair<MemEntry, mlir::Value>, 32> memEntries;
|
||||
llvm::SmallDenseMap<mlir::Value, MemEntry, 32>& globalMemEntriesMap;
|
||||
struct MemoryValueKey {
|
||||
mlir::Value value;
|
||||
std::optional<unsigned> lane;
|
||||
|
||||
bool operator==(const MemoryValueKey& other) const { return value == other.value && lane == other.lane; }
|
||||
};
|
||||
|
||||
struct CompiledLocalMemoryEntry {
|
||||
mlir::Value value;
|
||||
MemEntry memory;
|
||||
};
|
||||
|
||||
struct CompiledCoreMemoryPlan {
|
||||
llvm::SmallVector<CompiledLocalMemoryEntry, 32> entries;
|
||||
uint64_t logicalAllocationCount = 0;
|
||||
uint64_t logicalBytes = 0;
|
||||
size_t arenaSize = 0;
|
||||
};
|
||||
|
||||
struct MemoryReportRow {
|
||||
uint64_t hostObjectCount = 0;
|
||||
uint64_t hostBytes = 0;
|
||||
uint64_t logicalLocalAllocationCount = 0;
|
||||
uint64_t logicalLocalBytes = 0;
|
||||
uint64_t physicalLocalBytes = 0;
|
||||
|
||||
bool operator==(const MemoryReportRow& other) const {
|
||||
return hostObjectCount == other.hostObjectCount && hostBytes == other.hostBytes
|
||||
&& logicalLocalAllocationCount == other.logicalLocalAllocationCount
|
||||
&& logicalLocalBytes == other.logicalLocalBytes && physicalLocalBytes == other.physicalLocalBytes;
|
||||
}
|
||||
};
|
||||
|
||||
enum class MemoryReportKind {
|
||||
None,
|
||||
Alloca,
|
||||
Global,
|
||||
Input
|
||||
};
|
||||
|
||||
struct PendingMemEntry {
|
||||
MemEntry memEntry;
|
||||
MemoryValueKey key;
|
||||
MemoryReportKind reportKind = MemoryReportKind::None;
|
||||
};
|
||||
|
||||
struct MemoryReportEntry {
|
||||
enum class Kind {
|
||||
Core,
|
||||
Batch
|
||||
};
|
||||
|
||||
Kind kind = Kind::Core;
|
||||
uint64_t id = 0;
|
||||
llvm::SmallVector<int32_t, 8> coreIds;
|
||||
MemoryReportRow row;
|
||||
};
|
||||
|
||||
class PimMemory {
|
||||
llvm::SmallVector<PendingMemEntry, 32> memEntries;
|
||||
llvm::SmallDenseMap<MemoryValueKey, MemEntry, 32>& globalMemEntriesMap;
|
||||
llvm::SmallDenseMap<MemoryValueKey, MemEntry, 32> ownedMemEntriesMap;
|
||||
MemoryReportRow reportRow;
|
||||
std::optional<size_t> localArenaSize;
|
||||
|
||||
size_t maxSize = 0; // 0 for unbounded memory
|
||||
size_t startAddress = 0;
|
||||
size_t minAlignment = 4;
|
||||
size_t firstAvailableAddress = 0;
|
||||
|
||||
MemEntry* gatherMemEntry(mlir::Value value);
|
||||
MemEntry* gatherMemEntry(mlir::Value value, std::optional<unsigned> lane = std::nullopt);
|
||||
size_t allocateAddress(size_t size, const MemoryValueKey& key);
|
||||
void allocateGatheredMemory();
|
||||
void allocateMemoryForValue(mlir::Value value, MemEntry& memEntry);
|
||||
void allocateMemoryForValue(const MemoryValueKey& key, MemEntry& memEntry, MemoryReportKind reportKind);
|
||||
|
||||
public:
|
||||
PimMemory(llvm::SmallDenseMap<mlir::Value, MemEntry, 32>& globalMemEntriesMap)
|
||||
PimMemory(llvm::SmallDenseMap<MemoryValueKey, MemEntry, 32>& globalMemEntriesMap)
|
||||
: globalMemEntriesMap(globalMemEntriesMap) {}
|
||||
|
||||
void allocateHost(mlir::ModuleOp moduleOp, mlir::func::FuncOp funcOp);
|
||||
void allocateCore(mlir::Operation* op);
|
||||
void report(llvm::raw_ostream& os);
|
||||
void remove(mlir::Value val);
|
||||
void allocateCore(const CompiledCoreMemoryPlan& plan, std::optional<unsigned> lane = std::nullopt);
|
||||
MemoryReportRow getReportRow() const;
|
||||
|
||||
size_t getFirstAvailableAddress() const { return firstAvailableAddress; }
|
||||
MemEntry getMemEntry(mlir::Value value) const;
|
||||
MemEntry getMemEntry(const MemoryValueKey& key) const;
|
||||
};
|
||||
|
||||
class PimAcceleratorMemory {
|
||||
public:
|
||||
llvm::SmallDenseMap<mlir::Value, MemEntry, 32> memEntriesMap;
|
||||
llvm::SmallDenseMap<MemoryValueKey, MemEntry, 32> memEntriesMap;
|
||||
PimMemory hostMem;
|
||||
|
||||
private:
|
||||
llvm::SmallDenseMap<size_t, PimMemory> deviceMem;
|
||||
std::fstream fileReport;
|
||||
std::optional<MemoryReportRow> hostReportRow;
|
||||
llvm::SmallVector<MemoryReportEntry, 32> reportEntries;
|
||||
uint64_t totalWeightBytes = 0;
|
||||
mutable llvm::DenseMap<mlir::Value, CompiledIndexExpr> compiledIndexExprs;
|
||||
mutable llvm::DenseMap<mlir::Value, CompiledAddressExpr> compiledAddressExprs;
|
||||
|
||||
public:
|
||||
PimAcceleratorMemory()
|
||||
: hostMem(memEntriesMap) {
|
||||
|
||||
std::string outputDir = getOutputDir();
|
||||
if (outputDir.empty())
|
||||
return;
|
||||
|
||||
std::string dialectsDir = outputDir + "/reports/";
|
||||
createDirectory(dialectsDir);
|
||||
std::fstream file(dialectsDir + "/memory_report.txt", std::ios::out);
|
||||
fileReport = std::move(file);
|
||||
}
|
||||
PimAcceleratorMemory();
|
||||
PimAcceleratorMemory(const llvm::SmallDenseMap<MemoryValueKey, MemEntry, 32>& initialMemEntries, bool enableReport);
|
||||
|
||||
PimMemory& getOrCreateDeviceMem(size_t id);
|
||||
|
||||
size_t getValueAddress(mlir::Value value, const StaticValueKnowledge& knowledge = {}) const;
|
||||
size_t getValueAddress(mlir::Value value,
|
||||
const StaticValueKnowledge& knowledge = {},
|
||||
std::optional<unsigned> lane = std::nullopt) const;
|
||||
llvm::FailureOr<int64_t> getIndexValue(mlir::Value value, const StaticValueKnowledge& knowledge = {}) const;
|
||||
void reportHost();
|
||||
void reportCore(size_t coreId);
|
||||
void clean(mlir::Operation* op);
|
||||
void recordCoreReport(size_t coreId, const MemoryReportRow& row);
|
||||
void recordBatchReport(uint64_t batchId, llvm::ArrayRef<int32_t> coreIds, const MemoryReportRow& perCoreRow);
|
||||
void setTotalWeightBytes(uint64_t bytes) { totalWeightBytes = bytes; }
|
||||
void flushReport();
|
||||
};
|
||||
|
||||
struct CoreEmissionJob {
|
||||
mlir::Operation* coreLikeOp = nullptr;
|
||||
const CompiledCoreProgram* program = nullptr;
|
||||
const CompiledCoreMemoryPlan* memoryPlan = nullptr;
|
||||
size_t physicalCoreId = 0;
|
||||
llvm::SmallVector<unsigned, 4> lanes;
|
||||
std::optional<uint64_t> batchReportId;
|
||||
};
|
||||
|
||||
class PimCodeGen {
|
||||
PimAcceleratorMemory& memory;
|
||||
llvm::raw_fd_ostream& coreFileStream;
|
||||
const llvm::DenseMap<size_t, size_t>& emittedCoreIds;
|
||||
PimInstructionWriter& instructionWriter;
|
||||
llvm::raw_fd_ostream* coreJsonStream;
|
||||
std::optional<unsigned> batchLane;
|
||||
mutable std::array<std::optional<int32_t>, 256> scalarRegisterValues = {};
|
||||
mutable std::optional<std::array<int32_t, 2>> vectorBitwidths;
|
||||
|
||||
size_t addressOf(mlir::Value value, const StaticValueKnowledge& knowledge) const {
|
||||
return memory.getValueAddress(value, knowledge);
|
||||
return memory.getValueAddress(value, knowledge, batchLane);
|
||||
}
|
||||
size_t remapCoreId(size_t coreId) const;
|
||||
|
||||
static llvm::json::Object createEmptyOffset();
|
||||
void emitInstruction(llvm::json::Object instruction) const;
|
||||
void emitInstruction(const pim_binary::InstructionRecord& instruction) const;
|
||||
void updateScalarRegisterCache(const pim_binary::InstructionRecord& instruction) const;
|
||||
void ensureVectorBitwidth(int32_t inputBitwidth, int32_t outputBitwidth) const;
|
||||
|
||||
void genSetRegisterImmediate(uint8_t registerNumber, int32_t immediate) const;
|
||||
void genSetRegisterImmediateUnsigned(size_t registerNumber, size_t immediate) 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 setupRdRs1Rs2(
|
||||
size_t rdAddress, size_t rdOffset, size_t rs1Address, size_t rs1Offset, size_t rs2Address, size_t rs2Offset) const;
|
||||
|
||||
void emitMemCopyOp(mlir::StringRef opName,
|
||||
void emitMemCopyOp(pim_binary::Opcode opcode,
|
||||
size_t rdAddr,
|
||||
size_t rdOffset,
|
||||
size_t rs1Addr,
|
||||
size_t rs1Offset,
|
||||
size_t size,
|
||||
mlir::StringRef sizeFieldName = "size") const;
|
||||
void emitCommunicationOp(mlir::StringRef opName, size_t bufferAddr, size_t coreId, size_t size) const;
|
||||
void emitCommunicationOp(pim_binary::Opcode opcode, size_t bufferAddr, size_t coreId, size_t size) const;
|
||||
void emitMvmOp(size_t groupId, size_t rdAddr, size_t rdOffset, size_t rs1Addr, size_t rs1Offset) const;
|
||||
|
||||
public:
|
||||
PimCodeGen(PimAcceleratorMemory& memory,
|
||||
llvm::raw_fd_ostream& coreJson,
|
||||
const llvm::DenseMap<size_t, size_t>& emittedCoreIds)
|
||||
: memory(memory), coreFileStream(coreJson), emittedCoreIds(emittedCoreIds) {}
|
||||
void emitBinaryVectorOp(pim_binary::Opcode opcode,
|
||||
mlir::Value output,
|
||||
mlir::Value lhs,
|
||||
mlir::Value rhs,
|
||||
const StaticValueKnowledge& knowledge) const;
|
||||
void emitUnaryVectorOp(pim_binary::Opcode opcode,
|
||||
mlir::Value output,
|
||||
mlir::Value input,
|
||||
const StaticValueKnowledge& knowledge,
|
||||
int32_t r2OrImm = 0,
|
||||
int32_t generic1 = 0) const;
|
||||
|
||||
PimCodeGen(PimAcceleratorMemory& memory, PimInstructionWriter& instructionWriter, llvm::raw_fd_ostream* coreJson)
|
||||
: memory(memory), instructionWriter(instructionWriter), coreJsonStream(coreJson) {}
|
||||
|
||||
void setBatchLane(std::optional<unsigned> lane) { batchLane = lane; }
|
||||
llvm::FailureOr<int64_t> indexOf(mlir::Value value, const StaticValueKnowledge& knowledge) const {
|
||||
return memory.getIndexValue(value, knowledge);
|
||||
}
|
||||
|
||||
void codeGenLoadOp(pim::PimMemCopyHostToDevOp loadOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenStoreOp(pim::PimMemCopyDevToHostOp storeOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenLmvOp(pim::PimMemCopyOp lmvOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVMVOp(pim::PimVMVOp vmvOp, const StaticValueKnowledge& knowledge) const;
|
||||
|
||||
void codeGenReceiveOp(pim::PimReceiveOp receiveOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenSendOp(pim::PimSendOp sendOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenWaitOp(pim::PimWaitOp waitOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenSyncOp(pim::PimSyncOp syncOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenConcatOp(pim::PimConcatOp concatOp, const StaticValueKnowledge& knowledge) const;
|
||||
|
||||
template <typename MVMTy>
|
||||
void codeGenMVMLikeOp(size_t mvmId, MVMTy mvmLikeOp, bool transposeMatrix, const StaticValueKnowledge& knowledge);
|
||||
|
||||
void codeGenVVAddOp(pim::PimVVAddOp vvaddOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVVSubOp(pim::PimVVSubOp vvsubOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVVMulOp(pim::PimVVMulOp vvmulOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVVMaxOp(pim::PimVVMaxOp vvmaxOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVVDMulOp(pim::PimVVDMulOp vvdmulOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVAvgOp(pim::PimVAvgOp vavgOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVReluOp(pim::PimVReluOp vreluOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVTanhOp(pim::PimVTanhOp vtanhOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVSigmOp(pim::PimVSigmOp vsigmOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenVSoftmaxOp(pim::PimVSoftmaxOp vsoftmaxOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGetGlobalOp(mlir::memref::GetGlobalOp getGlobalOp, const StaticValueKnowledge& knowledge) const;
|
||||
void codeGenTransposeOp(pim::PimTransposeOp transposeOp, const StaticValueKnowledge& knowledge) const;
|
||||
};
|
||||
|
||||
OnnxMlirCompilerErrorCodes compileToPimJson(mlir::ModuleOp& moduleOpRef, std::string& outputDirName);
|
||||
OnnxMlirCompilerErrorCodes compileToPimCode(mlir::ModuleOp& moduleOpRef, std::string& outputDirName);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
namespace llvm {
|
||||
|
||||
template <>
|
||||
struct DenseMapInfo<onnx_mlir::MemoryValueKey> {
|
||||
static onnx_mlir::MemoryValueKey getEmptyKey() { return {DenseMapInfo<mlir::Value>::getEmptyKey(), 0}; }
|
||||
|
||||
static onnx_mlir::MemoryValueKey getTombstoneKey() { return {DenseMapInfo<mlir::Value>::getTombstoneKey(), 0}; }
|
||||
|
||||
static unsigned getHashValue(const onnx_mlir::MemoryValueKey& key) {
|
||||
return hash_combine(key.value, key.lane.value_or(std::numeric_limits<unsigned>::max()));
|
||||
}
|
||||
|
||||
static bool isEqual(const onnx_mlir::MemoryValueKey& lhs, const onnx_mlir::MemoryValueKey& rhs) { return lhs == rhs; }
|
||||
};
|
||||
|
||||
} // namespace llvm
|
||||
|
||||
@@ -1,18 +1,9 @@
|
||||
/*
|
||||
* SPDX-License-Identifier: Apache-2.0
|
||||
*/
|
||||
#include "llvm/Support/ErrorHandling.h"
|
||||
|
||||
//===------------------------- PimCompilerOptions.cpp --------------------===//
|
||||
//
|
||||
// Copyright 2022 The IBM Research Authors.
|
||||
//
|
||||
// =============================================================================
|
||||
//
|
||||
// Compiler Options for PIM
|
||||
//
|
||||
//===----------------------------------------------------------------------===//
|
||||
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
|
||||
|
||||
#include <limits>
|
||||
|
||||
#define DEBUG_TYPE "PimCompilerOptions"
|
||||
|
||||
namespace onnx_mlir {
|
||||
@@ -26,36 +17,134 @@ llvm::cl::opt<PimEmissionTargetType> pimEmissionTarget(
|
||||
llvm::cl::init(EmitPimCodegen),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<PimMemoryReportLevel> pimMemoryReport(
|
||||
"pim-memory-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(PimMemoryReportSummary, "summary", "Emit a concise PIM memory summary")),
|
||||
llvm::cl::init(PimMemoryReportSummary),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<PimConvLoweringType> pimConvLowering(
|
||||
"pim-conv-lowering",
|
||||
llvm::cl::desc("Convolution lowering strategy for PIM"),
|
||||
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(PimConvLoweringDepthwise, "depthwise", "Force the depthwise-specialized Conv lowering")),
|
||||
llvm::cl::values(
|
||||
clEnumValN(PimConvLoweringPackedIm2Col, "packed-im2col", "Use explicit im2col with packed multi-position GEMM")),
|
||||
llvm::cl::values(clEnumValN(PimConvLoweringStreamedPatch,
|
||||
"streamed-patch",
|
||||
"Use streamed/chunked im2col rows without multi-position packing")),
|
||||
llvm::cl::values(clEnumValN(PimConvLoweringStreamedPacked,
|
||||
"streamed-packed",
|
||||
"Use streamed/chunked im2col rows with packed multi-position GEMM")),
|
||||
llvm::cl::values(clEnumValN(PimConvLoweringOutputChannelTiled,
|
||||
"output-channel-tiled",
|
||||
"Force Conv lowering that relies on Gemm output-channel tiling")),
|
||||
llvm::cl::values(
|
||||
clEnumValN(PimConvLoweringInputKTiled, "input-k-tiled", "Force Conv lowering that relies on Gemm K tiling")),
|
||||
llvm::cl::values(clEnumValN(PimConvLoweringTiled2D,
|
||||
"tiled-2d",
|
||||
"Force Conv lowering that relies on Gemm 2D K/C tiling")),
|
||||
llvm::cl::init(PimConvLoweringAuto),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<PimSpatialDataflowExportType> pimExportSpatialDataflow(
|
||||
"pim-export-spatial-dataflow",
|
||||
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(SpatialDataflowExportSpatial1, "spatial1", "Emit spatial1 graph dataflow CSV reports")),
|
||||
llvm::cl::values(
|
||||
clEnumValN(SpatialDataflowExportSpatial2, "spatial2", "Emit spatial2 trivially merged graph dataflow CSV reports")),
|
||||
llvm::cl::values(
|
||||
clEnumValN(SpatialDataflowExportSpatial3, "spatial3", "Emit spatial3 scheduled dataflow CSV reports")),
|
||||
llvm::cl::values(
|
||||
clEnumValN(SpatialDataflowExportSpatial4, "spatial4", "Emit spatial4 realized dataflow CSV reports")),
|
||||
llvm::cl::values(clEnumValN(SpatialDataflowExportAll, "all", "Emit all Spatial dataflow CSV reports")),
|
||||
llvm::cl::init(SpatialDataflowExportNone),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<bool>
|
||||
pimOnlyCodegen("pim-only-codegen",
|
||||
llvm::cl::desc("Only generate code for PIM (assume input is already in bufferized PIM IR)"),
|
||||
llvm::cl::init(false),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<bool> useExperimentalConvImpl("use-experimental-conv-impl",
|
||||
llvm::cl::desc("Use experimental implementation for convolution"),
|
||||
llvm::cl::init(false),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
llvm::cl::opt<uint64_t> pimConvIm2colMaxElements(
|
||||
"pim-conv-im2col-max-elements",
|
||||
llvm::cl::desc("Maximum number of im2col elements to materialize globally for one Conv before streaming/chunking"),
|
||||
llvm::cl::init(1ull << 20),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<uint64_t> pimConvStreamChunkPositions(
|
||||
"pim-conv-stream-chunk-positions",
|
||||
llvm::cl::desc("Maximum number of Conv output positions to materialize in one streamed chunk"),
|
||||
llvm::cl::init(1024),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<bool> pimReportConvLowering("pim-report-conv-lowering",
|
||||
llvm::cl::desc("Emit a bounded Conv lowering report"),
|
||||
llvm::cl::init(true),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<bool> pimEmitJson("pim-emit-json",
|
||||
llvm::cl::desc("Also emit per-core JSON instruction files alongside binary .pim files"),
|
||||
llvm::cl::init(false),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<bool> pimDetectCommunicationDeadlock(
|
||||
"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::init(false),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<bool> pimVerifyBufferizationCopyFreedom(
|
||||
"pim-verify-bufferization-copy-freedom",
|
||||
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<size_t>
|
||||
crossbarSize("crossbar-size", llvm::cl::desc("Width and heigth of a single crossbar"), llvm::cl::init(2));
|
||||
crossbarSize("crossbar-size", llvm::cl::desc("Width and height of a single crossbar"), llvm::cl::init(128));
|
||||
|
||||
llvm::cl::opt<size_t>
|
||||
crossbarCountInCore("crossbar-count", llvm::cl::desc("Number of crossbars in each core"), llvm::cl::init(256));
|
||||
crossbarCountInCore("crossbar-count", llvm::cl::desc("Number of crossbars in each core"), llvm::cl::init(64));
|
||||
|
||||
llvm::cl::opt<size_t> pipelineStages(
|
||||
"pipeline",
|
||||
llvm::cl::desc("Number of throughput pipeline stages (1 preserves latency scheduling)"),
|
||||
llvm::cl::init(1),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<long> coresCount("core-count",
|
||||
llvm::cl::desc("Number of cores in the chip. `-1` to use the minimum amount of cores."),
|
||||
llvm::cl::desc("Number of cores in the chip. Required for PIM compilation."),
|
||||
llvm::cl::init(-1));
|
||||
|
||||
llvm::cl::opt<size_t> dcpCriticalWindowSize(
|
||||
"dcp-critical-window-size",
|
||||
llvm::cl::desc("Number of lowest-slack virtual nodes considered by each DCP coarsening iteration. "
|
||||
"Use 0 to run the legacy full-graph DCP analysis."),
|
||||
llvm::cl::init(4000));
|
||||
llvm::cl::opt<std::string> pimTargetConfig(
|
||||
"pim-target-config",
|
||||
llvm::cl::desc("PIM target configuration used to construct the Spatial scheduling cost model"),
|
||||
llvm::cl::init(""),
|
||||
llvm::cl::cat(OnnxMlirOptions));
|
||||
|
||||
llvm::cl::opt<bool>
|
||||
ignoreConcatError("ignore-concat-error",
|
||||
llvm::cl::desc("Ignore ConcatOp corner case: do not assert and do a simplification"),
|
||||
llvm::cl::init(false));
|
||||
bool hasExplicitPimCoreCount() { return coresCount.getNumOccurrences() != 0; }
|
||||
|
||||
void verifyExplicitPimCoreCount() {
|
||||
if (!hasExplicitPimCoreCount())
|
||||
llvm::report_fatal_error("PIM compilation requires an explicit --core-count=<positive integer>");
|
||||
if (coresCount.getValue() <= 0)
|
||||
llvm::report_fatal_error("PIM compilation requires --core-count to be a positive integer");
|
||||
}
|
||||
|
||||
void verifyPimPipelineStages() {
|
||||
if (pipelineStages.getValue() == 0)
|
||||
llvm::report_fatal_error("PIM compilation requires --pipeline to be positive");
|
||||
if (static_cast<size_t>(coresCount.getValue()) % pipelineStages.getValue() != 0)
|
||||
llvm::report_fatal_error("PIM compilation requires --core-count to be divisible by --pipeline");
|
||||
if (crossbarCountInCore.getValue()
|
||||
> std::numeric_limits<size_t>::max() / pipelineStages.getValue())
|
||||
llvm::report_fatal_error("PIM compilation --crossbar-count * --pipeline overflows");
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -2,6 +2,8 @@
|
||||
|
||||
#include "llvm/Support/CommandLine.h"
|
||||
|
||||
#include <string>
|
||||
|
||||
#define INSTRUMENTSTAGE_ENUM_PIM
|
||||
|
||||
#define INSTRUMENTSTAGE_CL_ENUM_PIM
|
||||
@@ -20,23 +22,54 @@ typedef enum {
|
||||
EmitPimCodegen = 3
|
||||
} PimEmissionTargetType;
|
||||
|
||||
typedef enum {
|
||||
PimMemoryReportNone = 0,
|
||||
PimMemoryReportSummary = 1,
|
||||
} PimMemoryReportLevel;
|
||||
|
||||
typedef enum {
|
||||
PimConvLoweringAuto = 0,
|
||||
PimConvLoweringLegacy = 1,
|
||||
PimConvLoweringDepthwise = 2,
|
||||
PimConvLoweringPackedIm2Col = 3,
|
||||
PimConvLoweringStreamedPatch = 4,
|
||||
PimConvLoweringStreamedPacked = 5,
|
||||
PimConvLoweringOutputChannelTiled = 6,
|
||||
PimConvLoweringInputKTiled = 7,
|
||||
PimConvLoweringTiled2D = 8,
|
||||
} PimConvLoweringType;
|
||||
|
||||
typedef enum {
|
||||
SpatialDataflowExportNone = 0,
|
||||
SpatialDataflowExportSpatial1 = 1,
|
||||
SpatialDataflowExportSpatial2 = 2,
|
||||
SpatialDataflowExportSpatial3 = 3,
|
||||
SpatialDataflowExportSpatial4 = 4,
|
||||
SpatialDataflowExportAll = 5,
|
||||
} PimSpatialDataflowExportType;
|
||||
|
||||
extern llvm::cl::OptionCategory OnnxMlirOptions;
|
||||
extern llvm::cl::opt<PimEmissionTargetType> pimEmissionTarget;
|
||||
extern llvm::cl::opt<PimMemoryReportLevel> pimMemoryReport;
|
||||
extern llvm::cl::opt<PimConvLoweringType> pimConvLowering;
|
||||
extern llvm::cl::opt<PimSpatialDataflowExportType> pimExportSpatialDataflow;
|
||||
|
||||
extern llvm::cl::opt<bool> pimOnlyCodegen;
|
||||
extern llvm::cl::opt<bool> useExperimentalConvImpl;
|
||||
extern llvm::cl::opt<bool> pimEmitJson;
|
||||
extern llvm::cl::opt<bool> pimReportConvLowering;
|
||||
extern llvm::cl::opt<bool> pimDetectCommunicationDeadlock;
|
||||
extern llvm::cl::opt<bool> pimVerifyBufferizationCopyFreedom;
|
||||
|
||||
extern llvm::cl::opt<size_t> crossbarSize;
|
||||
extern llvm::cl::opt<size_t> crossbarCountInCore;
|
||||
extern llvm::cl::opt<size_t> pipelineStages;
|
||||
extern llvm::cl::opt<long> coresCount;
|
||||
extern llvm::cl::opt<size_t> dcpCriticalWindowSize;
|
||||
extern llvm::cl::opt<std::string> pimTargetConfig;
|
||||
extern llvm::cl::opt<uint64_t> pimConvIm2colMaxElements;
|
||||
extern llvm::cl::opt<uint64_t> pimConvStreamChunkPositions;
|
||||
|
||||
// This option, by default set to false, will ignore an error when resolving a
|
||||
// specific tiles of the operands of a concat. This specific case is when the
|
||||
// wanted tile is generated by two separate operands of the concat. If this is
|
||||
// set to false, this corner case will assert an error. If this is set to true,
|
||||
// a simplification is performed and only the tile from the first operand is
|
||||
// taken.
|
||||
extern llvm::cl::opt<bool> ignoreConcatError;
|
||||
bool hasExplicitPimCoreCount();
|
||||
void verifyExplicitPimCoreCount();
|
||||
void verifyPimPipelineStages();
|
||||
|
||||
} // namespace onnx_mlir
|
||||
|
||||
@@ -1,9 +1,25 @@
|
||||
#include "mlir/Conversion/AffineToStandard/AffineToStandard.h"
|
||||
#include "mlir/Transforms/Passes.h"
|
||||
|
||||
#include "llvm/Support/Error.h"
|
||||
#include "llvm/Support/ErrorHandling.h"
|
||||
#include "llvm/Support/JSON.h"
|
||||
#include "llvm/Support/MemoryBuffer.h"
|
||||
#include "llvm/Support/Path.h"
|
||||
|
||||
#include "llvm/ADT/SmallString.h"
|
||||
#include <cmath>
|
||||
#include <limits>
|
||||
#include <tuple>
|
||||
|
||||
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimCompilerUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/ONNXToSpatialOptions.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp"
|
||||
#include "src/Accelerators/PIM/Pass/PIMPasses.h"
|
||||
#include "src/Accelerators/PIM/Dialect/Spatial/SpatialTargetResources.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Spatial/Passes/Transforms/MergeComputeNodes/SpatialDataflowCsvExporter.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Spatial/Passes/Transforms/MergeComputeNodes/Scheduling/SchedulingTarget.hpp"
|
||||
#include "src/Accelerators/PIM/Passes/PIMPasses.h"
|
||||
#include "src/Compiler/CompilerPasses.hpp"
|
||||
|
||||
#define DEBUG_TYPE "PimCompilerUtils"
|
||||
@@ -13,13 +29,316 @@ using namespace onnx_mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
namespace {
|
||||
|
||||
void setDefaultPimInterProcessorLatencies(
|
||||
spatial::SchedulingTarget& target) {
|
||||
size_t rows = static_cast<size_t>(
|
||||
std::sqrt(static_cast<long double>(target.processorCount)));
|
||||
while (rows > 1 && target.processorCount % rows != 0)
|
||||
--rows;
|
||||
size_t columns = (target.processorCount + rows - 1) / rows;
|
||||
|
||||
target.interProcessorLatencyNs.assign(
|
||||
target.processorCount * target.processorCount, 0);
|
||||
Cost latencySum = 0;
|
||||
size_t pairCount = 0;
|
||||
for (size_t source = 0; source < target.processorCount; ++source) {
|
||||
for (size_t destination = 0;
|
||||
destination < target.processorCount; ++destination) {
|
||||
if (source == destination)
|
||||
continue;
|
||||
size_t sourceRow = source / columns;
|
||||
size_t sourceColumn = source % columns;
|
||||
size_t destinationRow = destination / columns;
|
||||
size_t destinationColumn = destination % columns;
|
||||
size_t rowDistance = sourceRow > destinationRow
|
||||
? sourceRow - destinationRow
|
||||
: destinationRow - sourceRow;
|
||||
size_t columnDistance = sourceColumn > destinationColumn
|
||||
? sourceColumn - destinationColumn
|
||||
: destinationColumn - sourceColumn;
|
||||
Cost latency = static_cast<Cost>(2 + rowDistance + columnDistance);
|
||||
target.interProcessorLatencyNs[
|
||||
source * target.processorCount + destination] = latency;
|
||||
latencySum = checkedAdd(latencySum, latency);
|
||||
++pairCount;
|
||||
}
|
||||
}
|
||||
target.averageInterProcessorLatencyNs =
|
||||
pairCount == 0
|
||||
? 0
|
||||
: (latencySum + static_cast<Cost>(pairCount) - 1)
|
||||
/ static_cast<Cost>(pairCount);
|
||||
}
|
||||
|
||||
spatial::SchedulingTarget getDefaultPimSchedulingTarget() {
|
||||
spatial::SchedulingTarget target;
|
||||
target.processorCount = static_cast<size_t>(coresCount.getValue());
|
||||
target.residentWeightCapacity = crossbarCountInCore.getValue();
|
||||
target.matrixRows = crossbarSize.getValue();
|
||||
target.matrixColumns = crossbarSize.getValue();
|
||||
|
||||
setDefaultPimInterProcessorLatencies(target);
|
||||
return target;
|
||||
}
|
||||
|
||||
spatial::ConvLoweringStrategy getSpatialConvLoweringStrategy(PimConvLoweringType strategy) {
|
||||
switch (strategy) {
|
||||
case PimConvLoweringAuto: return spatial::ConvLoweringStrategy::Auto;
|
||||
case PimConvLoweringLegacy: return spatial::ConvLoweringStrategy::Legacy;
|
||||
case PimConvLoweringDepthwise: return spatial::ConvLoweringStrategy::Depthwise;
|
||||
case PimConvLoweringPackedIm2Col: return spatial::ConvLoweringStrategy::PackedIm2Col;
|
||||
case PimConvLoweringStreamedPatch: return spatial::ConvLoweringStrategy::StreamedPatch;
|
||||
case PimConvLoweringStreamedPacked: return spatial::ConvLoweringStrategy::StreamedPacked;
|
||||
case PimConvLoweringOutputChannelTiled: return spatial::ConvLoweringStrategy::OutputChannelTiled;
|
||||
case PimConvLoweringInputKTiled: return spatial::ConvLoweringStrategy::InputKTiled;
|
||||
case PimConvLoweringTiled2D: return spatial::ConvLoweringStrategy::Tiled2D;
|
||||
}
|
||||
llvm_unreachable("unknown PIM Conv lowering strategy");
|
||||
}
|
||||
|
||||
spatial::SpatialDataflowExportStage getPimSpatialDataflowExportStage(
|
||||
PimSpatialDataflowExportType stage) {
|
||||
switch (stage) {
|
||||
case SpatialDataflowExportNone: return spatial::SpatialDataflowExportStage::None;
|
||||
case SpatialDataflowExportSpatial1: return spatial::SpatialDataflowExportStage::Spatial1;
|
||||
case SpatialDataflowExportSpatial2: return spatial::SpatialDataflowExportStage::Spatial2;
|
||||
case SpatialDataflowExportSpatial3: return spatial::SpatialDataflowExportStage::Spatial3;
|
||||
case SpatialDataflowExportSpatial4: return spatial::SpatialDataflowExportStage::Spatial4;
|
||||
case SpatialDataflowExportAll: return spatial::SpatialDataflowExportStage::All;
|
||||
}
|
||||
llvm_unreachable("unknown PIM Spatial dataflow export stage");
|
||||
}
|
||||
|
||||
spatial::SpatialTargetResources getPimSpatialTargetResources(const spatial::SchedulingTarget& target) {
|
||||
spatial::SpatialTargetResources resources;
|
||||
resources.matrixShape = {target.matrixRows, target.matrixColumns};
|
||||
resources.matrixUnitsPerProcessor = target.residentWeightCapacity;
|
||||
resources.processorCount = target.processorCount;
|
||||
resources.vectorWidth = target.vectorWidth;
|
||||
if (failed(resources.verify()))
|
||||
llvm::report_fatal_error("PIM target resources are incomplete");
|
||||
return resources;
|
||||
}
|
||||
|
||||
ONNXToSpatialPlanningOptions getPimONNXToSpatialPlanningOptions() {
|
||||
ONNXToSpatialPlanningOptions options;
|
||||
options.convIm2colMaxElements = pimConvIm2colMaxElements.getValue();
|
||||
options.convStreamChunkPositions = pimConvStreamChunkPositions.getValue();
|
||||
options.forcedConvStrategy = getSpatialConvLoweringStrategy(pimConvLowering.getValue());
|
||||
options.reportConvLowering = pimReportConvLowering.getValue();
|
||||
return options;
|
||||
}
|
||||
|
||||
const llvm::json::Object& requireObject(const llvm::json::Object& object,
|
||||
llvm::StringRef key,
|
||||
llvm::StringRef path) {
|
||||
const llvm::json::Object* nested = object.getObject(key);
|
||||
if (!nested)
|
||||
llvm::report_fatal_error("PIM target config is missing object '" + path + "." + key + "'");
|
||||
return *nested;
|
||||
}
|
||||
|
||||
Cost getConfigCost(const llvm::json::Object& object,
|
||||
llvm::StringRef key,
|
||||
Cost fallback,
|
||||
bool allowZero = false) {
|
||||
std::optional<double> number = object.getNumber(key);
|
||||
if (!number)
|
||||
return fallback;
|
||||
if (!std::isfinite(*number) || *number < 0.0 || (!allowZero && *number == 0.0)
|
||||
|| *number > static_cast<double>(std::numeric_limits<Cost>::max()))
|
||||
llvm::report_fatal_error("PIM target config field '" + key + "' must be a valid positive number");
|
||||
return static_cast<Cost>(std::ceil(*number));
|
||||
}
|
||||
|
||||
std::pair<size_t, size_t> getConfigPair(const llvm::json::Object& object,
|
||||
llvm::StringRef key) {
|
||||
const llvm::json::Array* values = object.getArray(key);
|
||||
if (!values || values->size() != 2)
|
||||
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> second = (*values)[1].getAsInteger();
|
||||
if (!first || !second || *first <= 0 || *second <= 0)
|
||||
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)};
|
||||
}
|
||||
|
||||
void loadPimInterProcessorLatencies(
|
||||
spatial::SchedulingTarget& target,
|
||||
const llvm::json::Object& network) {
|
||||
std::optional<llvm::StringRef> filename =
|
||||
network.getString("net_config_file_path");
|
||||
if (!filename)
|
||||
llvm::report_fatal_error(
|
||||
"PIM target config is missing network latency file path");
|
||||
|
||||
llvm::SmallString<256> networkPath(*filename);
|
||||
if (!llvm::sys::path::is_absolute(networkPath)) {
|
||||
llvm::SmallString<256> configDirectory(pimTargetConfig.getValue());
|
||||
llvm::sys::path::remove_filename(configDirectory);
|
||||
llvm::sys::path::append(configDirectory, networkPath);
|
||||
networkPath = configDirectory;
|
||||
}
|
||||
|
||||
auto buffer = llvm::MemoryBuffer::getFile(networkPath);
|
||||
if (!buffer)
|
||||
llvm::report_fatal_error(
|
||||
llvm::Twine("failed to read PIM network config '")
|
||||
+ networkPath + "': " + buffer.getError().message());
|
||||
auto parsed = llvm::json::parse(buffer.get()->getBuffer());
|
||||
if (!parsed)
|
||||
llvm::report_fatal_error(
|
||||
llvm::Twine("failed to parse PIM network config '")
|
||||
+ networkPath + "': " + llvm::toString(parsed.takeError()));
|
||||
const llvm::json::Object* root = parsed->getAsObject();
|
||||
const llvm::json::Object* latencies =
|
||||
root ? root->getObject("latency") : nullptr;
|
||||
if (!latencies)
|
||||
llvm::report_fatal_error(
|
||||
"PIM network config is missing its latency matrix");
|
||||
|
||||
target.interProcessorLatencyNs.assign(
|
||||
target.processorCount * target.processorCount, 0);
|
||||
Cost latencySum = 0;
|
||||
size_t pairCount = 0;
|
||||
for (size_t source = 0; source < target.processorCount; ++source) {
|
||||
std::string sourceKey = std::to_string(source);
|
||||
const llvm::json::Object* row = latencies->getObject(sourceKey);
|
||||
if (!row)
|
||||
llvm::report_fatal_error(
|
||||
llvm::Twine("PIM network config is missing latency row ")
|
||||
+ sourceKey);
|
||||
for (size_t destination = 0;
|
||||
destination < target.processorCount; ++destination) {
|
||||
if (source == destination)
|
||||
continue;
|
||||
std::string destinationKey = std::to_string(destination);
|
||||
std::optional<double> latency = row->getNumber(destinationKey);
|
||||
if (!latency || !std::isfinite(*latency) || *latency <= 0.0)
|
||||
llvm::report_fatal_error(
|
||||
llvm::Twine("PIM network config is missing latency ")
|
||||
+ sourceKey + " -> " + destinationKey);
|
||||
Cost roundedLatency = static_cast<Cost>(std::ceil(*latency));
|
||||
target.interProcessorLatencyNs[
|
||||
source * target.processorCount + destination] = roundedLatency;
|
||||
latencySum = checkedAdd(latencySum, roundedLatency);
|
||||
++pairCount;
|
||||
}
|
||||
}
|
||||
target.averageInterProcessorLatencyNs =
|
||||
pairCount == 0
|
||||
? 0
|
||||
: (latencySum + static_cast<Cost>(pairCount) - 1)
|
||||
/ static_cast<Cost>(pairCount);
|
||||
}
|
||||
|
||||
spatial::SchedulingTarget getPimSchedulingTarget() {
|
||||
spatial::SchedulingTarget target = getDefaultPimSchedulingTarget();
|
||||
if (pimTargetConfig.empty())
|
||||
return target;
|
||||
|
||||
auto buffer = llvm::MemoryBuffer::getFile(pimTargetConfig);
|
||||
if (!buffer)
|
||||
llvm::report_fatal_error(
|
||||
llvm::Twine("failed to read PIM target config '")
|
||||
+ pimTargetConfig.getValue() + "': " + buffer.getError().message());
|
||||
auto parsed = llvm::json::parse(buffer.get()->getBuffer());
|
||||
if (!parsed)
|
||||
llvm::report_fatal_error(
|
||||
llvm::Twine("failed to parse PIM target config '")
|
||||
+ pimTargetConfig.getValue() + "': "
|
||||
+ llvm::toString(parsed.takeError()));
|
||||
const llvm::json::Object* root = parsed->getAsObject();
|
||||
if (!root)
|
||||
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& core = requireObject(chip, "core_config", "chip_config");
|
||||
const llvm::json::Object& matrix =
|
||||
requireObject(core, "matrix_config", "chip_config.core_config");
|
||||
const llvm::json::Object& localMemory =
|
||||
requireObject(core, "local_memory_config", "chip_config.core_config");
|
||||
const llvm::json::Object& network =
|
||||
requireObject(chip, "network_config", "chip_config");
|
||||
|
||||
std::optional<int64_t> coreCount = chip.getInteger("core_cnt");
|
||||
if (!coreCount || *coreCount <= 0)
|
||||
llvm::report_fatal_error("PIM target config field 'core_cnt' must be a positive integer");
|
||||
target.processorCount = static_cast<size_t>(*coreCount);
|
||||
target.residentWeightCapacity =
|
||||
getConfigCost(matrix, "xbar_array_count", target.residentWeightCapacity);
|
||||
std::tie(target.matrixRows, target.matrixColumns) =
|
||||
getConfigPair(matrix, "xbar_size");
|
||||
|
||||
if (target.processorCount != static_cast<size_t>(coresCount.getValue())
|
||||
|| target.residentWeightCapacity != crossbarCountInCore.getValue()
|
||||
|| target.matrixRows != crossbarSize.getValue()
|
||||
|| target.matrixColumns != crossbarSize.getValue())
|
||||
llvm::report_fatal_error("PIM target config resources do not match --core-count, "
|
||||
"--crossbar-count, and --crossbar-size");
|
||||
loadPimInterProcessorLatencies(target, network);
|
||||
|
||||
target.processorPeriodNs =
|
||||
getConfigCost(core, "period", target.processorPeriodNs);
|
||||
target.localMemoryWidthBytes =
|
||||
getConfigCost(localMemory, "data_width", target.localMemoryWidthBytes);
|
||||
target.localMemoryReadLatencyCycles =
|
||||
getConfigCost(localMemory, "read_latency_cycle", target.localMemoryReadLatencyCycles);
|
||||
target.localMemoryWriteLatencyCycles =
|
||||
getConfigCost(localMemory, "write_latency_cycle", target.localMemoryWriteLatencyCycles);
|
||||
target.transferWidthBytes =
|
||||
getConfigCost(network, "bus_width", target.transferWidthBytes);
|
||||
target.vectorWidth = getConfigCost(core, "vector_width", target.vectorWidth);
|
||||
target.vectorLatencyCycles =
|
||||
getConfigCost(core, "vector_latency_cycle", target.vectorLatencyCycles);
|
||||
|
||||
target.matrixPeriodNs =
|
||||
getConfigCost(matrix, "period", target.matrixPeriodNs);
|
||||
target.matrixInputResolutionBits =
|
||||
getConfigCost(matrix, "dac_resolution", target.matrixInputResolutionBits);
|
||||
target.matrixInputLatencyCycles =
|
||||
getConfigCost(matrix, "dac_latency_cycle", target.matrixInputLatencyCycles);
|
||||
target.matrixInputParallelism =
|
||||
getConfigCost(matrix, "dac_count", target.matrixInputParallelism);
|
||||
target.matrixReadLatencyNs =
|
||||
getConfigCost(matrix, "xbar_latency", target.matrixReadLatencyNs);
|
||||
target.matrixSampleLatencyCycles =
|
||||
getConfigCost(matrix, "sample_hold_latency_cycle", target.matrixSampleLatencyCycles);
|
||||
target.matrixOutputLatencyCycles =
|
||||
getConfigCost(matrix, "adc_latency_cycle", target.matrixOutputLatencyCycles);
|
||||
target.matrixOutputParallelism =
|
||||
getConfigCost(matrix, "adc_count", target.matrixOutputParallelism);
|
||||
target.matrixShiftLatencyCycles =
|
||||
getConfigCost(matrix, "shift_adder_latency_cycle", target.matrixShiftLatencyCycles);
|
||||
target.matrixBufferLatencyCycles =
|
||||
getConfigCost(matrix, "output_buffer_latency_cycle", target.matrixBufferLatencyCycles);
|
||||
target.matrixInputBufferLatencyCycles =
|
||||
getConfigCost(matrix,
|
||||
"input_buffer_latency_cycle",
|
||||
target.matrixInputBufferLatencyCycles,
|
||||
/*allowZero=*/true);
|
||||
target.matrixPipeline = matrix.getBoolean("pipeline_mode").value_or(target.matrixPipeline);
|
||||
return target;
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
void addPassesPim(OwningOpRef<ModuleOp>& module,
|
||||
PassManager& pm,
|
||||
EmissionTargetType& emissionTarget,
|
||||
std::string outputNameNoExt) {
|
||||
verifyExplicitPimCoreCount();
|
||||
verifyPimPipelineStages();
|
||||
spatial::SchedulingTarget schedulingTarget = getPimSchedulingTarget();
|
||||
spatial::SpatialTargetResources targetResources = getPimSpatialTargetResources(schedulingTarget);
|
||||
|
||||
if (pimOnlyCodegen) {
|
||||
// Skip all the lowering passes and directly generate code for PIM.
|
||||
pm.addPass(createPimInstructionSelectionPass());
|
||||
pm.addPass(createPimLocalMemoryPlanningPass());
|
||||
pm.addPass(createPimVerificationPass(targetResources, pimDetectCommunicationDeadlock.getValue()));
|
||||
pm.addPass(createEmitPimCodePass());
|
||||
return;
|
||||
}
|
||||
|
||||
@@ -27,33 +346,44 @@ void addPassesPim(OwningOpRef<ModuleOp>& module,
|
||||
addONNXToMLIRPasses(pm, /*target CPU*/ false);
|
||||
|
||||
if (pimEmissionTarget >= EmitSpatial) {
|
||||
pm.addPass(createONNXToSpatialPass());
|
||||
pm.addPass(createMergeComputeNodesPass());
|
||||
// pm.addPass(createCountInstructionPass());
|
||||
ONNXToSpatialPlanningOptions planningOptions = getPimONNXToSpatialPlanningOptions();
|
||||
spatial::SpatialDataflowExportStage exportStage =
|
||||
getPimSpatialDataflowExportStage(pimExportSpatialDataflow.getValue());
|
||||
pm.addPass(createONNXToSpatialPass(targetResources, planningOptions));
|
||||
pm.addPass(createSpatialLayoutPlanningPass(targetResources));
|
||||
pm.addPass(createLowerSpatialPlansPass(targetResources, planningOptions, exportStage));
|
||||
pm.addPass(createTrivialGraphComputeMergePass(
|
||||
schedulingTarget.residentWeightCapacity, exportStage));
|
||||
pm.addPass(spatial::createScheduleAndRealizeSpatialPass(
|
||||
schedulingTarget, exportStage, pipelineStages.getValue()));
|
||||
pm.addPass(createMessagePass("Onnx lowered to Spatial"));
|
||||
}
|
||||
|
||||
if (pimEmissionTarget >= EmitPim) {
|
||||
pm.addPass(createSpatialToPimPass());
|
||||
// pm.addPass(createCountInstructionPass());
|
||||
pm.addPass(createSpatialToPimPass(targetResources));
|
||||
pm.addPass(createMessagePass("Spatial lowered to Pim"));
|
||||
}
|
||||
|
||||
if (pimEmissionTarget >= EmitPimBufferized) {
|
||||
pm.addPass(createPimBufferizationPass());
|
||||
// pm.addPass(createCountInstructionPass());
|
||||
pm.addPass(createPimBufferizationPreparationPass(pimVerifyBufferizationCopyFreedom.getValue()));
|
||||
pm.addPass(createPimOneShotBufferizationPass());
|
||||
pm.addPass(createPimMemoryNormalizationPass());
|
||||
pm.addPass(createPimBufferizationVerificationPass());
|
||||
pm.addPass(createMessagePass("Pim bufferized"));
|
||||
}
|
||||
|
||||
if (pimEmissionTarget >= EmitPimCodegen) {
|
||||
pm.addPass(mlir::createLowerAffinePass());
|
||||
pm.addPass(createPimHostConstantFoldingPass());
|
||||
pm.addPass(createMessagePass("Pim host constants folded"));
|
||||
pm.addPass(createPimMaterializeHostConstantsPass());
|
||||
pm.addPass(createPimVerificationPass());
|
||||
pm.addPass(createPimInstructionSelectionPass());
|
||||
pm.addPass(createMessagePass("Pim instructions selected"));
|
||||
pm.addPass(createPimLocalMemoryPlanningPass());
|
||||
pm.addPass(createMessagePass("Pim local memory planned"));
|
||||
pm.addPass(createPimVerificationPass(targetResources, pimDetectCommunicationDeadlock.getValue()));
|
||||
pm.addPass(createMessagePass("Pim verified"));
|
||||
pm.addPass(createEmitPimJsonPass());
|
||||
// pm.addPass(createCountInstructionPass());
|
||||
pm.addPass(createMessagePass("Pim json code emitted"));
|
||||
pm.addPass(createEmitPimCodePass());
|
||||
pm.addPass(createMessagePass("Pim code emitted"));
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
#include "mlir/Dialect/MemRef/IR/MemRef.h"
|
||||
#include "mlir/Dialect/SCF/IR/SCF.h"
|
||||
#include "src/Accelerators/PIM/Common/IR/CoreBlockUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimCoreProgram.hpp"
|
||||
#include "src/Accelerators/PIM/Dialect/Pim/PimOps.hpp"
|
||||
|
||||
using namespace llvm;
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
namespace {
|
||||
|
||||
static FailureOr<CompiledCoreOpKind> classifyCompiledCoreOpKind(Operation& op) {
|
||||
if (isa<pim::PimMemCopyHostToDevOp>(op)) return CompiledCoreOpKind::Load;
|
||||
if (isa<pim::PimMemCopyDevToHostOp>(op)) return CompiledCoreOpKind::Store;
|
||||
if (isa<pim::PimMemCopyOp>(op)) return CompiledCoreOpKind::Lmv;
|
||||
if (isa<pim::PimVMVOp>(op)) return CompiledCoreOpKind::VMV;
|
||||
if (isa<pim::PimReceiveOp>(op)) return CompiledCoreOpKind::Receive;
|
||||
if (isa<pim::PimSendOp>(op)) return CompiledCoreOpKind::Send;
|
||||
if (isa<pim::PimWaitOp>(op)) return CompiledCoreOpKind::Wait;
|
||||
if (isa<pim::PimSyncOp>(op)) return CompiledCoreOpKind::Sync;
|
||||
if (isa<pim::PimConcatOp>(op)) return CompiledCoreOpKind::Concat;
|
||||
if (isa<pim::PimVMMOp>(op)) return CompiledCoreOpKind::Vmm;
|
||||
if (isa<pim::PimVVAddOp>(op)) return CompiledCoreOpKind::VVAdd;
|
||||
if (isa<pim::PimVVSubOp>(op)) return CompiledCoreOpKind::VVSub;
|
||||
if (isa<pim::PimVVMulOp>(op)) return CompiledCoreOpKind::VVMul;
|
||||
if (isa<pim::PimVVMaxOp>(op)) return CompiledCoreOpKind::VVMax;
|
||||
if (isa<pim::PimVVDMulOp>(op)) return CompiledCoreOpKind::VVDMul;
|
||||
if (isa<pim::PimVAvgOp>(op)) return CompiledCoreOpKind::VAvg;
|
||||
if (isa<pim::PimVReluOp>(op)) return CompiledCoreOpKind::VRelu;
|
||||
if (isa<pim::PimVTanhOp>(op)) return CompiledCoreOpKind::VTanh;
|
||||
if (isa<pim::PimVSigmOp>(op)) return CompiledCoreOpKind::VSigm;
|
||||
if (isa<pim::PimVSoftmaxOp>(op)) return CompiledCoreOpKind::VSoftmax;
|
||||
return failure();
|
||||
}
|
||||
|
||||
static LogicalResult compileCoreEmissionPlan(Block& block, SmallVectorImpl<CompiledCoreNode>& plan) {
|
||||
for (Operation& op : block) {
|
||||
if (isa<pim::PimHaltOp, scf::YieldOp, memref::GetGlobalOp>(op) || isCoreStaticAddressOp(&op))
|
||||
continue;
|
||||
if (auto loadOp = dyn_cast<memref::LoadOp>(op); loadOp && succeeded(compileIndexExpr(loadOp.getResult())))
|
||||
continue;
|
||||
|
||||
if (auto forOp = dyn_cast<scf::ForOp>(op)) {
|
||||
auto lower = compileIndexExpr(forOp.getLowerBound());
|
||||
auto upper = compileIndexExpr(forOp.getUpperBound());
|
||||
auto step = compileIndexExpr(forOp.getStep());
|
||||
if (failed(lower) || failed(upper) || failed(step)) {
|
||||
forOp.emitOpError("requires statically evaluable scf.for bounds for PIM codegen");
|
||||
return failure();
|
||||
}
|
||||
CompiledCoreNode node;
|
||||
node.kind = CompiledCoreNode::Kind::Loop;
|
||||
node.op = forOp;
|
||||
node.lowerBound = *lower;
|
||||
node.upperBound = *upper;
|
||||
node.step = *step;
|
||||
node.loopBody = std::make_unique<SmallVector<CompiledCoreNode, 8>>();
|
||||
if (failed(compileCoreEmissionPlan(forOp.getRegion().front(), *node.loopBody))) return failure();
|
||||
plan.push_back(std::move(node));
|
||||
continue;
|
||||
}
|
||||
if (auto ifOp = dyn_cast<scf::IfOp>(op)) {
|
||||
auto condition = compileIndexExpr(ifOp.getCondition());
|
||||
if (failed(condition)) {
|
||||
ifOp.emitOpError("requires statically evaluable scf.if condition for PIM codegen");
|
||||
return failure();
|
||||
}
|
||||
CompiledCoreNode node;
|
||||
node.kind = CompiledCoreNode::Kind::If;
|
||||
node.op = ifOp;
|
||||
node.condition = *condition;
|
||||
node.thenBody = std::make_unique<SmallVector<CompiledCoreNode, 8>>();
|
||||
node.elseBody = std::make_unique<SmallVector<CompiledCoreNode, 8>>();
|
||||
if (failed(compileCoreEmissionPlan(ifOp.getThenRegion().front(), *node.thenBody))) return failure();
|
||||
if (!ifOp.getElseRegion().empty()
|
||||
&& failed(compileCoreEmissionPlan(ifOp.getElseRegion().front(), *node.elseBody)))
|
||||
return failure();
|
||||
plan.push_back(std::move(node));
|
||||
continue;
|
||||
}
|
||||
if (auto switchOp = dyn_cast<scf::IndexSwitchOp>(op)) {
|
||||
auto selector = compileIndexExpr(switchOp.getArg());
|
||||
if (failed(selector)) {
|
||||
switchOp.emitOpError("requires a statically evaluable scf.index_switch selector for PIM codegen");
|
||||
return failure();
|
||||
}
|
||||
CompiledCoreNode node;
|
||||
node.kind = CompiledCoreNode::Kind::IndexSwitch;
|
||||
node.op = switchOp;
|
||||
node.condition = *selector;
|
||||
llvm::append_range(node.caseValues, switchOp.getCases());
|
||||
for (Region& region : switchOp.getCaseRegions()) {
|
||||
auto body = std::make_unique<SmallVector<CompiledCoreNode, 8>>();
|
||||
if (failed(compileCoreEmissionPlan(region.front(), *body))) return failure();
|
||||
node.caseBodies.push_back(std::move(body));
|
||||
}
|
||||
node.defaultBody = std::make_unique<SmallVector<CompiledCoreNode, 8>>();
|
||||
if (failed(compileCoreEmissionPlan(switchOp.getDefaultRegion().front(), *node.defaultBody))) return failure();
|
||||
plan.push_back(std::move(node));
|
||||
continue;
|
||||
}
|
||||
|
||||
auto opKind = classifyCompiledCoreOpKind(op);
|
||||
if (failed(opKind)) {
|
||||
InFlightDiagnostic diagnostic = op.emitError() << "unsupported codegen for op '" << op.getName() << "'";
|
||||
if (auto coreOp = op.getParentOfType<pim::PimCoreOp>())
|
||||
diagnostic << " inside pim.core " << coreOp.getCoreId();
|
||||
else if (auto batchOp = op.getParentOfType<pim::PimCoreBatchOp>())
|
||||
diagnostic << " inside pim.core_batch with laneCount " << batchOp.getLaneCount();
|
||||
return failure();
|
||||
}
|
||||
CompiledCoreNode node;
|
||||
node.op = &op;
|
||||
node.opKind = *opKind;
|
||||
plan.push_back(std::move(node));
|
||||
}
|
||||
return success();
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
LogicalResult compileCoreProgram(Operation* coreLikeOp, CompiledCoreProgram& program) {
|
||||
Block& block = isa<pim::PimCoreOp>(coreLikeOp) ? cast<pim::PimCoreOp>(coreLikeOp).getBody().front()
|
||||
: cast<pim::PimCoreBatchOp>(coreLikeOp).getBody().front();
|
||||
return compileCoreEmissionPlan(block, program.nodes);
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,60 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/IR/Operation.h"
|
||||
#include "mlir/Support/LogicalResult.h"
|
||||
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
|
||||
#include <memory>
|
||||
#include "src/Accelerators/PIM/Common/IR/AddressAnalysis.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
enum class CompiledCoreOpKind : uint8_t {
|
||||
Load,
|
||||
Store,
|
||||
Lmv,
|
||||
VMV,
|
||||
Receive,
|
||||
Send,
|
||||
Wait,
|
||||
Sync,
|
||||
Concat,
|
||||
Vmm,
|
||||
VVAdd,
|
||||
VVSub,
|
||||
VVMul,
|
||||
VVMax,
|
||||
VVDMul,
|
||||
VAvg,
|
||||
VRelu,
|
||||
VTanh,
|
||||
VSigm,
|
||||
VSoftmax
|
||||
};
|
||||
|
||||
struct CompiledCoreNode {
|
||||
enum class Kind : uint8_t { Op, Loop, If, IndexSwitch };
|
||||
|
||||
Kind kind = Kind::Op;
|
||||
mlir::Operation* op = nullptr;
|
||||
CompiledCoreOpKind opKind = CompiledCoreOpKind::Load;
|
||||
CompiledIndexExpr lowerBound;
|
||||
CompiledIndexExpr upperBound;
|
||||
CompiledIndexExpr step;
|
||||
CompiledIndexExpr condition;
|
||||
std::unique_ptr<llvm::SmallVector<CompiledCoreNode, 8>> loopBody;
|
||||
std::unique_ptr<llvm::SmallVector<CompiledCoreNode, 8>> thenBody;
|
||||
std::unique_ptr<llvm::SmallVector<CompiledCoreNode, 8>> elseBody;
|
||||
llvm::SmallVector<int64_t> caseValues;
|
||||
llvm::SmallVector<std::unique_ptr<llvm::SmallVector<CompiledCoreNode, 8>>> caseBodies;
|
||||
std::unique_ptr<llvm::SmallVector<CompiledCoreNode, 8>> defaultBody;
|
||||
};
|
||||
|
||||
struct CompiledCoreProgram {
|
||||
llvm::SmallVector<CompiledCoreNode, 32> nodes;
|
||||
};
|
||||
|
||||
mlir::LogicalResult compileCoreProgram(mlir::Operation* coreLikeOp, CompiledCoreProgram& program);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,95 @@
|
||||
#include "mlir/Dialect/MemRef/IR/MemRef.h"
|
||||
#include "mlir/IR/BuiltinAttributes.h"
|
||||
#include "mlir/IR/BuiltinTypes.h"
|
||||
|
||||
#include "llvm/Support/FileSystem.h"
|
||||
#include "llvm/Support/raw_ostream.h"
|
||||
|
||||
#include <cassert>
|
||||
|
||||
#include "Common/Support/CheckedArithmetic.hpp"
|
||||
#include "Conversion/ONNXToSpatial/Common/Common.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimCompilerOptions.hpp"
|
||||
#include "src/Accelerators/PIM/Compiler/PimWeightEmitter.hpp"
|
||||
|
||||
using namespace llvm;
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
namespace {} // namespace
|
||||
|
||||
WeightEmissionResult createAndPopulateWeightFolder(ArrayRef<WeightFileRequest> requests, StringRef outputDirPath) {
|
||||
auto coreWeightsDirPath = outputDirPath + "/weights";
|
||||
auto error = sys::fs::create_directory(coreWeightsDirPath);
|
||||
assert(!error && "Error creating weights directory");
|
||||
size_t indexFileName = 0;
|
||||
|
||||
int64_t xbarSize = crossbarSize.getValue();
|
||||
WeightEmissionResult result;
|
||||
llvm::SmallVector<std::pair<ResolvedWeightView, std::string>, 16> materializedWeights;
|
||||
|
||||
auto materializeWeight = [&](const ResolvedWeightView& weightView) -> std::string {
|
||||
if (auto it = llvm::find_if(materializedWeights, [&](const auto& entry) { return entry.first == weightView; });
|
||||
it != materializedWeights.end())
|
||||
return it->second;
|
||||
|
||||
auto globalOp = weightView.globalOp;
|
||||
auto denseAttr = mlir::dyn_cast<DenseElementsAttr>(*globalOp.getInitialValue());
|
||||
assert(denseAttr && "Weight global must have dense initial value");
|
||||
|
||||
ArrayRef<int64_t> shape = weightView.shape;
|
||||
assert(isMatrixShape(shape) && "Weight matrix must be 2-dimensional");
|
||||
int64_t numRows = shape[0];
|
||||
int64_t numCols = shape[1];
|
||||
assert(numRows <= xbarSize && numCols % xbarSize == 0
|
||||
&& numCols / xbarSize <= static_cast<int64_t>(crossbarCountInCore)
|
||||
&& "Weight dimensions must fit in one array group");
|
||||
|
||||
size_t elementByteWidth = getElementTypeSizeInBytes(denseAttr.getElementType());
|
||||
|
||||
std::string newFileName = "crossbar_" + std::to_string(indexFileName++) + ".bin";
|
||||
auto weightFilePath = (coreWeightsDirPath + "/" + newFileName).str();
|
||||
std::error_code errorCode;
|
||||
raw_fd_ostream weightFileStream(weightFilePath, errorCode, sys::fs::OF_None);
|
||||
if (errorCode) {
|
||||
errs() << "Error while opening weight file `" << weightFilePath << "`: " << errorCode.message() << '\n';
|
||||
assert(errorCode);
|
||||
}
|
||||
|
||||
uint64_t zero = 0;
|
||||
for (int64_t row = 0; row < xbarSize; row++) {
|
||||
for (int64_t col = 0; col < numCols; col++) {
|
||||
if (row < numRows && col < numCols) {
|
||||
int64_t elementIndex = weightView.offset + row * weightView.strides[0] + col * weightView.strides[1];
|
||||
APInt bits = denseAttr.getValues<APFloat>()[elementIndex].bitcastToAPInt();
|
||||
uint64_t word = bits.getZExtValue();
|
||||
weightFileStream.write(reinterpret_cast<const char*>(&word), elementByteWidth);
|
||||
}
|
||||
else {
|
||||
weightFileStream.write(reinterpret_cast<const char*>(&zero), elementByteWidth);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
weightFileStream.close();
|
||||
materializedWeights.push_back({weightView, newFileName});
|
||||
uint64_t weightBytes = pim::checkedMulOrCrash(
|
||||
pim::checkedMulOrCrash(static_cast<size_t>(xbarSize), static_cast<size_t>(numCols), "weight element count"),
|
||||
elementByteWidth,
|
||||
"weight byte size");
|
||||
result.totalWeightBytes = pim::checkedAddOrCrash(result.totalWeightBytes, weightBytes, "total weight bytes");
|
||||
return newFileName;
|
||||
};
|
||||
|
||||
for (const WeightFileRequest& request : requests) {
|
||||
auto& coreFiles = result.mapCoreWeightToFileName[request.coreId];
|
||||
coreFiles.reserve(request.weights.size());
|
||||
for (const ResolvedWeightView& weight : request.weights)
|
||||
coreFiles.push_back(materializeWeight(weight));
|
||||
}
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,29 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/Dialect/Func/IR/FuncOps.h"
|
||||
|
||||
#include "llvm/ADT/DenseMap.h"
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
#include "llvm/ADT/StringRef.h"
|
||||
|
||||
#include <cstdint>
|
||||
#include <string>
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/WeightUtils.hpp"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
struct WeightFileRequest {
|
||||
size_t coreId = 0;
|
||||
llvm::SmallVector<ResolvedWeightView, 8> weights;
|
||||
};
|
||||
|
||||
struct WeightEmissionResult {
|
||||
llvm::DenseMap<size_t, llvm::SmallVector<std::string, 8>> mapCoreWeightToFileName;
|
||||
uint64_t totalWeightBytes = 0;
|
||||
};
|
||||
|
||||
WeightEmissionResult createAndPopulateWeightFolder(llvm::ArrayRef<WeightFileRequest> requests,
|
||||
llvm::StringRef outputDirPath);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -1,3 +1,2 @@
|
||||
add_subdirectory(ONNXToSpatial)
|
||||
add_subdirectory(SpatialToGraphviz)
|
||||
add_subdirectory(SpatialToPim)
|
||||
add_subdirectory(SpatialToPim)
|
||||
|
||||
@@ -3,7 +3,13 @@ mlir_tablegen(ONNXToSpatial.hpp.inc -gen-rewriters "-I${ONNX_MLIR_SRC_ROOT}")
|
||||
add_public_tablegen_target(ONNXToSpatialIncGen)
|
||||
|
||||
add_pim_library(OMONNXToSpatial
|
||||
Patterns.cpp
|
||||
CompileTime.cpp
|
||||
Passes/Analyses/ONNXToSpatialVerifier.cpp
|
||||
Patterns/Pre.cpp
|
||||
Patterns/Post.cpp
|
||||
Patterns/Math/Conv.cpp
|
||||
Patterns/Math/ConvGeometry.cpp
|
||||
Patterns/Math/Elementwise.cpp
|
||||
Patterns/Math/Gemm.cpp
|
||||
Patterns/Math/MatMul.cpp
|
||||
@@ -13,12 +19,24 @@ add_pim_library(OMONNXToSpatial
|
||||
Patterns/NN/Sigmoid.cpp
|
||||
Patterns/NN/Softmax.cpp
|
||||
Patterns/Tensor/Concat.cpp
|
||||
Patterns/Tensor/Flatten.cpp
|
||||
Patterns/Tensor/Gather.cpp
|
||||
Patterns/Tensor/Resize.cpp
|
||||
Patterns/Tensor/Reshape.cpp
|
||||
Patterns/Tensor/Slice.cpp
|
||||
Patterns/Tensor/Split.cpp
|
||||
ONNXToSpatialPass.cpp
|
||||
Patterns/Tensor/Transpose.cpp
|
||||
Passes/Transforms/ONNXToSpatialPass.cpp
|
||||
Passes/Analyses/SpatialLayoutCapabilities.cpp
|
||||
Passes/Transforms/SpatialLayoutPlanningPass.cpp
|
||||
Passes/Transforms/SpatialPlanLoweringPatterns.cpp
|
||||
Passes/Transforms/LowerSpatialPlansPass.cpp
|
||||
Common/AttributeUtils.cpp
|
||||
Common/BiasAddUtils.cpp
|
||||
Common/ComputeRegionBuilder.cpp
|
||||
Common/ContractionPlanning.cpp
|
||||
Common/MatrixProductLowering.cpp
|
||||
Common/RowStripLayoutUtils.cpp
|
||||
Common/ShapeTilingUtils.cpp
|
||||
Common/WeightMaterialization.cpp
|
||||
|
||||
@@ -28,10 +46,9 @@ add_pim_library(OMONNXToSpatial
|
||||
ONNXToSpatialIncGen
|
||||
|
||||
LINK_LIBS PUBLIC
|
||||
MLIRLinalgDialect
|
||||
MLIRSCFDialect
|
||||
MLIRTosaDialect
|
||||
OMCompilerOptions
|
||||
OMPimCompilerOptions
|
||||
OMONNXOps
|
||||
SpatialOps
|
||||
OMPimCommon
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
#include "mlir/IR/BuiltinAttributes.h"
|
||||
|
||||
#include "AttributeUtils.hpp"
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
int64_t getI64Attr(ArrayAttr attr, size_t index) { return cast<IntegerAttr>(attr[index]).getInt(); }
|
||||
|
||||
int64_t getOptionalI64Attr(std::optional<ArrayAttr> attr, size_t index, int64_t defaultValue) {
|
||||
return attr ? getI64Attr(*attr, index) : defaultValue;
|
||||
}
|
||||
|
||||
llvm::SmallVector<int64_t> getI64ArrayAttrValues(ArrayAttr attr) {
|
||||
llvm::SmallVector<int64_t> values;
|
||||
values.reserve(attr.size());
|
||||
for (Attribute value : attr)
|
||||
values.push_back(cast<IntegerAttr>(value).getInt());
|
||||
return values;
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,18 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/IR/BuiltinAttributes.h"
|
||||
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
|
||||
#include <cstddef>
|
||||
#include <optional>
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
int64_t getI64Attr(mlir::ArrayAttr attr, size_t index);
|
||||
|
||||
int64_t getOptionalI64Attr(std::optional<mlir::ArrayAttr> attr, size_t index, int64_t defaultValue);
|
||||
|
||||
llvm::SmallVector<int64_t> getI64ArrayAttrValues(mlir::ArrayAttr attr);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,112 @@
|
||||
#include "mlir/IR/BuiltinAttributes.h"
|
||||
#include "mlir/IR/BuiltinTypes.h"
|
||||
|
||||
#include "llvm/ADT/SmallVector.h"
|
||||
|
||||
#include "src/Accelerators/PIM/Common/IR/ConstantUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Common/IR/ShapeUtils.hpp"
|
||||
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/CompileTime.hpp"
|
||||
#include "src/Accelerators/PIM/Conversion/ONNXToSpatial/Common/BiasAddUtils.hpp"
|
||||
|
||||
using namespace mlir;
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
LogicalResult isSupportedBiasAddShape(RankedTensorType biasType, RankedTensorType resultType) {
|
||||
if (!biasType || !resultType || !biasType.hasStaticShape() || !resultType.hasStaticShape())
|
||||
return failure();
|
||||
if (resultType.getRank() != 4)
|
||||
return failure();
|
||||
if (biasType.getElementType() != resultType.getElementType())
|
||||
return failure();
|
||||
|
||||
const int64_t channels = resultType.getDimSize(1);
|
||||
ArrayRef<int64_t> shape = biasType.getShape();
|
||||
if (shape.empty())
|
||||
return success();
|
||||
if (shape.size() == 1)
|
||||
return success(shape[0] == channels);
|
||||
if (shape.size() == 2)
|
||||
return success(shape[0] == 1 && shape[1] == channels);
|
||||
if (shape.size() == 4)
|
||||
return success(shape[0] == 1 && shape[1] == channels && shape[2] == 1 && shape[3] == 1);
|
||||
return failure();
|
||||
}
|
||||
|
||||
FailureOr<SmallVector<Attribute>> getBiasChannelValues(DenseElementsAttr denseAttr, RankedTensorType resultType) {
|
||||
auto biasType = dyn_cast<RankedTensorType>(denseAttr.getType());
|
||||
if (!biasType || failed(isSupportedBiasAddShape(biasType, resultType)))
|
||||
return failure();
|
||||
|
||||
const int64_t channels = resultType.getDimSize(1);
|
||||
if (denseAttr.isSplat()) {
|
||||
return SmallVector<Attribute>(channels, denseAttr.getSplatValue<Attribute>());
|
||||
}
|
||||
|
||||
SmallVector<Attribute> flattened(denseAttr.getValues<Attribute>());
|
||||
if (biasType.getRank() == 1)
|
||||
return flattened;
|
||||
if (biasType.getRank() == 2)
|
||||
return flattened;
|
||||
|
||||
SmallVector<Attribute> channelValues;
|
||||
channelValues.reserve(channels);
|
||||
const int64_t channelStride = biasType.getDimSize(2) * biasType.getDimSize(3);
|
||||
for (int64_t channel = 0; channel < channels; ++channel)
|
||||
channelValues.push_back(flattened[channel * channelStride]);
|
||||
return channelValues;
|
||||
}
|
||||
|
||||
bool isSupportedBiasAddValue(Value bias, RankedTensorType resultType, DenseElementsAttr* denseAttr) {
|
||||
auto attr = getHostConstDenseElementsAttr(bias);
|
||||
if (!attr)
|
||||
return false;
|
||||
auto biasType = dyn_cast<RankedTensorType>(attr.getType());
|
||||
if (!biasType || failed(isSupportedBiasAddShape(biasType, resultType)))
|
||||
return false;
|
||||
if (failed(getBiasChannelValues(attr, resultType)))
|
||||
return false;
|
||||
if (denseAttr)
|
||||
*denseAttr = attr;
|
||||
return true;
|
||||
}
|
||||
|
||||
FailureOr<BiasAddPlanCandidate> classifyBiasAddPlanCandidate(Value lhs, Value rhs, RankedTensorType resultType) {
|
||||
auto lhsType = dyn_cast<RankedTensorType>(lhs.getType());
|
||||
auto rhsType = dyn_cast<RankedTensorType>(rhs.getType());
|
||||
if (!lhsType || !rhsType)
|
||||
return failure();
|
||||
if (lhsType == resultType && isSupportedBiasAddValue(rhs, resultType))
|
||||
return BiasAddPlanCandidate {lhs, rhs};
|
||||
if (rhsType == resultType && isSupportedBiasAddValue(lhs, resultType))
|
||||
return BiasAddPlanCandidate {rhs, lhs};
|
||||
return failure();
|
||||
}
|
||||
|
||||
FailureOr<Value>
|
||||
materializeDenseBiasAddTensor(Value bias, RankedTensorType resultType, RewriterBase& rewriter, Location loc) {
|
||||
DenseElementsAttr denseAttr;
|
||||
if (!isSupportedBiasAddValue(bias, resultType, &denseAttr))
|
||||
return failure();
|
||||
|
||||
FailureOr<SmallVector<Attribute>> channelValues = getBiasChannelValues(denseAttr, resultType);
|
||||
if (failed(channelValues))
|
||||
return failure();
|
||||
|
||||
SmallVector<Attribute> resultValues;
|
||||
resultValues.reserve(resultType.getNumElements());
|
||||
const int64_t batches = resultType.getDimSize(0);
|
||||
const int64_t channels = resultType.getDimSize(1);
|
||||
const int64_t height = resultType.getDimSize(2);
|
||||
const int64_t width = resultType.getDimSize(3);
|
||||
for (int64_t n = 0; n < batches; ++n)
|
||||
for (int64_t c = 0; c < channels; ++c)
|
||||
for (int64_t h = 0; h < height; ++h)
|
||||
for (int64_t w = 0; w < width; ++w)
|
||||
resultValues.push_back((*channelValues)[c]);
|
||||
|
||||
auto resultAttr = DenseElementsAttr::get(resultType, resultValues);
|
||||
return getOrCreateConstant(rewriter, rewriter.getInsertionBlock()->getParentOp(), resultAttr, resultType);
|
||||
}
|
||||
|
||||
} // namespace onnx_mlir
|
||||
@@ -0,0 +1,30 @@
|
||||
#pragma once
|
||||
|
||||
#include "mlir/IR/BuiltinAttributes.h"
|
||||
#include "mlir/IR/BuiltinTypes.h"
|
||||
#include "mlir/IR/PatternMatch.h"
|
||||
#include "mlir/IR/Value.h"
|
||||
#include "mlir/Support/LogicalResult.h"
|
||||
|
||||
namespace onnx_mlir {
|
||||
|
||||
struct BiasAddPlanCandidate {
|
||||
mlir::Value data;
|
||||
mlir::Value bias;
|
||||
};
|
||||
|
||||
mlir::LogicalResult isSupportedBiasAddShape(mlir::RankedTensorType biasType, mlir::RankedTensorType resultType);
|
||||
bool isSupportedBiasAddValue(mlir::Value bias,
|
||||
mlir::RankedTensorType resultType,
|
||||
mlir::DenseElementsAttr* denseAttr = nullptr);
|
||||
mlir::FailureOr<llvm::SmallVector<mlir::Attribute>>
|
||||
getBiasChannelValues(mlir::DenseElementsAttr denseAttr, mlir::RankedTensorType resultType);
|
||||
mlir::FailureOr<BiasAddPlanCandidate> classifyBiasAddPlanCandidate(mlir::Value lhs,
|
||||
mlir::Value rhs,
|
||||
mlir::RankedTensorType resultType);
|
||||
mlir::FailureOr<mlir::Value> materializeDenseBiasAddTensor(mlir::Value bias,
|
||||
mlir::RankedTensorType resultType,
|
||||
mlir::RewriterBase& rewriter,
|
||||
mlir::Location loc);
|
||||
|
||||
} // namespace onnx_mlir
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user