pytorch-hexagdly 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,1041 @@
1
+ r"""
2
+ This file contains utilities to set up hexagonal convolution and pooling
3
+ kernels in PyTorch. The size of the input is abitrary, whereas the layout
4
+ from top to bottom (along tensor index 2) has to be of zig-zag-edge shape
5
+ and from left to right (along tensor index 3) of armchair-edge shape as
6
+ shown below.
7
+ __ __ __ __ __ __
8
+ /11\__/31\__ . . . |11|21|31|41| . . .
9
+ \__/21\__/41\ |__|__|__|__|
10
+ /12\__/32\__/ . . . _______|\ |12|22|32|42| . . .
11
+ \__/22\__/42\ | \ |__|__|__|__|
12
+ \__/ \__/ |_______ /
13
+ . . . . . |/ . . . . .
14
+ . . . . . . . . . .
15
+ . . . . . . . . . .
16
+
17
+ For more information visit https://github.com/ai4iacts/hexagdly
18
+
19
+ """
20
+
21
+ __version__ = "0.1.0"
22
+
23
+ __all__ = [
24
+ "Conv2d",
25
+ "Conv2d_CustomKernel",
26
+ "Conv3d",
27
+ "Conv3d_CustomKernel",
28
+ "MaxPool2d",
29
+ "MaxPool3d",
30
+ "ring_maps_2d",
31
+ ]
32
+
33
+ import numpy as np
34
+ import torch
35
+ import torch.nn as nn
36
+ import torch.nn.functional as F
37
+ from torch.nn.parameter import Parameter
38
+
39
+
40
+ class HexBase:
41
+ def __init__(self):
42
+ super(HexBase, self).__init__()
43
+ self.hexbase_size = None
44
+ self.depth_size = None
45
+ self.hexbase_stride = None
46
+ self.depth_stride = None
47
+ self.input_size_is_known = False
48
+ self.odd_columns_slices = []
49
+ self.odd_columns_pads = []
50
+ self.even_columns_slices = []
51
+ self.even_columns_pads = []
52
+ self.dimensions = None
53
+ self.combine = None
54
+ self.process = None
55
+ self.kwargs = dict()
56
+
57
+ def shape_for_odd_columns(self, input_size, kernel_number):
58
+ slices = [None, None, None, None]
59
+ pads = [0, 0, 0, 0]
60
+ # left
61
+ pads[0] = kernel_number
62
+ # right
63
+ pads[1] = max(0, kernel_number - ((input_size[-1] - 1) % (2 * self.hexbase_stride)))
64
+ # top
65
+ pads[2] = self.hexbase_size - int(kernel_number / 2)
66
+ # bottom
67
+ constraint = (
68
+ input_size[-2]
69
+ - 1
70
+ - int((input_size[-2] - 1 - int(self.hexbase_stride / 2)) / self.hexbase_stride)
71
+ * self.hexbase_stride
72
+ )
73
+ bottom = (self.hexbase_size - int((kernel_number + 1) / 2)) - constraint
74
+ if bottom >= 0:
75
+ pads[3] = bottom
76
+ else:
77
+ slices[1] = bottom
78
+
79
+ return slices, pads
80
+
81
+ def shape_for_even_columns(self, input_size, kernel_number):
82
+ slices = [None, None, None, None]
83
+ pads = [0, 0, 0, 0]
84
+ # left
85
+ left = kernel_number - self.hexbase_stride
86
+ if left >= 0:
87
+ pads[0] = left
88
+ else:
89
+ slices[2] = -left
90
+ # right
91
+ pads[1] = max(
92
+ 0,
93
+ kernel_number
94
+ - ((input_size[-1] - 1 - self.hexbase_stride) % (2 * self.hexbase_stride)),
95
+ )
96
+ # top
97
+ top_shift = -(kernel_number % 2) if (self.hexbase_stride % 2) == 1 else 0
98
+ top = (
99
+ (self.hexbase_size - int(kernel_number / 2)) + top_shift - int(self.hexbase_stride / 2)
100
+ )
101
+ if top >= 0:
102
+ pads[2] = top
103
+ else:
104
+ slices[0] = -top
105
+ # bottom
106
+ bottom_shift = 0 if (self.hexbase_stride % 2) == 1 else -(kernel_number % 2)
107
+ pads[3] = max(
108
+ 0,
109
+ self.hexbase_size
110
+ - int(kernel_number / 2)
111
+ + bottom_shift
112
+ - ((input_size[-2] - int(self.hexbase_stride / 2) - 1) % self.hexbase_stride),
113
+ )
114
+
115
+ return slices, pads
116
+
117
+ def get_padded_input(self, input, pads):
118
+ if self.dimensions == 2:
119
+ return nn.ZeroPad2d(tuple(pads))(input)
120
+ elif self.dimensions == 3:
121
+ return nn.ConstantPad3d(tuple(pads + [0, 0]), 0)(input)
122
+
123
+ def get_sliced_input(self, input, slices):
124
+ if self.dimensions == 2:
125
+ return input[:, :, slices[0] : slices[1], slices[2] : slices[3]]
126
+ elif self.dimensions == 3:
127
+ return input[:, :, :, slices[0] : slices[1], slices[2] : slices[3]]
128
+
129
+ def get_dilation(self, dilation_2d):
130
+ if self.dimensions == 2:
131
+ return dilation_2d
132
+ elif self.dimensions == 3:
133
+ return tuple([1] + list(dilation_2d))
134
+
135
+ def get_stride(self):
136
+ if self.dimensions == 2:
137
+ return (self.hexbase_stride, 2 * self.hexbase_stride)
138
+ elif self.dimensions == 3:
139
+ return (self.depth_stride, self.hexbase_stride, 2 * self.hexbase_stride)
140
+
141
+ def get_ordered_output(self, input, order):
142
+ if self.dimensions == 2:
143
+ return input[:, :, :, order]
144
+ elif self.dimensions == 3:
145
+ return input[:, :, :, :, order]
146
+
147
+ # general implementation of an operation with a hexagonal kernel
148
+ def operation_with_arbitrary_stride(self, input):
149
+ assert input.size(-2) - (self.hexbase_stride // 2) >= 0, (
150
+ "Too few rows to apply hex conv with the stide that is set"
151
+ )
152
+ odd_columns = None
153
+ even_columns = None
154
+
155
+ for i in range(self.hexbase_size + 1):
156
+ dilation_base = (1, 1) if i == 0 else (1, 2 * i)
157
+
158
+ if not self.input_size_is_known:
159
+ slices, pads = self.shape_for_odd_columns(input.size(), i)
160
+ self.odd_columns_slices.append(slices)
161
+ self.odd_columns_pads.append(pads)
162
+ slices, pads = self.shape_for_even_columns(input.size(), i)
163
+ self.even_columns_slices.append(slices)
164
+ self.even_columns_pads.append(pads)
165
+ if i == self.hexbase_size:
166
+ self.input_size_is_known = True
167
+
168
+ if odd_columns is None:
169
+ odd_columns = self.process(
170
+ self.get_padded_input(
171
+ self.get_sliced_input(input, self.odd_columns_slices[i]),
172
+ self.odd_columns_pads[i],
173
+ ),
174
+ getattr(self, "kernel" + str(i)),
175
+ dilation=self.get_dilation(dilation_base),
176
+ stride=self.get_stride(),
177
+ **self.kwargs,
178
+ )
179
+ else:
180
+ odd_columns = self.combine(
181
+ odd_columns,
182
+ self.process(
183
+ self.get_padded_input(
184
+ self.get_sliced_input(input, self.odd_columns_slices[i]),
185
+ self.odd_columns_pads[i],
186
+ ),
187
+ getattr(self, "kernel" + str(i)),
188
+ dilation=self.get_dilation(dilation_base),
189
+ stride=self.get_stride(),
190
+ ),
191
+ )
192
+
193
+ if even_columns is None:
194
+ even_columns = self.process(
195
+ self.get_padded_input(
196
+ self.get_sliced_input(input, self.even_columns_slices[i]),
197
+ self.even_columns_pads[i],
198
+ ),
199
+ getattr(self, "kernel" + str(i)),
200
+ dilation=self.get_dilation(dilation_base),
201
+ stride=self.get_stride(),
202
+ **self.kwargs,
203
+ )
204
+ else:
205
+ even_columns = self.combine(
206
+ even_columns,
207
+ self.process(
208
+ self.get_padded_input(
209
+ self.get_sliced_input(input, self.even_columns_slices[i]),
210
+ self.even_columns_pads[i],
211
+ ),
212
+ getattr(self, "kernel" + str(i)),
213
+ dilation=self.get_dilation(dilation_base),
214
+ stride=self.get_stride(),
215
+ ),
216
+ )
217
+
218
+ concatenated_columns = torch.cat((odd_columns, even_columns), 1 + self.dimensions)
219
+
220
+ n_odd_columns = odd_columns.size(-1)
221
+ n_even_columns = even_columns.size(-1)
222
+ if n_odd_columns == n_even_columns:
223
+ order = [int(i + x * n_even_columns) for i in range(n_even_columns) for x in range(2)]
224
+ else:
225
+ order = [int(i + x * n_odd_columns) for i in range(n_even_columns) for x in range(2)]
226
+ order.append(n_even_columns)
227
+
228
+ return self.get_ordered_output(concatenated_columns, order)
229
+
230
+ # a slightly faster, case specific implementation of the hexagonal convolution
231
+ def operation_with_single_hexbase_stride(self, input):
232
+ columns_mod2 = input.size(-1) % 2
233
+ odd_kernels_odd_columns = []
234
+ odd_kernels_even_columns = []
235
+ even_kernels_all_columns = []
236
+
237
+ even_kernels_all_columns = self.process(
238
+ self.get_padded_input(input, [0, 0, self.hexbase_size, self.hexbase_size]),
239
+ self.kernel0,
240
+ stride=(1, 1) if self.dimensions == 2 else (self.depth_stride, 1, 1),
241
+ **self.kwargs,
242
+ )
243
+ if self.hexbase_size >= 1:
244
+ odd_kernels_odd_columns = self.process(
245
+ self.get_padded_input(
246
+ input, [1, columns_mod2, self.hexbase_size, self.hexbase_size - 1]
247
+ ),
248
+ self.kernel1,
249
+ dilation=self.get_dilation((1, 2)),
250
+ stride=self.get_stride(),
251
+ )
252
+ odd_kernels_even_columns = self.process(
253
+ self.get_padded_input(
254
+ input,
255
+ [0, 1 - columns_mod2, self.hexbase_size - 1, self.hexbase_size],
256
+ ),
257
+ self.kernel1,
258
+ dilation=self.get_dilation((1, 2)),
259
+ stride=self.get_stride(),
260
+ )
261
+
262
+ if self.hexbase_size > 1:
263
+ for i in range(2, self.hexbase_size + 1):
264
+ if i % 2 == 0:
265
+ even_kernels_all_columns = self.combine(
266
+ even_kernels_all_columns,
267
+ self.process(
268
+ self.get_padded_input(
269
+ input,
270
+ [
271
+ i,
272
+ i,
273
+ self.hexbase_size - int(i / 2),
274
+ self.hexbase_size - int(i / 2),
275
+ ],
276
+ ),
277
+ getattr(self, "kernel" + str(i)),
278
+ dilation=self.get_dilation((1, 2 * i)),
279
+ stride=(1, 1) if self.dimensions == 2 else (self.depth_stride, 1, 1),
280
+ ),
281
+ )
282
+ else:
283
+ x = self.hexbase_size + int((1 - i) / 2)
284
+ odd_kernels_odd_columns = self.combine(
285
+ odd_kernels_odd_columns,
286
+ self.process(
287
+ self.get_padded_input(input, [i, i - 1 + columns_mod2, x, x - 1]),
288
+ getattr(self, "kernel" + str(i)),
289
+ dilation=self.get_dilation((1, 2 * i)),
290
+ stride=self.get_stride(),
291
+ ),
292
+ )
293
+ odd_kernels_even_columns = self.combine(
294
+ odd_kernels_even_columns,
295
+ self.process(
296
+ self.get_padded_input(input, [i - 1, i - columns_mod2, x - 1, x]),
297
+ getattr(self, "kernel" + str(i)),
298
+ dilation=self.get_dilation((1, 2 * i)),
299
+ stride=self.get_stride(),
300
+ ),
301
+ )
302
+
303
+ odd_kernels_concatenated_columns = torch.cat(
304
+ (odd_kernels_odd_columns, odd_kernels_even_columns), 1 + self.dimensions
305
+ )
306
+
307
+ n_odd_columns = odd_kernels_odd_columns.size(-1)
308
+ n_even_columns = odd_kernels_even_columns.size(-1)
309
+ if n_odd_columns == n_even_columns:
310
+ order = [int(i + x * n_even_columns) for i in range(n_even_columns) for x in range(2)]
311
+ else:
312
+ order = [int(i + x * n_odd_columns) for i in range(n_even_columns) for x in range(2)]
313
+ order.append(n_even_columns)
314
+
315
+ return self.combine(
316
+ even_kernels_all_columns,
317
+ self.get_ordered_output(odd_kernels_concatenated_columns, order),
318
+ )
319
+
320
+
321
+ # ----------------------------------------------------------------------------
322
+ # Ring sharing (share_neighbors): tie weights by hexagonal ring, like TDSCAN.
323
+ # The hexagdly offset layout has no clean closed-form hex distance, so the ring
324
+ # index of every kernel cell is derived EMPIRICALLY: a single-tap impulse through
325
+ # the conv reveals each cell's physical (row, col) offset, and the ring is the
326
+ # smallest kernel size whose support contains it. Exact, framework-self-consistent.
327
+ # ----------------------------------------------------------------------------
328
+
329
+ _RING_MAP_CACHE = {}
330
+
331
+
332
+ def _tap_offset(n, i, r, c):
333
+ """Physical (dr, dc) offset of sub-kernel cell (i, r, c) for kernel size n."""
334
+ g = 6 * n + 11
335
+ cen = g // 2
336
+ imp = torch.zeros(1, 1, g, g)
337
+ imp[0, 0, cen, cen] = 1.0
338
+ sub_kernels = []
339
+ for k in range(n + 1):
340
+ kh = 2 * n + 1 - k
341
+ kw = 1 if k == 0 else 2
342
+ a = np.zeros((1, 1, kh, kw), dtype=np.float32)
343
+ if k == i:
344
+ a[0, 0, r, c] = 1.0
345
+ sub_kernels.append(a)
346
+ layer = Conv2d_CustomKernel(sub_kernels=sub_kernels, stride=1)
347
+ out = layer(imp).detach().numpy()[0, 0]
348
+ pos = np.argwhere(np.isclose(out, 1.0))
349
+ return int(pos[0][0] - cen), int(pos[0][1] - cen)
350
+
351
+
352
+ def ring_maps_2d(n):
353
+ """Return ``(ring_maps, num_rings)`` for kernel size ``n`` (see hexagdly_tf)."""
354
+ if n in _RING_MAP_CACHE:
355
+ return _RING_MAP_CACHE[n]
356
+ support = {}
357
+ for ks in range(1, n + 1):
358
+ offs = set()
359
+ for i in range(ks + 1):
360
+ rows = 2 * ks + 1 - i
361
+ cols = 1 if i == 0 else 2
362
+ for r in range(rows):
363
+ for c in range(cols):
364
+ offs.add(_tap_offset(ks, i, r, c))
365
+ support[ks] = offs
366
+
367
+ def ring_of(off):
368
+ if off == (0, 0):
369
+ return 0
370
+ for ks in range(1, n + 1):
371
+ if off in support[ks]:
372
+ return ks
373
+ raise ValueError(f"offset {off} not within kernel size {n}")
374
+
375
+ ring_maps = []
376
+ for i in range(n + 1):
377
+ rows = 2 * n + 1 - i
378
+ cols = 1 if i == 0 else 2
379
+ m = np.zeros((rows, cols), dtype=np.int64)
380
+ for r in range(rows):
381
+ for c in range(cols):
382
+ m[r, c] = ring_of(_tap_offset(n, i, r, c))
383
+ ring_maps.append(m)
384
+ result = (ring_maps, n + 1)
385
+ _RING_MAP_CACHE[n] = result
386
+ return result
387
+
388
+
389
+ class Conv2d(HexBase, nn.Module):
390
+ r"""Applies a 2D hexagonal convolution`
391
+
392
+ Args:
393
+ in_channels: int: number of input channels
394
+ out_channels: int: number of output channels
395
+ kernel_size: int: number of layers with neighbouring pixels
396
+ covered by the pooling kernel
397
+ stride: int: length of strides
398
+ bias: bool: add bias if True (default)
399
+ debug: bool: switch to debug mode
400
+ False: weights are initalised with
401
+ kaiming normal, bias with 0.01 (default)
402
+ True: weights / bias are set to 1.
403
+ share_neighbors: bool: tie weights by hexagonal ring (default: False)
404
+
405
+ Examples::
406
+
407
+ >>> conv2d = pytorch_hexagdly.Conv2d(1,3,2,1)
408
+ >>> input = torch.randn(1, 1, 4, 2)
409
+ >>> output = conv2d(input)
410
+ >>> print(output)
411
+ """
412
+
413
+ def __init__(
414
+ self,
415
+ in_channels,
416
+ out_channels,
417
+ kernel_size=1,
418
+ stride=1,
419
+ bias=True,
420
+ debug=False,
421
+ share_neighbors=False,
422
+ ):
423
+ super(Conv2d, self).__init__()
424
+ self.in_channels = in_channels
425
+ self.out_channels = out_channels
426
+ self.hexbase_size = kernel_size
427
+ self.hexbase_stride = stride
428
+ self.debug = debug
429
+ self.bias = bias
430
+ self.share_neighbors = share_neighbors
431
+ self.dimensions = 2
432
+ self.process = F.conv2d
433
+ self.combine = torch.add
434
+
435
+ if share_neighbors:
436
+ # One weight per hex ring; broadcast to every cell at forward time.
437
+ self._ring_maps, self.num_rings = ring_maps_2d(self.hexbase_size)
438
+ self._ring_idx = [torch.as_tensor(m, dtype=torch.long) for m in self._ring_maps]
439
+ self.ring_weights = Parameter(torch.Tensor(out_channels, in_channels, self.num_rings))
440
+ else:
441
+ for i in range(self.hexbase_size + 1):
442
+ setattr(
443
+ self,
444
+ "kernel" + str(i),
445
+ Parameter(
446
+ torch.Tensor(
447
+ out_channels,
448
+ in_channels,
449
+ 1 + 2 * self.hexbase_size - i,
450
+ 1 if i == 0 else 2,
451
+ )
452
+ ),
453
+ )
454
+ if self.bias:
455
+ self.bias_tensor = Parameter(torch.Tensor(out_channels))
456
+ self.kwargs = {"bias": self.bias_tensor}
457
+ else:
458
+ self.kwargs = {"bias": None}
459
+ self.init_parameters(self.debug)
460
+
461
+ def init_parameters(self, debug):
462
+ if self.share_neighbors:
463
+ if debug:
464
+ nn.init.constant_(self.ring_weights, 1)
465
+ else:
466
+ nn.init.kaiming_normal_(self.ring_weights)
467
+ if self.bias:
468
+ nn.init.constant_(self.kwargs["bias"], 1.0 if debug else 0.01)
469
+ return
470
+ if debug:
471
+ for i in range(self.hexbase_size + 1):
472
+ nn.init.constant_(getattr(self, "kernel" + str(i)), 1)
473
+ if self.bias:
474
+ nn.init.constant_(getattr(self, "kwargs")["bias"], 1.0)
475
+ else:
476
+ for i in range(self.hexbase_size + 1):
477
+ nn.init.kaiming_normal_(getattr(self, "kernel" + str(i)))
478
+ if self.bias:
479
+ nn.init.constant_(getattr(self, "kwargs")["bias"], 0.01)
480
+
481
+ def _materialize_shared_kernels(self):
482
+ """kernel{i}[:, :, r, c] = ring_weights[:, :, ring_map[i][r, c]].
483
+
484
+ Gathers along the ring axis -> dense (out, in, rows, cols) kernels, so
485
+ the forward pass is unchanged and gradients flow back into ring_weights
486
+ (all cells of a ring share one weight), exactly like TDSCAN.
487
+ """
488
+ for i in range(self.hexbase_size + 1):
489
+ idx = self._ring_idx[i].to(self.ring_weights.device) # (rows, cols)
490
+ # ring_weights: (out, in, num_rings) -> index_select on last axis,
491
+ # then reshape to (out, in, rows, cols).
492
+ flat = torch.index_select(self.ring_weights, 2, idx.reshape(-1))
493
+ setattr(
494
+ self,
495
+ "kernel" + str(i),
496
+ flat.reshape(self.out_channels, self.in_channels, *idx.shape),
497
+ )
498
+
499
+ def forward(self, input):
500
+ if self.share_neighbors:
501
+ self._materialize_shared_kernels()
502
+ if self.hexbase_stride == 1:
503
+ return self.operation_with_single_hexbase_stride(input)
504
+ else:
505
+ return self.operation_with_arbitrary_stride(input)
506
+
507
+ def __repr__(self):
508
+ s = (
509
+ "{name}({in_channels}, {out_channels}, kernel_size={hexbase_size}"
510
+ ", stride={hexbase_stride}"
511
+ )
512
+ if self.bias is False:
513
+ s += ", bias=False"
514
+ if self.debug is True:
515
+ s += ", debug=True"
516
+ s += ")"
517
+ return s.format(name=self.__class__.__name__, **self.__dict__)
518
+
519
+
520
+ class Conv2d_CustomKernel(HexBase, nn.Module):
521
+ r"""Applies a 2D hexagonal convolution with custom kernels`
522
+
523
+ Args:
524
+ sub_kernels: list: list containing sub-kernels as numpy arrays
525
+ stride: int: length of strides
526
+ bias: array: numpy array with biases (default: None)
527
+ requires_grad: bool: trainable parameters if True (default: False)
528
+ debug: bool: If True a kernel of size one with all values
529
+ set to 1 will be applied as well as no bias
530
+ (default: False)
531
+
532
+ Examples::
533
+
534
+ Given in the online repository https://github.com/ai4iacts/hexagdly
535
+ """
536
+
537
+ def __init__(self, sub_kernels=[], stride=1, bias=None, requires_grad=False, debug=False):
538
+ super(Conv2d_CustomKernel, self).__init__()
539
+ self.sub_kernels = sub_kernels
540
+ self.bias_array = bias
541
+ self.hexbase_stride = stride
542
+ self.requires_grad = requires_grad
543
+ self.debug = debug
544
+ self.dimensions = 2
545
+ self.process = F.conv2d
546
+ self.combine = torch.add
547
+
548
+ self.init_parameters(self.debug)
549
+
550
+ def init_parameters(self, debug):
551
+ if debug or len(self.sub_kernels) == 0:
552
+ print("The debug kernel is used for {name}!".format(name=self.__class__.__name__))
553
+ self.sub_kernels = [
554
+ np.array([[[[1], [1], [1]]]]),
555
+ np.array([[[[1, 1], [1, 1]]]]),
556
+ ]
557
+ self.hexbase_size = len(self.sub_kernels) - 1
558
+ self.check_sub_kernels()
559
+
560
+ for i in range(self.hexbase_size + 1):
561
+ setattr(
562
+ self,
563
+ "kernel" + str(i),
564
+ Parameter(
565
+ torch.from_numpy(self.sub_kernels[i]).type(torch.FloatTensor),
566
+ requires_grad=self.requires_grad,
567
+ ),
568
+ )
569
+
570
+ if not debug and self.bias_array is not None:
571
+ self.check_bias()
572
+ self.bias_tensor = Parameter(
573
+ torch.from_numpy(self.bias_array).type(torch.FloatTensor),
574
+ requires_grad=self.requires_grad,
575
+ )
576
+ self.kwargs = {"bias": self.bias_tensor}
577
+ self.bias = True
578
+ else:
579
+ self.bias = False
580
+ if self.bias_array is not None:
581
+ print(
582
+ "{name}: Bias is not used in debug mode!".format(name=self.__class__.__name__)
583
+ )
584
+
585
+ def check_sub_kernels(self):
586
+ for i in range(self.hexbase_size + 1):
587
+ assert type(self.sub_kernels[i]).__module__ == np.__name__, (
588
+ "sub-kernels must be given as numpy arrays"
589
+ )
590
+ assert len(self.sub_kernels[i].shape) == 4, (
591
+ "sub-kernels must be of rank 4 for a 2d convolution"
592
+ )
593
+ if i == 0:
594
+ assert self.sub_kernels[i].shape[3] == 1, "first sub-kernel must have only 1 column"
595
+ assert self.sub_kernels[i].shape[2] == 2 * self.hexbase_size + 1, (
596
+ "first sub-kernel must have 2* (kernel size) + 1 rows"
597
+ )
598
+ self.out_channels = self.sub_kernels[i].shape[0]
599
+ self.in_channels = self.sub_kernels[i].shape[1]
600
+ else:
601
+ assert self.sub_kernels[i].shape[3] == 2, (
602
+ "sub-kernel {}: all but the first sub-kernel must have 2 columns".format(i)
603
+ )
604
+ assert self.sub_kernels[i].shape[2] == 2 * self.hexbase_size + 1 - i, (
605
+ "{}. sub-kernel must have 2* (kernel size) + 1 - {} rows".format(i, i)
606
+ )
607
+ assert self.sub_kernels[i].shape[0] == self.out_channels, (
608
+ "sub-kernel {}: out channels are not consistent".format(i)
609
+ )
610
+ assert self.sub_kernels[i].shape[1] == self.in_channels, (
611
+ "sub-kernel {}: in channels are not consistent".format(i)
612
+ )
613
+
614
+ def check_bias(self):
615
+ assert type(self.bias_array).__module__ == np.__name__, (
616
+ "bias must be given as a numpy array"
617
+ )
618
+ assert len(self.bias_array.shape) == 1, "bias must be of rank 1"
619
+ assert self.bias_array.shape[0] == self.out_channels, (
620
+ "bias must have length equal to number of out channels"
621
+ )
622
+
623
+ def forward(self, input):
624
+ if self.hexbase_stride == 1:
625
+ return self.operation_with_single_hexbase_stride(input)
626
+ else:
627
+ return self.operation_with_arbitrary_stride(input)
628
+
629
+ def __repr__(self):
630
+ s = (
631
+ "{name}({in_channels}, {out_channels}, kernel_size={hexbase_size}"
632
+ ", stride={hexbase_stride}"
633
+ )
634
+ if self.bias is False:
635
+ s += ", bias=False"
636
+ if self.debug is True:
637
+ s += ", debug=True"
638
+ s += ")"
639
+ return s.format(name=self.__class__.__name__, **self.__dict__)
640
+
641
+
642
+ class Conv3d(HexBase, nn.Module):
643
+ r"""Applies a 3D hexagonal convolution`
644
+
645
+ Args:
646
+ in_channels: int: number of input channels
647
+ out_channels: int: number of output channels
648
+ kernel_size: int, tuple: number of layers with neighbouring pixels
649
+ covered by the pooling kernel
650
+ int: same number of layers in all dimensions
651
+ tuple of two ints:
652
+ 1st int: layers in depth
653
+ 2nd int: layers in hexagonal base
654
+ stride: int, tuple: length of strides
655
+ int: same lenght of strides in each dimension
656
+ tuple of two ints:
657
+ 1st int: length of strides in depth
658
+ 2nd int: length of strides in hexagonal base
659
+ bias: bool: add bias if True (default)
660
+ debug: bool: switch to debug mode
661
+ False: weights are initalised with
662
+ kaiming normal, bias with 0.01 (default)
663
+ True: weights / bias are set to 1.
664
+ share_neighbors: bool: tie weights by hexagonal ring (default: False)
665
+ depth_padding: str: 'valid' (default) or 'same' — 'same' zero-pads
666
+ the depth axis so output depth equals input depth
667
+
668
+ Examples::
669
+
670
+ >>> conv3d = pytorch_hexagdly.Conv3d((1,1), (2,2))
671
+ >>> input = torch.randn(1, 1, 6, 5, 4)
672
+ >>> output = conv3d(input)
673
+ >>> print(output)
674
+ """
675
+
676
+ def __init__(
677
+ self,
678
+ in_channels,
679
+ out_channels,
680
+ kernel_size=1,
681
+ stride=1,
682
+ bias=True,
683
+ debug=False,
684
+ share_neighbors=False,
685
+ depth_padding="valid",
686
+ ):
687
+ super(Conv3d, self).__init__()
688
+ if depth_padding not in ("valid", "same"):
689
+ raise ValueError("depth_padding must be 'valid' or 'same'.")
690
+ self.depth_padding = depth_padding
691
+ self.in_channels = in_channels
692
+ self.out_channels = out_channels
693
+ if isinstance(kernel_size, int):
694
+ self.hexbase_size = kernel_size
695
+ self.depth_size = kernel_size
696
+ elif isinstance(kernel_size, tuple):
697
+ assert len(kernel_size) == 2, "Need a tuple of two ints to set kernel size"
698
+ self.hexbase_size = kernel_size[1]
699
+ self.depth_size = kernel_size[0]
700
+ if isinstance(stride, int):
701
+ self.hexbase_stride = stride
702
+ self.depth_stride = stride
703
+ elif isinstance(stride, tuple):
704
+ assert len(stride) == 2, "Need a tuple of two ints to set stride"
705
+ self.hexbase_stride = stride[1]
706
+ self.depth_stride = stride[0]
707
+ self.debug = debug
708
+ self.bias = bias
709
+ self.share_neighbors = share_neighbors
710
+ self.dimensions = 3
711
+ self.process = F.conv3d
712
+ self.combine = torch.add
713
+
714
+ if share_neighbors:
715
+ # Share over the hex axes only; depth (time) stays independent ->
716
+ # ring_weights is (out, in, depth, num_rings) (cf. TDSCAN L x rings).
717
+ self._ring_maps, self.num_rings = ring_maps_2d(self.hexbase_size)
718
+ self._ring_idx = [torch.as_tensor(m, dtype=torch.long) for m in self._ring_maps]
719
+ self.ring_weights = Parameter(
720
+ torch.Tensor(out_channels, in_channels, self.depth_size, self.num_rings)
721
+ )
722
+ else:
723
+ for i in range(self.hexbase_size + 1):
724
+ setattr(
725
+ self,
726
+ "kernel" + str(i),
727
+ Parameter(
728
+ torch.Tensor(
729
+ out_channels,
730
+ in_channels,
731
+ self.depth_size,
732
+ 1 + 2 * self.hexbase_size - i,
733
+ 1 if i == 0 else 2,
734
+ )
735
+ ),
736
+ )
737
+ if self.bias:
738
+ self.bias_tensor = Parameter(torch.Tensor(out_channels))
739
+ self.kwargs = {"bias": self.bias_tensor}
740
+ else:
741
+ self.kwargs = {"bias": None}
742
+
743
+ self.init_parameters(self.debug)
744
+
745
+ def init_parameters(self, debug):
746
+ if self.share_neighbors:
747
+ if debug:
748
+ nn.init.constant_(self.ring_weights, 1)
749
+ else:
750
+ nn.init.kaiming_normal_(self.ring_weights)
751
+ if self.bias:
752
+ nn.init.constant_(self.kwargs["bias"], 1.0 if debug else 0.01)
753
+ return
754
+ if debug:
755
+ for i in range(self.hexbase_size + 1):
756
+ nn.init.constant_(getattr(self, "kernel" + str(i)), 1)
757
+ if self.bias:
758
+ nn.init.constant_(getattr(self, "kwargs")["bias"], 1.0)
759
+ else:
760
+ for i in range(self.hexbase_size + 1):
761
+ nn.init.kaiming_normal_(getattr(self, "kernel" + str(i)))
762
+ if self.bias:
763
+ nn.init.constant_(getattr(self, "kwargs")["bias"], 0.01)
764
+
765
+ def _materialize_shared_kernels(self):
766
+ """kernel{i} = ring_weights gathered along the ring axis into
767
+ (out, in, depth, rows, cols); depth left independent."""
768
+ for i in range(self.hexbase_size + 1):
769
+ idx = self._ring_idx[i].to(self.ring_weights.device) # (rows, cols)
770
+ flat = torch.index_select(self.ring_weights, 3, idx.reshape(-1))
771
+ setattr(
772
+ self,
773
+ "kernel" + str(i),
774
+ flat.reshape(self.out_channels, self.in_channels, self.depth_size, *idx.shape),
775
+ )
776
+
777
+ def forward(self, input):
778
+ if self.share_neighbors:
779
+ self._materialize_shared_kernels()
780
+ if self.depth_padding == "same":
781
+ # Symmetric zero-pad the depth axis (NCDHW -> axis 2) so the temporal
782
+ # kernel is centred and output depth == input depth, like TDSCAN.
783
+ pad = (self.depth_size - 1) // 2
784
+ top = pad
785
+ bot = self.depth_size - 1 - pad
786
+ input = F.pad(input, [0, 0, 0, 0, top, bot]) # pads last dims; here D
787
+ if self.hexbase_stride == 1:
788
+ return self.operation_with_single_hexbase_stride(input)
789
+ else:
790
+ return self.operation_with_arbitrary_stride(input)
791
+
792
+ def __repr__(self):
793
+ s = (
794
+ "{name}({in_channels}, {out_channels}, kernel_size=({depth_size}, {hexbase_size})"
795
+ ", stride=({depth_stride}, {hexbase_stride})"
796
+ )
797
+ if self.bias is False:
798
+ s += ", bias=False"
799
+ if self.debug is True:
800
+ s += ", debug=True"
801
+ s += ")"
802
+ return s.format(name=self.__class__.__name__, **self.__dict__)
803
+
804
+
805
+ class Conv3d_CustomKernel(HexBase, nn.Module):
806
+ r"""Applies a 3D hexagonal convolution with custom kernels`
807
+
808
+ Args:
809
+ sub_kernels: list: list containing sub-kernels as numpy arrays
810
+ stride: stride: int, tuple: length of strides
811
+ int: same lenght of strides in each dimension
812
+ tuple of two ints:
813
+ 1st int: length of strides in depth
814
+ 2nd int: length of strides in hexagonal base
815
+ requires_grad: bool: trainable parameters if True (default: False)
816
+ debug: bool: If True a kernel of size one with all values
817
+ set to 1 will be applied as well as no bias
818
+ (default: False)
819
+
820
+ Examples::
821
+
822
+ Given in the online repository https://github.com/ai4iacts/hexagdly
823
+ """
824
+
825
+ def __init__(self, sub_kernels=[], stride=1, bias=None, requires_grad=False, debug=False):
826
+ super(Conv3d_CustomKernel, self).__init__()
827
+ self.sub_kernels = sub_kernels
828
+ self.bias_array = bias
829
+ if isinstance(stride, int):
830
+ self.hexbase_stride = stride
831
+ self.depth_stride = stride
832
+ elif isinstance(stride, tuple):
833
+ assert len(stride) == 2, "Need a tuple of two ints to set stride"
834
+ self.hexbase_stride = stride[1]
835
+ self.depth_stride = stride[0]
836
+ self.requires_grad = requires_grad
837
+ self.debug = debug
838
+ self.dimensions = 3
839
+ self.process = F.conv3d
840
+ self.combine = torch.add
841
+
842
+ self.init_parameters(self.debug)
843
+
844
+ def init_parameters(self, debug):
845
+ if debug or len(self.sub_kernels) == 0:
846
+ print("The debug kernel is used for {name}!".format(name=self.__class__.__name__))
847
+ self.sub_kernels = [
848
+ np.array([[[[[1], [1], [1]]]]]),
849
+ np.array([[[[[1, 1], [1, 1]]]]]),
850
+ ]
851
+ self.hexbase_size = len(self.sub_kernels) - 1
852
+ self.check_sub_kernels()
853
+
854
+ for i in range(self.hexbase_size + 1):
855
+ setattr(
856
+ self,
857
+ "kernel" + str(i),
858
+ Parameter(
859
+ torch.from_numpy(self.sub_kernels[i]).type(torch.FloatTensor),
860
+ requires_grad=self.requires_grad,
861
+ ),
862
+ )
863
+
864
+ if not debug and self.bias_array is not None:
865
+ self.check_bias()
866
+ self.bias_tensor = Parameter(
867
+ torch.from_numpy(self.bias_array).type(torch.FloatTensor),
868
+ requires_grad=self.requires_grad,
869
+ )
870
+ self.kwargs = {"bias": self.bias_tensor}
871
+ self.bias = True
872
+ else:
873
+ self.bias = False
874
+ print("No bias is used for {name}!".format(name=self.__class__.__name__))
875
+
876
+ def check_sub_kernels(self):
877
+ for i in range(self.hexbase_size + 1):
878
+ assert type(self.sub_kernels[i]).__module__ == np.__name__, (
879
+ "sub-kernels must be given as numpy arrays"
880
+ )
881
+ assert len(self.sub_kernels[i].shape) == 5, (
882
+ "sub-kernels must be of rank 5 for a 3d convolution"
883
+ )
884
+ if i == 0:
885
+ assert self.sub_kernels[i].shape[4] == 1, "first sub-kernel must have only 1 column"
886
+ assert self.sub_kernels[i].shape[3] == 2 * self.hexbase_size + 1, (
887
+ "first sub-kernel must have 2* (kernel size) + 1 rows"
888
+ )
889
+ self.out_channels = self.sub_kernels[i].shape[0]
890
+ self.in_channels = self.sub_kernels[i].shape[1]
891
+ self.depth_size = self.sub_kernels[i].shape[2]
892
+ else:
893
+ assert self.sub_kernels[i].shape[4] == 2, (
894
+ "sub-kernel {}: all but the first sub-kernel must have 2 columns".format(i)
895
+ )
896
+ assert self.sub_kernels[i].shape[3] == 2 * self.hexbase_size + 1 - i, (
897
+ "{}th sub-kernel must have 2* (kernel size) + 1 - {} rows".format(i, i)
898
+ )
899
+ assert self.sub_kernels[i].shape[0] == self.out_channels, (
900
+ "sub-kernel {}: out channels are not consistent".format(i)
901
+ )
902
+ assert self.sub_kernels[i].shape[1] == self.in_channels, (
903
+ "sub-kernel {}: out channels are not consistent".format(i)
904
+ )
905
+ assert self.sub_kernels[i].shape[2] == self.depth_size, (
906
+ "sub-kernel {}: depths are not consistent".format(i)
907
+ )
908
+
909
+ def check_bias(self):
910
+ assert type(self.bias_array).__module__ == np.__name__, (
911
+ "bias must be given as a numpy array"
912
+ )
913
+ assert len(self.bias_array.shape) == 1, "bias must be of rank 1"
914
+ assert self.bias_array.shape[0] == self.out_channels, (
915
+ "bias must have length equal to number of out channels"
916
+ )
917
+
918
+ def forward(self, input):
919
+ if self.hexbase_stride == 1:
920
+ return self.operation_with_single_hexbase_stride(input)
921
+ else:
922
+ return self.operation_with_arbitrary_stride(input)
923
+
924
+ def __repr__(self):
925
+ s = (
926
+ "{name}({in_channels}, {out_channels}, kernel_size=({depth_size}, {hexbase_size})"
927
+ ", stride=({depth_stride}, {hexbase_stride})"
928
+ )
929
+ if self.bias is False:
930
+ s += ", bias=False"
931
+ if self.debug is True:
932
+ s += ", debug=True"
933
+ s += ")"
934
+ return s.format(name=self.__class__.__name__, **self.__dict__)
935
+
936
+
937
+ class MaxPool2d(HexBase, nn.Module):
938
+ r"""Applies a 2D hexagonal max pooling`
939
+
940
+ Args:
941
+ kernel_size: int: number of layers with neighbouring pixels
942
+ covered by the pooling kernel
943
+ stride: int: length of strides
944
+
945
+ Examples::
946
+
947
+ >>> maxpool2d = pytorch_hexagdly.MaxPool2d(1,2)
948
+ >>> input = torch.randn(1, 1, 4, 2)
949
+ >>> output = maxpool2d(input)
950
+ >>> print(output)
951
+ """
952
+
953
+ def __init__(self, kernel_size=1, stride=1):
954
+ super(MaxPool2d, self).__init__()
955
+ self.hexbase_size = kernel_size
956
+ self.hexbase_stride = stride
957
+ self.dimensions = 2
958
+ self.process = F.max_pool2d
959
+ self.combine = torch.max
960
+
961
+ for i in range(self.hexbase_size + 1):
962
+ setattr(
963
+ self,
964
+ "kernel" + str(i),
965
+ (1 + 2 * self.hexbase_size - i, 1 if i == 0 else 2),
966
+ )
967
+
968
+ def forward(self, input):
969
+ if self.hexbase_stride == 1:
970
+ return self.operation_with_single_hexbase_stride(input)
971
+ else:
972
+ return self.operation_with_arbitrary_stride(input)
973
+
974
+ def __repr__(self):
975
+ s = "{name}(kernel_size={hexbase_size}, stride={hexbase_stride})"
976
+ return s.format(name=self.__class__.__name__, **self.__dict__)
977
+
978
+
979
+ class MaxPool3d(HexBase, nn.Module):
980
+ r"""Applies a 3D hexagonal max pooling`
981
+
982
+ Args:
983
+ kernel_size: int, tuple: number of layers with neighbouring pixels
984
+ covered by the pooling kernel
985
+ int: same number of layers in all dimensions
986
+ tuple of two ints:
987
+ 1st int: layers in depth
988
+ 2nd int: layers in hexagonal base
989
+ stride: int, tuple: length of strides
990
+ int: same lenght of strides in each dimension
991
+ tuple of two ints:
992
+ 1st int: length of strides in depth
993
+ 2nd int: length of strides in hexagonal base
994
+
995
+ Examples::
996
+
997
+ >>> maxpool3d = pytorch_hexagdly.MaxPool3d((1,1), (2,2))
998
+ >>> input = torch.randn(1, 1, 6, 5, 4)
999
+ >>> output = maxpool3d(input)
1000
+ >>> print(output)
1001
+ """
1002
+
1003
+ def __init__(self, kernel_size=1, stride=1):
1004
+ super(MaxPool3d, self).__init__()
1005
+ if isinstance(kernel_size, int):
1006
+ self.hexbase_size = kernel_size
1007
+ self.depth_size = kernel_size
1008
+ elif isinstance(kernel_size, tuple):
1009
+ assert len(kernel_size) == 2, "Too many parameters"
1010
+ self.hexbase_size = kernel_size[1]
1011
+ self.depth_size = kernel_size[0]
1012
+ if isinstance(stride, int):
1013
+ self.hexbase_stride = stride
1014
+ self.depth_stride = stride
1015
+ elif isinstance(stride, tuple):
1016
+ assert len(stride) == 2, "Too many parameters"
1017
+ self.hexbase_stride = stride[1]
1018
+ self.depth_stride = stride[0]
1019
+ self.dimensions = 3
1020
+ self.process = F.max_pool3d
1021
+ self.combine = torch.max
1022
+
1023
+ for i in range(self.hexbase_size + 1):
1024
+ setattr(
1025
+ self,
1026
+ "kernel" + str(i),
1027
+ (self.depth_size, 1 + 2 * self.hexbase_size - i, 1 if i == 0 else 2),
1028
+ )
1029
+
1030
+ def forward(self, input):
1031
+ if self.hexbase_stride == 1:
1032
+ return self.operation_with_single_hexbase_stride(input)
1033
+ else:
1034
+ return self.operation_with_arbitrary_stride(input)
1035
+
1036
+ def __repr__(self):
1037
+ s = (
1038
+ "{name}(kernel_size=({depth_size}, {hexbase_size})"
1039
+ ", stride=({depth_stride}, {hexbase_stride}))"
1040
+ )
1041
+ return s.format(name=self.__class__.__name__, **self.__dict__)
@@ -0,0 +1,238 @@
1
+ Metadata-Version: 2.4
2
+ Name: pytorch-hexagdly
3
+ Version: 0.1.0
4
+ Summary: Hexagonal convolution and pooling layers for PyTorch — an extended HexagDLy fork
5
+ Project-URL: Homepage, https://github.com/YugnatD/pytorch-hexagdly
6
+ Project-URL: Source, https://github.com/YugnatD/pytorch-hexagdly
7
+ Project-URL: Issues, https://github.com/YugnatD/pytorch-hexagdly/issues
8
+ Project-URL: Changelog, https://github.com/YugnatD/pytorch-hexagdly/blob/master/CHANGELOG.md
9
+ Project-URL: Original (HexagDLy), https://github.com/ai4iacts/hexagdly
10
+ Author: Tanguy Dietrich
11
+ License: MIT License
12
+
13
+ Copyright (c) 2018 ai4iacts (HexagDLy authors: Tim Lukas Holch, Constantin Steppa)
14
+ Copyright (c) 2026 Tanguy Dietrich, HEPIA, SST-1M Collaboration (pytorch-hexagdly fork)
15
+
16
+ Permission is hereby granted, free of charge, to any person obtaining a copy
17
+ of this software and associated documentation files (the "Software"), to deal
18
+ in the Software without restriction, including without limitation the rights
19
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
20
+ copies of the Software, and to permit persons to whom the Software is
21
+ furnished to do so, subject to the following conditions:
22
+
23
+ The above copyright notice and this permission notice shall be included in all
24
+ copies or substantial portions of the Software.
25
+
26
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
27
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
28
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
29
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
30
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
31
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
32
+ SOFTWARE.
33
+ License-File: LICENSE
34
+ License-File: NOTICE.md
35
+ Keywords: astroparticle-physics,cherenkov,cnn,convolution,deep-learning,equivariant,geometric-deep-learning,hexagdly,hexagonal,hexagonal-convolution,hexagonal-grid,iact,neural-networks,pytorch,torch
36
+ Classifier: Development Status :: 4 - Beta
37
+ Classifier: Intended Audience :: Developers
38
+ Classifier: Intended Audience :: Science/Research
39
+ Classifier: License :: OSI Approved :: MIT License
40
+ Classifier: Operating System :: OS Independent
41
+ Classifier: Programming Language :: Python :: 3
42
+ Classifier: Programming Language :: Python :: 3.9
43
+ Classifier: Programming Language :: Python :: 3.10
44
+ Classifier: Programming Language :: Python :: 3.11
45
+ Classifier: Programming Language :: Python :: 3.12
46
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
47
+ Classifier: Topic :: Scientific/Engineering :: Image Recognition
48
+ Classifier: Topic :: Scientific/Engineering :: Physics
49
+ Requires-Python: >=3.9
50
+ Requires-Dist: numpy
51
+ Requires-Dist: torch
52
+ Provides-Extra: dev
53
+ Requires-Dist: pytest; extra == 'dev'
54
+ Description-Content-Type: text/markdown
55
+
56
+ # pytorch-hexagdly — Hexagonal Convolutions for PyTorch
57
+
58
+ `pytorch-hexagdly` is a fork of [HexagDLy](https://github.com/ai4iacts/hexagdly)
59
+ that extends the original hexagonal convolution and pooling layers for PyTorch with
60
+ two new features: **ring-shared weights** (`share_neighbors`) and **depth-axis same-padding**
61
+ (`depth_padding="same"` on `Conv3d`).
62
+
63
+ - [Getting Started](#getting-started)
64
+ - [New Features](#new-features)
65
+ - [Preparing the Data](#preparing-the-data)
66
+ - [How to use pytorch-hexagdly](#how-to-use-pytorch-hexagdly)
67
+ - [General Concept](#general-concept)
68
+ - [Disclaimer](#disclaimer)
69
+ - [Citing HexagDLy](#citation)
70
+
71
+
72
+ ## Getting Started
73
+
74
+ ### Pip Installation
75
+
76
+ ```
77
+ pip install pytorch-hexagdly
78
+ ```
79
+
80
+ ```python
81
+ import pytorch_hexagdly
82
+ ```
83
+
84
+ To get the dependencies needed to run the provided [unit tests](tests) and
85
+ [notebooks](notebooks), add the `dev` option:
86
+
87
+ ```
88
+ pip install pytorch-hexagdly[dev]
89
+ ```
90
+
91
+ ### Manual Installation
92
+
93
+ Requires a working installation of [PyTorch](https://github.com/pytorch/pytorch).
94
+ Clone the repository and install in editable mode:
95
+
96
+ ```
97
+ git clone https://github.com/YugnatD/pytorch-hexagdly
98
+ cd pytorch-hexagdly
99
+ pip install -e .
100
+ ```
101
+
102
+
103
+ ## New Features
104
+
105
+ ### `share_neighbors` — ring-shared kernel weights
106
+
107
+ Available on `Conv2d` and `Conv3d`. When set to `True`, all cells at the same
108
+ hexagonal ring distance share a single weight, reducing the number of learnable
109
+ parameters. Ring 0 is the center pixel; ring *r* covers the 6*r* cells at
110
+ hex-distance *r*. This mirrors the TDSCAN triggering approach.
111
+
112
+ ```python
113
+ import torch
114
+ import pytorch_hexagdly
115
+
116
+ conv = pytorch_hexagdly.Conv2d(1, 8, kernel_size=2, stride=1, share_neighbors=True)
117
+ x = torch.randn(1, 1, 21, 21)
118
+ print(conv(x).shape)
119
+ ```
120
+
121
+ ### `depth_padding="same"` — temporal same-padding for `Conv3d`
122
+
123
+ When `depth_padding="same"`, the depth/time axis is zero-padded symmetrically so
124
+ the output depth equals the input depth. The default is `"valid"` (upstream behaviour).
125
+
126
+ ```python
127
+ conv3d = pytorch_hexagdly.Conv3d(1, 4, kernel_size=(3, 1), stride=1,
128
+ depth_padding="same")
129
+ x = torch.randn(1, 1, 10, 21, 21)
130
+ print(conv3d(x).shape) # depth dimension preserved
131
+ ```
132
+
133
+
134
+ ## How to use pytorch-hexagdly
135
+
136
+ As `pytorch-hexagdly` is based on PyTorch, it is of advantage to be familiar with
137
+ PyTorch's functionalities and concepts. Before applying it, ensure that the input
138
+ data has the correct hexagonal layout. An [example notebook](notebooks/how_to_apply_adressing_scheme.ipynb)
139
+ illustrates the steps to get data into the correct format.
140
+
141
+ Basic example:
142
+
143
+ ```python
144
+ import torch
145
+ import pytorch_hexagdly
146
+
147
+ kernel_size, stride = 1, 4
148
+ in_channels, out_channels = 1, 3
149
+
150
+ hexconv = pytorch_hexagdly.Conv2d(in_channels, out_channels, kernel_size, stride)
151
+ input = torch.rand(1, 1, 21, 21)
152
+ output = hexconv(input)
153
+ ```
154
+
155
+ HexagDLy uses an addressing scheme to map hexagonal grid data to a square tensor.
156
+ The layout from top to bottom (along tensor index 2) must be of zig-zag-edge shape
157
+ and from left to right (along tensor index 3) of armchair-edge shape.
158
+
159
+ Additional examples for basic use-cases are shown in the [notebooks](notebooks) folder.
160
+
161
+
162
+ ## General Concept
163
+
164
+ As common deep learning frameworks process data on square grids, hexagonally sampled
165
+ data must be mapped to a square tensor. This conversion is non-trivial due to the
166
+ different symmetries of square (4-fold) vs hexagonal (6-fold) grids.
167
+
168
+ HexagDLy solves this by splitting each convolution kernel into sub-kernels that
169
+ together cover the true neighbours of a data point in the hexagonal grid. A full
170
+ hexagonal convolution with size 1 (next-neighbour kernel) decomposes into three
171
+ sub-convolutions with two different sub-kernels applied to three differently padded
172
+ versions of the input.
173
+
174
+ ![kerne size+stride](figures/kernel_size+stride.png "Examples of different kernels of different size and strides.")
175
+
176
+ **Please note**: Operations are only performed where the center point of a kernel is
177
+ located within the input tensor. This could result in output columns of different
178
+ length; in such cases the output will be sliced according to the shortest column.
179
+
180
+ ![violating_symmetry](figures/violating_symmetry.png "Squeezing hexagonal data in a square grid and applying square convolution kernels disregards the symmetry of the hexagonal lattice.")
181
+
182
+ ![explicit_next_neighbour_conv](figures/explicit_next_neighbour_conv.png "Schematic description of the individual sub-convolutions and combination of the individual outputs to perform a hexagonal convolution.")
183
+
184
+
185
+ ## Disclaimer
186
+
187
+ `pytorch-hexagdly` is built as an easy-to-use prototyping tool to design convolutional
188
+ neural networks for hexagonally sampled data. The implemented methods aim for
189
+ flexibility rather than performance. Once a model is optimized, hard-coding kernel
190
+ size, stride and input dimensions will make the implementation faster.
191
+
192
+
193
+ ## Authors
194
+
195
+ **Fork (`pytorch-hexagdly`)**
196
+ * **Tanguy Dietrich** — HEPIA / SST-1M Collaboration
197
+
198
+ **Original HexagDLy**
199
+ * **Tim Lukas Holch**
200
+ * **Constantin Steppa**
201
+
202
+ See [NOTICE.md](NOTICE.md) for full attribution.
203
+
204
+
205
+ ## License
206
+
207
+ MIT license — see [LICENSE](LICENSE).
208
+
209
+
210
+ ## Citation
211
+
212
+ If this work has helped your research, please cite the original HexagDLy paper:
213
+
214
+ ```bibtex
215
+ @article{hexagdly_paper,
216
+ title = "HexagDLy—Processing hexagonally sampled data with CNNs in PyTorch",
217
+ author = "Constantin Steppa and Tim L. Holch",
218
+ journal = "SoftwareX",
219
+ volume = "9",
220
+ pages = "193 - 198",
221
+ year = "2019",
222
+ issn = "2352-7110",
223
+ doi = "https://doi.org/10.1016/j.softx.2019.02.010",
224
+ url = "https://www.sciencedirect.com/science/article/pii/S2352711018302723",
225
+ keywords = "Convolutional neural networks, Hexagonal grid, PyTorch, Astroparticle physics",
226
+ abstract = "HexagDLy is a Python-library extending the PyTorch deep learning framework with convolution and pooling operations on hexagonal grids. It aims to ease the access to convolutional neural networks for applications that rely on hexagonally sampled data as, for example, commonly found in ground-based astroparticle physics experiments."
227
+ }
228
+ ```
229
+
230
+ HexagDLy was developed as part of a research study in ground-based gamma-ray astronomy
231
+ published in [Astroparticle Physics](https://doi.org/10.1016/j.astropartphys.2018.10.003).
232
+
233
+
234
+ ## Acknowledgments
235
+
236
+ The original HexagDLy project evolved by exploring new analysis techniques for Imaging
237
+ Atmospheric Cherenkov Telescopes with H.E.S.S. The fork was developed in the context of
238
+ the SST-1M Collaboration / HEPIA TDSCAN triggering project.
@@ -0,0 +1,6 @@
1
+ pytorch_hexagdly/__init__.py,sha256=zTdyIR_ahxonhlnS7ocx5b-MOZerp5kP6sbTQT8aJgM,41226
2
+ pytorch_hexagdly-0.1.0.dist-info/METADATA,sha256=37z0dQB9nbF0QB0-u0UcUnIGGnYTBV7ZqekiEFaNs7I,9554
3
+ pytorch_hexagdly-0.1.0.dist-info/WHEEL,sha256=mffPy8wBnZQn2VnJUU5jE99KsxaSfiyMHV9Yt0aLVxs,87
4
+ pytorch_hexagdly-0.1.0.dist-info/licenses/LICENSE,sha256=XBXQ51d01xYORI8cXZa89M3L9RzrfySS0khSDSypNRM,1208
5
+ pytorch_hexagdly-0.1.0.dist-info/licenses/NOTICE.md,sha256=rTymQYhl8cpShbJotr8JgN1G2ye-hjIonYJ5cDbi4n0,1688
6
+ pytorch_hexagdly-0.1.0.dist-info/RECORD,,
@@ -0,0 +1,4 @@
1
+ Wheel-Version: 1.0
2
+ Generator: hatchling 1.30.1
3
+ Root-Is-Purelib: true
4
+ Tag: py3-none-any
@@ -0,0 +1,22 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2018 ai4iacts (HexagDLy authors: Tim Lukas Holch, Constantin Steppa)
4
+ Copyright (c) 2026 Tanguy Dietrich, HEPIA, SST-1M Collaboration (pytorch-hexagdly fork)
5
+
6
+ Permission is hereby granted, free of charge, to any person obtaining a copy
7
+ of this software and associated documentation files (the "Software"), to deal
8
+ in the Software without restriction, including without limitation the rights
9
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
10
+ copies of the Software, and to permit persons to whom the Software is
11
+ furnished to do so, subject to the following conditions:
12
+
13
+ The above copyright notice and this permission notice shall be included in all
14
+ copies or substantial portions of the Software.
15
+
16
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
17
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
18
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
19
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
20
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
21
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
22
+ SOFTWARE.
@@ -0,0 +1,45 @@
1
+ # Notice
2
+
3
+ `pytorch-hexagdly` is a fork of
4
+ [HexagDLy](https://github.com/ai4iacts/hexagdly) by Tim Lukas Holch and
5
+ Constantin Steppa (ai4iacts), originally developed for hexagonal convolution
6
+ and pooling on PyTorch in the context of Imaging Atmospheric Cherenkov
7
+ Telescope analysis with H.E.S.S.
8
+
9
+ This package extends the original `HexBase` sub-kernel decomposition with two
10
+ new features that have no equivalent in upstream HexagDLy:
11
+
12
+ - `share_neighbors`: ties the weights of a hexagonal kernel by ring (ring 0 =
13
+ center, ring *r* = the 6*r* cells at hex-distance *r*), instead of giving
14
+ every cell its own weight.
15
+ - `depth_padding` (`Conv3d` only): `"same"` zero-pads the depth/time axis so
16
+ the temporal kernel is centred and output depth equals input depth, instead of
17
+ HexagDLy's `"valid"`-only behaviour.
18
+
19
+ This work was developed as part of the SST-1M Collaboration / HEPIA TDSCAN
20
+ triggering project.
21
+
22
+ ## License
23
+
24
+ Both the original HexagDLy code and this fork are distributed under the MIT
25
+ license; see [LICENSE](LICENSE). The original copyright notice
26
+ (Copyright (c) 2018 ai4iacts) is preserved alongside the copyright notice for
27
+ this fork, as required by the MIT license.
28
+
29
+ ## Citing
30
+
31
+ If you use this package, please cite the original HexagDLy paper:
32
+
33
+ ```bibtex
34
+ @article{hexagdly_paper,
35
+ title = "HexagDLy—Processing hexagonally sampled data with CNNs in PyTorch",
36
+ author = "Constantin Steppa and Tim L. Holch",
37
+ journal = "SoftwareX",
38
+ volume = "9",
39
+ pages = "193 - 198",
40
+ year = "2019",
41
+ issn = "2352-7110",
42
+ doi = "https://doi.org/10.1016/j.softx.2019.02.010",
43
+ url = "https://www.sciencedirect.com/science/article/pii/S2352711018302723",
44
+ }
45
+ ```