Skip to content

Commit 838e68d

Browse files
committed
addign support for hls4ml (only 2d and only io_stream), maybe need further test but enough for basic support. updating readme, correcting jax and torch backend computation imprecision
1 parent 9ac4411 commit 838e68d

11 files changed

Lines changed: 251 additions & 82 deletions

File tree

‎README.md‎

Lines changed: 122 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -5,23 +5,32 @@ and pooling methods for hexagonally sampled data, originally written for
55
PyTorch by Tim Lukas Holch and Constantin Steppa (ai4iacts).
66

77
This port reproduces HexagDLy's hexagonal addressing scheme and sub-kernel
8-
decomposition exactly (bit-for-bit equivalent outputs, see
9-
[tests/test_vs_pytorch_hexagdly.py](tests/test_vs_pytorch_hexagdly.py)), but
8+
decomposition exactly (bit-for-bit equivalent outputs, cross-checked against the
9+
upstream PyTorch package in
10+
[hexagdly-oracle](https://github.com/YugnatD/hexagdly-oracle)), but
1011
is built on [Keras 3](https://keras.io) so it runs on any backend
1112
(TensorFlow, JAX, PyTorch) and uses a channels-last (`NHWC`/`NDHWC`) tensor
1213
layout instead of PyTorch's channels-first.
1314

14-
It also adds two functionalities that do not exist in upstream HexagDLy:
15+
It also adds three functionalities that do not exist in upstream HexagDLy:
1516

16-
- **`share_neighbors`** (`Conv2d`, `Conv3d`): ties the weights of a hexagonal
17-
kernel by ring (ring 0 = center cell, ring *r* = the `6*r` cells at hex
18-
distance *r*), instead of every cell having its own independent weight.
19-
Reduces parameter count and enforces exact 6-fold rotational symmetry of
20-
the learned kernel.
17+
- **`share_neighbors`** (`Conv2d`, `Conv3d`): ties the weights of several cells
18+
of a hexagonal kernel together instead of giving every cell its own
19+
independent weight, which cuts the parameter count and imposes a geometric
20+
symmetry on the learned kernel. Three modes -- `"ring"`, `"diag"`, `"sym"` --
21+
are [illustrated below](#new-share_neighbors----weight-sharing-across-kernel-cells).
2122
- **`depth_padding="same"`** (`Conv3d` only): zero-pads the depth/time axis
2223
so the temporal kernel is centred on each time step and the output depth
2324
equals the input depth, instead of HexagDLy's `"valid"`-only behaviour
2425
(output depth shrinks by `kernel - 1`).
26+
- **[hls4ml export](#fpga-export-via-hls4ml)** for FPGA synthesis, currently
27+
limited to `io_stream` and the 2D layers.
28+
29+
Weights trained with [pytorch-hexagdly](https://github.com/YugnatD/pytorch-hexagdly)
30+
can be loaded into these layers -- see
31+
[`keras_hexagdly.torch_interop`](src/keras_hexagdly/torch_interop.py) and
32+
[`notebooks/pytorch_to_keras_example.ipynb`](notebooks/pytorch_to_keras_example.ipynb).
33+
No torch dependency is required: a plain `.npz` of the state dict works.
2534

2635
See [NOTICE.md](NOTICE.md) for attribution details and citation information.
2736

@@ -55,12 +64,40 @@ y = hexconv(x)
5564
`in_channels` can be omitted; it is then inferred from the input on first
5665
call, like a standard Keras layer: `hgly.Conv2d(out_channels, kernel_size=kernel_size, stride=stride)`.
5766

58-
### New: weight sharing by hexagonal ring
67+
### New: `share_neighbors` -- weight sharing across kernel cells
68+
69+
Available on `Conv2d` and `Conv3d`. `share_neighbors` reduces the number of
70+
learnable parameters by grouping kernel cells that share a single weight.
71+
Three modes are available, illustrated below for `kernel_size=2` (19 cells):
72+
73+
| `share_neighbors="ring"` | `share_neighbors="diag"` | `share_neighbors="sym"` |
74+
|:---:|:---:|:---:|
75+
| ![ring](figures/share_ring_k2.png) | ![diag](figures/share_diag_k2.png) | ![sym](figures/share_sym_k2.png) |
76+
| **3 weights** -- cells at the same hex distance from center share one weight (concentric rings). | **10 weights** -- visually opposite (antipodal) cells share one weight. | **10 weights** -- geometrically adjacent 60 degree pairs share one weight. |
77+
78+
- **`"ring"`**: the most aggressive reduction. All 6 direct neighbours share one
79+
weight, all 12 outer cells share another. Enforces exact 6-fold rotational
80+
symmetry of the learned kernel.
81+
- **`"diag"`**: antipodal symmetry -- each cell and its mirror image through the
82+
center share a weight. Useful when the kernel should be point-symmetric.
83+
- **`"sym"`**: 60 degree adjacent pairs -- consecutive neighbours along the kernel
84+
boundary share a weight. Useful when the kernel should reflect local
85+
rotational symmetry.
86+
87+
For `kernel_size=1` (7 cells): ring=2 weights, diag=4, sym=4.
88+
For `kernel_size=2` (19 cells): ring=3 weights, diag=10, sym=10.
89+
Compare to the default `share_neighbors=None`, which gives 7 and 19 independent
90+
weights. `"ring"` works at any kernel size; `"diag"` and `"sym"` are defined for
91+
`kernel_size` 1 and 2 only, matching
92+
[pytorch-hexagdly](https://github.com/YugnatD/pytorch-hexagdly), which does not
93+
define them beyond n=2 either.
5994

6095
```python
61-
hexconv = hgly.Conv2d(in_channels, out_channels, kernel_size=3, share_neighbors=True)
96+
hexconv = hgly.Conv2d(in_channels, out_channels, kernel_size=2, share_neighbors="ring")
6297
```
6398

99+
`share_neighbors=True` is accepted as an alias for `"ring"`.
100+
64101
### New: same-padded temporal convolution (Conv3d)
65102

66103
```python
@@ -85,6 +122,53 @@ Ported from [HexagDLy's own notebooks](https://github.com/ai4iacts/hexagdly/tree
85122
- [`keras_hexagdly_cnn_example.ipynb`](notebooks/keras_hexagdly_cnn_example.ipynb) -- a small CNN classifying toy hexagonal shapes, trained with `model.fit`.
86123
- [`keras_hexagdly_hex_vs_square.ipynb`](notebooks/keras_hexagdly_hex_vs_square.ipynb) -- parameter-count and timing benchmark of hex vs. square kernels, plus a hex-CNN-vs-square-CNN classification comparison.
87124

125+
## FPGA export via hls4ml
126+
127+
The hex layers can be synthesised to HLS C++ through
128+
[hls4ml](https://github.com/fastmachinelearning/hls4ml). The layers are replaced
129+
by a fused line-buffer kernel that keeps only the resident rows of the frame in
130+
a shift register, rather than materialising the whole gathered tensor.
131+
132+
```python
133+
import hls4ml
134+
from keras_hexagdly.hls4ml_handler import register_hex_gather_layers
135+
from keras_hexagdly.hls4ml_ext import patch_model_for_hls, hex_reuse_config, check_hls_config
136+
137+
register_hex_gather_layers("Vitis")
138+
hls_ready = patch_model_for_hls(model) # strategy="linebuffer"
139+
140+
config = hls4ml.utils.config_from_keras_model(hls_ready, granularity="name")
141+
hex_reuse_config(config, hls_ready) # per-layer ReuseFactor
142+
check_hls_config(config, hls_ready, io_type="io_stream")
143+
144+
hls_model = hls4ml.converters.convert_from_keras_model(
145+
hls_ready, hls_config=config, io_type="io_stream", backend="Vitis",
146+
)
147+
```
148+
149+
### Supported scope
150+
151+
**Only `io_stream` and the 2D layers (`Conv2d`, `MaxPool2d`) are supported for
152+
now.** That is the combination that is covered by C-simulation and validated by
153+
RTL co-simulation.
154+
155+
| | `io_stream` | `io_parallel` |
156+
|---|---|---|
157+
| `Conv2d`, `MaxPool2d` | **supported** | not supported |
158+
| `Conv3d`, `MaxPool3d` | not supported | not supported |
159+
160+
The unsupported paths are not silently wrong -- they raise. `patch_model_for_hls`
161+
raises `NotImplementedError` on a 3D layer, and `check_hls_config` raises on
162+
`io_type="io_parallel"`. Both accept `allow_unvalidated=True` if you want to
163+
experiment with them anyway, but nothing about their numerics or resource usage
164+
is guaranteed.
165+
166+
`hex_reuse_config` matters more than it looks: hls4ml's
167+
`config_from_keras_model(granularity="name")` writes `ReuseFactor=1` into every
168+
layer entry it recognises, which overrides the model-level value -- but custom
169+
layers get no entry and inherit the model-level one instead. Without an explicit
170+
per-layer setting, a single model ends up mixing two different reuse factors.
171+
88172
## Testing
89173

90174
```
@@ -99,42 +183,42 @@ that mis-names the built wheel `UNKNOWN`. Verified clean with a modern pip
99183
Either way, `pytest tests/` works without installing anything -- `conftest.py`
100184
puts `src/` and `tests/` on `sys.path`.)
101185

186+
Most of the suite no longer lives here. It has moved to
187+
**[hexagdly-oracle](https://github.com/YugnatD/hexagdly-oracle)**, the shared
188+
test repo for this library and
189+
[pytorch-hexagdly](https://github.com/YugnatD/pytorch-hexagdly): hand-verified
190+
layer outputs, `share_neighbors` weight-sharing oracles, `depth_padding`,
191+
mixed precision, serialization, edge cases, indexed-equivalence and the hls4ml
192+
export tests (including C-simulation). `tests/` here keeps only what is
193+
genuinely local.
194+
195+
To run the full suite, check the oracle out as a sibling directory:
196+
197+
```
198+
git clone https://github.com/YugnatD/hexagdly-oracle
199+
PYTHONPATH=hexagdly-oracle/src pytest tests/ hexagdly-oracle/tests/
200+
```
201+
102202
Verified to pass on all three Keras 3 backends (set `KERAS_BACKEND=tensorflow|torch|jax`
103203
before importing keras; tensorflow is the default if unset):
104204

105205
```
106-
KERAS_BACKEND=tensorflow pytest tests/ # 272 passed, 7 skipped
107-
KERAS_BACKEND=torch pytest tests/ # 269 passed, 10 skipped
108-
KERAS_BACKEND=jax pytest tests/ # 269 passed, 10 skipped (slower: per-shape JIT compile)
206+
KERAS_BACKEND=tensorflow # 891 passed, 7 skipped
207+
KERAS_BACKEND=torch # 881 passed, 17 skipped
208+
KERAS_BACKEND=jax # 806 passed, 92 skipped (slower: per-shape JIT compile)
109209
```
110210

111211
A GitHub Actions workflow ([.github/workflows/test.yml](.github/workflows/test.yml))
112212
runs this matrix (3 backends x 3 Python versions) plus `ruff check`/`ruff format --check`
113-
on every push and PR.
114-
115-
The test suite has six parts:
116-
117-
- `test_Conv2d.py`, `test_Conv3d.py`, `test_*_CustomKernel.py`,
118-
`test_MaxPool2d.py`, `test_MaxPool3d.py`: standalone tests with
119-
hand-computed expected outputs, ported from
120-
[HexagDLy's own test suite](https://github.com/ai4iacts/hexagdly/tree/master/tests).
121-
- `test_vs_pytorch_hexagdly.py`: cross-checks every layer against the
122-
upstream PyTorch `hexagdly` PyPI package (random inputs, weights copied
123-
across frameworks, gradients, batch independence, odd/even column parity,
124-
large strides, asymmetric 3D kernels) -- the oracle that proves this port
125-
faithful. Also includes two *forward-compatibility* tests that stay
126-
skipped today (upstream hexagdly 2.0.2 has neither feature) but will
127-
automatically start cross-checking `share_neighbors`/`depth_padding`
128-
against upstream the day a future hexagdly release adds them.
129-
- `test_mixed_precision.py`: `keras.mixed_precision` / per-layer `dtype=`
130-
policies -- variables stay float32, compute happens in float16, on every
131-
backend.
132-
- `test_share_neighbors.py`, `test_depth_padding.py`: standalone tests for
133-
the two new functionalities, which have no PyTorch equivalent to check
134-
against.
135-
- `test_edge_cases.py`, `test_serialization.py`, `test_geometry.py`: input
136-
validation, minimum viable sizes, dtype handling, NaN/Inf checks, and
137-
`get_config`/`from_config` + full `model.save`/`load_model` round-trips.
213+
on every push and PR, checking out the oracle repo as part of the job.
214+
215+
Note for GPU users on the torch backend: the Keras torch backend runs on CUDA
216+
when a GPU is visible, and PyTorch defaults to `cudnn.allow_tf32 = True`, so
217+
convolutions are computed in TF32 (~1e-3 relative precision). That is enough to
218+
break equivalence assertions on hex kernels, which are wide by construction
219+
because dilation is done by zero insertion. The oracle's `conftest.py` pins the
220+
flag off for the test session; the library itself never touches global torch
221+
settings.
138222

139223
## Disclaimer
140224

‎figures/share_diag_k2.png‎

135 KB
Loading

‎figures/share_ring_k2.png‎

125 KB
Loading

‎figures/share_sym_k2.png‎

134 KB
Loading

‎notebooks/addressing_utils.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,10 +5,10 @@
55
of which deep learning framework consumes the result.
66
"""
77

8-
import numpy as np
98
import matplotlib.colors as mcolors
10-
from matplotlib.patches import RegularPolygon
9+
import numpy as np
1110
from matplotlib.collections import PatchCollection
11+
from matplotlib.patches import RegularPolygon
1212

1313

1414
class Detector:

‎notebooks/hexplot.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,12 +5,12 @@
55
(addressing scheme), just reading (N, H, W, C) tensors instead of (N, C, H, W).
66
"""
77

8-
import numpy as np
98
import keras
109
import matplotlib.pyplot as plt
11-
from matplotlib.patches import RegularPolygon
12-
from matplotlib.collections import PatchCollection
10+
import numpy as np
1311
from matplotlib import gridspec
12+
from matplotlib.collections import PatchCollection
13+
from matplotlib.patches import RegularPolygon
1414

1515

1616
def plot_hextensor(tensor, image_range=(0, None), channel_range=(0, None),

‎notebooks/toy_data.py‎

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,6 @@
1717

1818

1919
def put_shape(nx, ny, cx, cy, params):
20-
d = np.zeros((nx, ny))
2120
i = np.indices((nx, ny)).astype(float)
2221
i[0] -= cx
2322
i[1] -= cy

‎src/keras_hexagdly/hex_gather.py‎

Lines changed: 30 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
import keras
2222
import numpy as np
2323

24+
2425
@keras.saving.register_keras_serializable(package="hexagdly_tf")
2526
class HexConvLineBuffer(keras.layers.Layer):
2627
"""Fused hex neighbor gather + ring-MAC on the 2D grid.
@@ -107,8 +108,8 @@ def call(self, x):
107108
Cin = x.shape[-1]
108109
x_flat = keras.ops.reshape(x, (-1, self.H * self.W, Cin)) # (B, N_in, Cin)
109110

110-
valid = self.neighbor_idx >= 0 # (N_out, K)
111-
safe_idx = keras.ops.where(valid, self.neighbor_idx, 0)
111+
valid = keras.ops.greater_equal(self.neighbor_idx, 0) # (N_out, K)
112+
safe_idx = keras.ops.maximum(self.neighbor_idx, 0)
112113
gathered = keras.ops.take(x_flat, safe_idx, axis=1) # (B, N_out, K, Cin)
113114
mask = keras.ops.cast(keras.ops.reshape(valid, (1, self.N_out, self.K, 1)), gathered.dtype)
114115
gathered = gathered * mask
@@ -250,8 +251,15 @@ def call(self, x):
250251
D_padded = D_in + pad_top + pad_bot
251252
D_eff = D_padded - self.depth_size + 1
252253

253-
valid = self.neighbor_idx >= 0
254-
safe_idx = keras.ops.where(valid, self.neighbor_idx, 0)
254+
# maximum(), not where(): the border sentinel is -1, so clamping at 0
255+
# picks the same placeholder slot the mask below zeroes out anyway. The
256+
# where() form is what it replaced, and on the jax backend that one
257+
# raises during predict() -- ops.where over a Variable index forces a
258+
# read outside the trainer's stateless scope, where the variable has
259+
# been purged (TypeError: 'NoneType' object is not callable). Direct
260+
# __call__ and the TensorFlow backend never hit it.
261+
valid = keras.ops.greater_equal(self.neighbor_idx, 0)
262+
safe_idx = keras.ops.maximum(self.neighbor_idx, 0)
255263
G = keras.ops.take(x_flat, safe_idx, axis=2) # (B, D_padded, N_out, K, Cin)
256264
mask = keras.ops.cast(keras.ops.reshape(valid, (1, 1, self.N_out, self.K, 1)), G.dtype)
257265
G = G * mask
@@ -345,8 +353,15 @@ def call(self, x):
345353
# x: (B, H, W, C)
346354
C = x.shape[-1]
347355
x_flat = keras.ops.reshape(x, (-1, self.H * self.W, C))
348-
valid = self.neighbor_idx >= 0
349-
safe_idx = keras.ops.where(valid, self.neighbor_idx, 0)
356+
# maximum(), not where(): the border sentinel is -1, so clamping at 0
357+
# picks the same placeholder slot the mask below zeroes out anyway. The
358+
# where() form is what it replaced, and on the jax backend that one
359+
# raises during predict() -- ops.where over a Variable index forces a
360+
# read outside the trainer's stateless scope, where the variable has
361+
# been purged (TypeError: 'NoneType' object is not callable). Direct
362+
# __call__ and the TensorFlow backend never hit it.
363+
valid = keras.ops.greater_equal(self.neighbor_idx, 0)
364+
safe_idx = keras.ops.maximum(self.neighbor_idx, 0)
350365
gathered = keras.ops.take(x_flat, safe_idx, axis=1) # (B, N_out, K, C)
351366
mask = keras.ops.cast(keras.ops.reshape(valid, (1, self.N_out, self.K, 1)), gathered.dtype)
352367
gathered = gathered * mask
@@ -414,8 +429,15 @@ def call(self, x):
414429
C = x.shape[-1]
415430
d_in = x.shape[1]
416431
x_flat = keras.ops.reshape(x, (-1, d_in, self.H * self.W, C))
417-
valid = self.neighbor_idx >= 0
418-
safe_idx = keras.ops.where(valid, self.neighbor_idx, 0)
432+
# maximum(), not where(): the border sentinel is -1, so clamping at 0
433+
# picks the same placeholder slot the mask below zeroes out anyway. The
434+
# where() form is what it replaced, and on the jax backend that one
435+
# raises during predict() -- ops.where over a Variable index forces a
436+
# read outside the trainer's stateless scope, where the variable has
437+
# been purged (TypeError: 'NoneType' object is not callable). Direct
438+
# __call__ and the TensorFlow backend never hit it.
439+
valid = keras.ops.greater_equal(self.neighbor_idx, 0)
440+
safe_idx = keras.ops.maximum(self.neighbor_idx, 0)
419441
G = keras.ops.take(x_flat, safe_idx, axis=2) # (B, D_in, N_out, K, C)
420442
mask = keras.ops.cast(keras.ops.reshape(valid, (1, 1, self.N_out, self.K, 1)), G.dtype)
421443
G = G * mask

0 commit comments

Comments
 (0)