@@ -5,23 +5,32 @@ and pooling methods for hexagonally sampled data, originally written for
55PyTorch by Tim Lukas Holch and Constantin Steppa (ai4iacts).
66
77This 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
1011is 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
1213layout 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
2635See [ 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
5665call, 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
99183Either way, ` pytest tests/ ` works without installing anything -- ` conftest.py `
100184puts ` 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+
102202Verified to pass on all three Keras 3 backends (set ` KERAS_BACKEND=tensorflow|torch|jax `
103203before 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
111211A GitHub Actions workflow ([ .github/workflows/test.yml] ( .github/workflows/test.yml ) )
112212runs 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
0 commit comments