espdlx 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.
espdlx/__init__.py ADDED
@@ -0,0 +1,7 @@
1
+ """espdlx: esp-dl friendly NN blocks for ESP32-S3 (PyTorch backend)."""
2
+
3
+ __version__ = "0.1.0"
4
+
5
+ from .graph import Model
6
+
7
+ __all__ = ["Model", "__version__"]
espdlx/convert.py ADDED
@@ -0,0 +1,314 @@
1
+ """ONNX export + esp-dl quantization for espdlx models.
2
+
3
+ This is the shipped version of the pipeline previously duplicated across the
4
+ experiment scripts (``flower_to_espdl.py``, ``stl10_to_espdl.py``,
5
+ ``micro_to_espdl.py`` — not shipped)::
6
+
7
+ from espdlx.convert import convert
8
+
9
+ report = convert(model, example_input, calib_loader, "model.espdl")
10
+
11
+ Steps: ``torch.onnx.export`` -> esp-dl friendly rewrite
12
+ (``Relu + Min(6)`` -> ``Clip``, ``Reshape`` -> ``Flatten``) ->
13
+ ``espdl_quantize_onnx`` (ESP-PPQ) -> ``.espdl`` + report.
14
+
15
+ Requires the ``convert`` extra (``pip install espdlx[convert]``).
16
+ ``onnx`` / ``esp_ppq`` are imported lazily so ``import espdlx`` stays light.
17
+ """
18
+
19
+ from __future__ import annotations
20
+
21
+ import json
22
+ from pathlib import Path
23
+ from typing import Any
24
+
25
+ __all__ = [
26
+ "convert",
27
+ "export_onnx",
28
+ "make_espdl_friendly",
29
+ "quantize_onnx",
30
+ ]
31
+
32
+ DEFAULT_TARGET = "esp32s3"
33
+ DEFAULT_QUANT_TYPE = "w8a8"
34
+ DEFAULT_OPSET = 13
35
+ DEFAULT_CALIB_STEPS = 16
36
+
37
+
38
+ def _require_onnx() -> Any:
39
+ try:
40
+ import onnx
41
+ except ImportError as exc:
42
+ raise ImportError(
43
+ "espdlx.convert needs the 'convert' extra "
44
+ "(pip install espdlx[convert]): missing package 'onnx'"
45
+ ) from exc
46
+ return onnx
47
+
48
+
49
+ def _require_espdl_quantize() -> Any:
50
+ try:
51
+ from esp_ppq.api import espdl_quantize_onnx
52
+ except ImportError as exc:
53
+ raise ImportError(
54
+ "espdlx.convert needs the 'convert' extra "
55
+ "(pip install espdlx[convert]): missing package 'esp-ppq'"
56
+ ) from exc
57
+ return espdl_quantize_onnx
58
+
59
+
60
+ def export_onnx(
61
+ model: Any,
62
+ example_input: Any,
63
+ onnx_path: str | Path,
64
+ *,
65
+ input_name: str = "input",
66
+ output_name: str = "output",
67
+ opset: int = DEFAULT_OPSET,
68
+ ) -> Path:
69
+ """Export ``model`` to ONNX (classic exporter, fixed opset).
70
+
71
+ Returns the resolved ``onnx_path``. The model is left untouched
72
+ (callers typically ``model.eval()`` first — see :func:`convert`).
73
+ """
74
+ import torch
75
+
76
+ onnx_path = Path(onnx_path)
77
+ onnx_path.parent.mkdir(parents=True, exist_ok=True)
78
+ torch.onnx.export(
79
+ model,
80
+ example_input,
81
+ str(onnx_path),
82
+ input_names=[input_name],
83
+ output_names=[output_name],
84
+ opset_version=opset,
85
+ do_constant_folding=True,
86
+ )
87
+ return onnx_path
88
+
89
+
90
+ def make_espdl_friendly(model: Any) -> Any:
91
+ """Rewrite ONNX exporter artifacts to ops esp-dl supports, in place.
92
+
93
+ - ``Relu + Min(6)`` -> ``Clip(0, 6)`` (esp-dl has Clip, no Min)
94
+ - ``Reshape`` (static) -> ``Flatten`` (old ppq executor chokes on Reshape)
95
+
96
+ Takes and returns the ``onnx.ModelProto``. Runs ``onnx.checker`` first
97
+ (import) and last (result) so failures surface here, not in quantization.
98
+ """
99
+ onnx = _require_onnx()
100
+ onnx.checker.check_model(model)
101
+
102
+ inits = {i.name: onnx.numpy_helper.to_array(i) for i in model.graph.initializer}
103
+ six = [n for n, v in inits.items() if v.shape == () and float(v) == 6.0]
104
+ if six:
105
+ six_name = six[0]
106
+ else:
107
+ import numpy as _np
108
+
109
+ six_name = "clip_max_6"
110
+ model.graph.initializer.append(
111
+ onnx.numpy_helper.from_array(_np.array(6.0, dtype=_np.float32), six_name)
112
+ )
113
+ zero_name = "clip_min_0"
114
+ if zero_name not in inits:
115
+ import numpy as _np
116
+
117
+ model.graph.initializer.append(
118
+ onnx.numpy_helper.from_array(_np.array(0.0, dtype=_np.float32), zero_name)
119
+ )
120
+
121
+ by_out = {}
122
+ for n in model.graph.node:
123
+ for o in n.output:
124
+ by_out[o] = n
125
+ relu_outs = set()
126
+ for n in model.graph.node:
127
+ if n.op_type == "Min" and six_name in n.input:
128
+ key = n.input[0] if n.input[1] == six_name else n.input[1]
129
+ relu = by_out.get(key, None)
130
+ if relu is not None and relu.op_type == "Relu" and key == relu.output[0]:
131
+ relu_outs.add(key)
132
+
133
+ new_nodes = []
134
+ for n in model.graph.node:
135
+ if n.op_type == "Relu" and n.output[0] in relu_outs:
136
+ continue # fused into the Clip below
137
+ if n.op_type == "Min" and six_name in n.input:
138
+ key = n.input[0] if n.input[1] == six_name else n.input[1]
139
+ relu = by_out.get(key, None)
140
+ if relu is not None and key in relu_outs:
141
+ new_nodes.append(
142
+ onnx.helper.make_node(
143
+ "Clip",
144
+ inputs=[relu.input[0], zero_name, six_name],
145
+ outputs=[n.output[0]],
146
+ name=(n.name + "_clip") if n.name else "",
147
+ )
148
+ )
149
+ continue
150
+ if n.op_type == "Reshape":
151
+ flat = onnx.helper.make_node(
152
+ "Flatten",
153
+ inputs=[n.input[0]],
154
+ outputs=[n.output[0]],
155
+ name=(n.name + "_flat") if n.name else "",
156
+ )
157
+ flat.attribute.append(onnx.helper.make_attribute("axis", 1))
158
+ new_nodes.append(flat)
159
+ continue
160
+ new_nodes.append(n)
161
+ del model.graph.node[:]
162
+ model.graph.node.extend(new_nodes)
163
+
164
+ used = {i for n in model.graph.node for i in n.input}
165
+ for i in list(model.graph.initializer):
166
+ if i.name == "val_1" and i.name not in used:
167
+ model.graph.initializer.remove(i)
168
+ onnx.checker.check_model(model)
169
+ return model
170
+
171
+
172
+ def _default_collate(batch: Any) -> Any:
173
+ """Default PPQ collate: ``DataLoader(TensorDataset(x))`` yields ``[x]``."""
174
+ if isinstance(batch, (list, tuple)) and len(batch) == 1:
175
+ batch = batch[0]
176
+ to_cpu = getattr(batch, "to", None)
177
+ return to_cpu("cpu") if callable(to_cpu) else batch
178
+
179
+
180
+ def quantize_onnx(
181
+ onnx_path: str | Path,
182
+ espdl_path: str | Path,
183
+ calib_loader: Any,
184
+ input_shape: list[int] | tuple[int, ...],
185
+ *,
186
+ target: str = DEFAULT_TARGET,
187
+ quant_type: str = DEFAULT_QUANT_TYPE,
188
+ calib_steps: int = DEFAULT_CALIB_STEPS,
189
+ collate_fn: Any = None,
190
+ device: str = "cpu",
191
+ error_report: bool = False,
192
+ export_test_values: bool = False,
193
+ verbose: int = 0,
194
+ ) -> Any:
195
+ """Quantize a (friendly) ONNX model via ESP-PPQ; returns the PPQ graph."""
196
+ espdl_quantize_onnx = _require_espdl_quantize()
197
+ espdl_path = Path(espdl_path)
198
+ espdl_path.parent.mkdir(parents=True, exist_ok=True)
199
+ return espdl_quantize_onnx(
200
+ onnx_import_file=str(onnx_path),
201
+ espdl_export_file=str(espdl_path),
202
+ calib_dataloader=calib_loader,
203
+ calib_steps=calib_steps,
204
+ input_shape=list(input_shape),
205
+ target=target,
206
+ quant_type=quant_type,
207
+ collate_fn=collate_fn or _default_collate,
208
+ device=device,
209
+ error_report=error_report,
210
+ export_test_values=export_test_values,
211
+ verbose=verbose,
212
+ )
213
+
214
+
215
+ def _input_quant_params(graph: Any, input_name: str) -> tuple[Any, Any]:
216
+ """Best-effort input scale / zero-point from a quantized PPQ graph."""
217
+ try:
218
+ g_in = graph.inputs[input_name]
219
+ cfg = g_in.dest_op_configs[0]
220
+ sc, off = cfg.scale, cfg.offset
221
+ import torch
222
+
223
+ sc = float(sc.reshape(-1)[0]) if torch.is_tensor(sc) else float(sc)
224
+ off = int(off.reshape(-1)[0]) if torch.is_tensor(off) else int(off)
225
+ return sc, off
226
+ except Exception:
227
+ return None, None
228
+
229
+
230
+ def convert(
231
+ model: Any,
232
+ example_input: Any,
233
+ calib_loader: Any,
234
+ espdl_path: str | Path,
235
+ *,
236
+ onnx_path: str | Path | None = None,
237
+ report_path: str | Path | None = "auto",
238
+ input_shape: list[int] | tuple[int, ...] | None = None,
239
+ input_name: str = "input",
240
+ output_name: str = "output",
241
+ opset: int = DEFAULT_OPSET,
242
+ target: str = DEFAULT_TARGET,
243
+ quant_type: str = DEFAULT_QUANT_TYPE,
244
+ calib_steps: int = DEFAULT_CALIB_STEPS,
245
+ collate_fn: Any = None,
246
+ device: str = "cpu",
247
+ verbose: int = 0,
248
+ ) -> dict:
249
+ """Export ``model`` and quantize to ``.espdl``; return a JSON-able report.
250
+
251
+ - ``example_input``: dummy tensor for ``torch.onnx.export``
252
+ (also derives ``input_shape`` when omitted).
253
+ - ``calib_loader``: calibration batches (e.g. ``DataLoader(TensorDataset)``).
254
+ - ``espdl_path``: where to write the ``.espdl`` file.
255
+ - ``onnx_path``: defaults to ``espdl_path`` with a ``.onnx`` suffix.
256
+ - ``report_path``: ``"auto"`` writes ``espdl_report.json`` next to the
257
+ ``.espdl`` file (scratch convention); pass a path to override,
258
+ ``None`` to skip writing (the report dict is still returned).
259
+ """
260
+ onnx = _require_onnx()
261
+ espdl_path = Path(espdl_path)
262
+ if onnx_path is None:
263
+ onnx_path = espdl_path.with_suffix(".onnx")
264
+ else:
265
+ onnx_path = Path(onnx_path)
266
+ if input_shape is None:
267
+ input_shape = [int(d) for d in example_input.shape]
268
+
269
+ model.eval()
270
+ export_onnx(
271
+ model,
272
+ example_input,
273
+ onnx_path,
274
+ input_name=input_name,
275
+ output_name=output_name,
276
+ opset=opset,
277
+ )
278
+ friendly = onnx.load(str(onnx_path))
279
+ make_espdl_friendly(friendly)
280
+ onnx.save(friendly, str(onnx_path))
281
+
282
+ graph = quantize_onnx(
283
+ onnx_path,
284
+ espdl_path,
285
+ calib_loader,
286
+ input_shape,
287
+ target=target,
288
+ quant_type=quant_type,
289
+ calib_steps=calib_steps,
290
+ collate_fn=collate_fn,
291
+ device=device,
292
+ verbose=verbose,
293
+ )
294
+ scale, zero_point = _input_quant_params(graph, input_name)
295
+ report = {
296
+ "onnx_path": str(onnx_path),
297
+ "espdl_path": str(espdl_path),
298
+ "espdl_bytes": espdl_path.stat().st_size,
299
+ "opset": opset,
300
+ "target": target,
301
+ "quant_type": quant_type,
302
+ "calib_steps": calib_steps,
303
+ "input_name": input_name,
304
+ "input_shape": list(input_shape),
305
+ "input_scale": scale,
306
+ "input_zero_point": zero_point,
307
+ }
308
+ if report_path == "auto":
309
+ report_path = espdl_path.parent / "espdl_report.json"
310
+ if report_path is not None:
311
+ report_path = Path(report_path)
312
+ report_path.parent.mkdir(parents=True, exist_ok=True)
313
+ report_path.write_text(json.dumps(report, indent=2))
314
+ return report
espdlx/fomo.py ADDED
@@ -0,0 +1,180 @@
1
+ """EXPERIMENTAL FOMO helpers: the two things Edge Impulse does better.
2
+
3
+ Status: synthetic unit tests only (`tests/test_fomo.py`); best real-data
4
+ run so far reached P=0.03/R=0.28 — not a working detector yet.
5
+
6
+ Kept deliberately small (no parity chase):
7
+
8
+ - logits-only training (softmax at export/decode, never in the model),
9
+ - imbalance handling: object-weighted cross-entropy + classifier bias init
10
+ from the dataset foreground prior.
11
+
12
+ Targets are per-cell class maps ``(N, G, G)`` with ``0`` = background,
13
+ ``1`` = object centroid cell (see :func:`encode_centroids`).
14
+ """
15
+
16
+ from __future__ import annotations
17
+
18
+ import math
19
+
20
+ import torch
21
+ import torch.nn.functional as F
22
+
23
+ __all__ = [
24
+ "encode_centroids",
25
+ "encode_heatmap",
26
+ "weighted_loss",
27
+ "soft_loss",
28
+ "init_bias_from_prior",
29
+ "decode",
30
+ "decode_peaks",
31
+ ]
32
+
33
+
34
+ def encode_centroids(
35
+ boxes: list[tuple[float, float, float, float]],
36
+ grid: int,
37
+ ) -> torch.Tensor:
38
+ """Map normalized ``(cx, cy, w, h)`` boxes to a ``(G, G)`` class map.
39
+
40
+ The cell containing each centroid is marked ``1``; everything else stays
41
+ ``0``. Zero/negative-area boxes are skipped.
42
+ """
43
+ target = torch.zeros(grid, grid, dtype=torch.long)
44
+ for cx, cy, w, h in boxes:
45
+ if w <= 0 or h <= 0:
46
+ continue
47
+ j = min(max(int(math.floor(cx * grid)), 0), grid - 1)
48
+ i = min(max(int(math.floor(cy * grid)), 0), grid - 1)
49
+ target[i, j] = 1
50
+ return target
51
+
52
+
53
+ def weighted_loss(
54
+ logits: torch.Tensor,
55
+ targets: torch.Tensor,
56
+ object_weight: float = 100.0,
57
+ ) -> torch.Tensor:
58
+ """Weighted cross-entropy over ``(N, 2, G, G)`` logits.
59
+
60
+ Background weight is ``1.0``; the object cell weight defaults to ``100.0``
61
+ (Edge Impulse default) to counter the ~1:143 foreground/background ratio
62
+ on a 12x12 grid.
63
+ """
64
+ weight = torch.tensor([1.0, float(object_weight)], device=logits.device)
65
+ return F.cross_entropy(logits, targets.long(), weight=weight)
66
+
67
+
68
+ def encode_heatmap(
69
+ boxes: list[tuple[float, float, float, float]],
70
+ grid: int,
71
+ sigma: float = 1.0,
72
+ radius: int = 2,
73
+ ) -> torch.Tensor:
74
+ """Splat normalized ``(cx, cy, w, h)`` boxes into a ``(G, G)`` heatmap.
75
+
76
+ Each centroid stamps a Gaussian bump (peak 1 at the exact centroid,
77
+ ``exp(-d^2 / 2*sigma^2)`` in cell units, ``radius`` cells wide);
78
+ overlapping bumps merge by max. Zero/negative-area boxes are skipped.
79
+ Pairs with :func:`soft_loss` / :func:`decode_peaks`.
80
+ """
81
+ heat = torch.zeros(grid, grid)
82
+ for cx, cy, w, h in boxes:
83
+ if w <= 0 or h <= 0:
84
+ continue
85
+ gx, gy = cx * grid, cy * grid
86
+ jc, ic = int(math.floor(gx)), int(math.floor(gy))
87
+ for i in range(max(ic - radius, 0), min(ic + radius, grid - 1) + 1):
88
+ for j in range(max(jc - radius, 0), min(jc + radius, grid - 1) + 1):
89
+ d2 = (gx - j) ** 2 + (gy - i) ** 2
90
+ v = math.exp(-d2 / (2 * sigma * sigma))
91
+ if v > heat[i, j]:
92
+ heat[i, j] = v
93
+ return heat
94
+
95
+
96
+ def soft_loss(
97
+ logits: torch.Tensor,
98
+ targets: torch.Tensor,
99
+ heatmaps: torch.Tensor,
100
+ object_weight: float = 20.0,
101
+ gamma: float = 4.0,
102
+ ) -> torch.Tensor:
103
+ """Localization-tolerant loss over ``(N, 2, G, G)`` logits.
104
+
105
+ - positives (centroid cells from :func:`encode_centroids`): NLL toward
106
+ the object logit, scaled by ``object_weight``;
107
+ - negatives: NLL toward background, scaled by ``(1 - heat)^gamma`` with
108
+ ``heat`` from :func:`encode_heatmap` — an adjacent cell (heat ~0.6)
109
+ pays ~2% of a far cell's penalty, so "roughly right" barely hurts.
110
+ Normalized by the positive-cell count (min 1), so the scale is stable
111
+ across images with different object counts.
112
+ """
113
+ logp = F.log_softmax(logits, dim=1)
114
+ pos = targets.long() == 1
115
+ n_pos = max(int(pos.sum().item()), 1)
116
+ loss_pos = -logp[:, 1][pos].sum() * float(object_weight)
117
+ w = (1.0 - heatmaps.float()).pow(float(gamma))
118
+ loss_neg = -(w * logp[:, 0])[~pos].sum()
119
+ return (loss_pos + loss_neg) / n_pos
120
+
121
+
122
+ def init_bias_from_prior(model: torch.nn.Module, obj_prior: float) -> None:
123
+ """Init the final head bias so the model starts at the dataset prior.
124
+
125
+ With a 2-logit softmax head, ``bias = [0, log(p / (1 - p))]`` makes the
126
+ initial object probability equal the foreground cell fraction ``p``
127
+ (Edge Impulse's ``set_classifier_biases_from_dataset`` idea, reduced to
128
+ one number). Finds the last ``espdlx`` conv layer in ``model.layers``.
129
+ """
130
+ from .layers.conv import Conv2d as EspConv2d
131
+
132
+ p = min(max(float(obj_prior), 1e-6), 1.0 - 1e-6)
133
+ head = None
134
+ for lyr in getattr(model, "layers", []):
135
+ if isinstance(lyr, EspConv2d) and lyr.out_channels == 2:
136
+ head = lyr
137
+ if head is None:
138
+ raise ValueError("init_bias_from_prior: no 2-channel head conv found")
139
+ with torch.no_grad():
140
+ bias = head.conv.bias
141
+ assert bias is not None and bias.numel() == 2
142
+ bias.zero_()
143
+ bias[1] = math.log(p / (1.0 - p))
144
+
145
+
146
+ @torch.no_grad()
147
+ def decode(
148
+ logits: torch.Tensor,
149
+ thresh: float = 0.5,
150
+ ) -> list[list[tuple[int, int, float]]]:
151
+ """Logits ``(N, 2, G, G)`` -> per image ``[(col, row, score)]`` centroids.
152
+
153
+ Softmax is applied here (or at export), never during training. ``score``
154
+ is the object-channel probability.
155
+ """
156
+ prob = torch.softmax(logits, dim=1)[:, 1]
157
+ out: list[list[tuple[int, int, float]]] = []
158
+ for n in range(logits.shape[0]):
159
+ idx = (prob[n] > thresh).nonzero(as_tuple=False)
160
+ out.append([(int(j), int(i), float(prob[n, i, j])) for i, j in idx.tolist()])
161
+ return out
162
+
163
+
164
+ @torch.no_grad()
165
+ def decode_peaks(
166
+ logits: torch.Tensor,
167
+ thresh: float = 0.5,
168
+ ) -> list[list[tuple[int, int, float]]]:
169
+ """Peak-NMS decode: 3x3 max-pool keeps local maxima above ``thresh``.
170
+
171
+ Pairs with :func:`soft_loss` — the loss lets bumps spread over adjacent
172
+ cells, this collapses each bump back to one centroid.
173
+ """
174
+ prob = torch.softmax(logits, dim=1)[:, 1]
175
+ peaks = F.max_pool2d(prob[:, None], 3, stride=1, padding=1)[:, 0] == prob
176
+ out: list[list[tuple[int, int, float]]] = []
177
+ for n in range(logits.shape[0]):
178
+ idx = ((prob[n] > thresh) & peaks[n]).nonzero(as_tuple=False)
179
+ out.append([(int(j), int(i), float(prob[n, i, j])) for i, j in idx.tolist()])
180
+ return out
espdlx/graph.py ADDED
@@ -0,0 +1,37 @@
1
+ """The `Model` container: an ordered stack of espdlx layers.
2
+
3
+ Deploy with :func:`espdlx.convert.convert` (``torch.onnx.export`` ->
4
+ ESP-PPQ ``espdl_quantize_onnx`` -> ``.espdl``).
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from torch import nn
10
+
11
+ from .layers.base import Layer
12
+
13
+
14
+ class Model(nn.Module):
15
+ def __init__(self, layers: list, name: str = "model") -> None:
16
+ super().__init__()
17
+ for lyr in layers:
18
+ if not isinstance(lyr, Layer):
19
+ raise TypeError(
20
+ f"{type(lyr).__name__} is not an espdlx layer; use espdlx.layers"
21
+ )
22
+ self.name = name
23
+ self.layers = nn.ModuleList(layers)
24
+
25
+ def forward(self, x):
26
+ for lyr in self.layers:
27
+ x = lyr(x)
28
+ return x
29
+
30
+ def validate_shapes(self, input_shape: tuple) -> tuple:
31
+ shape = tuple(input_shape)
32
+ for lyr in self.layers:
33
+ shape = lyr.validate_shapes(shape)
34
+ return shape
35
+
36
+ def __repr__(self) -> str:
37
+ return f"Model(name={self.name!r}, layers={len(self.layers)})"
@@ -0,0 +1,27 @@
1
+ """espdlx constrained layer blocks (torch nn.Module wrappers)."""
2
+
3
+ from .activations import BatchNorm2d, HardSwish, ReLU, ReLU6, Sigmoid, Softmax
4
+ from .base import Layer
5
+ from .conv import Conv2d, DepthwiseConv2d
6
+ from .linear import Linear
7
+ from .math import Add, Flatten, Mean, Mul
8
+ from .pool import AvgPool2d, MaxPool2d
9
+
10
+ __all__ = [
11
+ "Layer",
12
+ "Conv2d",
13
+ "DepthwiseConv2d",
14
+ "Linear",
15
+ "MaxPool2d",
16
+ "AvgPool2d",
17
+ "ReLU",
18
+ "ReLU6",
19
+ "HardSwish",
20
+ "Sigmoid",
21
+ "Softmax",
22
+ "BatchNorm2d",
23
+ "Add",
24
+ "Mul",
25
+ "Mean",
26
+ "Flatten",
27
+ ]
@@ -0,0 +1,70 @@
1
+ """Activation blocks (plain torch semantics, esp-dl deployable)."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import torch.nn.functional as F
6
+ from torch import nn
7
+
8
+ from .base import Layer
9
+
10
+
11
+ class ReLU(Layer):
12
+ def forward(self, x):
13
+ return F.relu(x)
14
+
15
+ def validate_shapes(self, in_shape):
16
+ return in_shape
17
+
18
+
19
+ class ReLU6(Layer):
20
+ def forward(self, x):
21
+ return F.relu6(x)
22
+
23
+ def validate_shapes(self, in_shape):
24
+ return in_shape
25
+
26
+
27
+ class HardSwish(Layer):
28
+ def forward(self, x):
29
+ return F.hardswish(x)
30
+
31
+ def validate_shapes(self, in_shape):
32
+ return in_shape
33
+
34
+
35
+ class Sigmoid(Layer):
36
+ def forward(self, x):
37
+ return F.sigmoid(x)
38
+
39
+ def validate_shapes(self, in_shape):
40
+ return in_shape
41
+
42
+
43
+ class Softmax(Layer):
44
+ def __init__(self, dim: int = -1):
45
+ super().__init__()
46
+ self.dim = dim
47
+
48
+ def forward(self, x):
49
+ return F.softmax(x, dim=self.dim)
50
+
51
+ def validate_shapes(self, in_shape):
52
+ return in_shape
53
+
54
+
55
+ class BatchNorm2d(Layer):
56
+ """Batch normalization (standard torch semantics, exports to ONNX)."""
57
+
58
+ def __init__(self, num_features: int, eps: float = 1e-5, momentum: float = 0.1):
59
+ super().__init__()
60
+ self.num_features = int(num_features)
61
+ self.bn = nn.BatchNorm2d(self.num_features, eps=eps, momentum=momentum)
62
+
63
+ def forward(self, x):
64
+ return self.bn(x)
65
+
66
+ def validate_shapes(self, in_shape):
67
+ return in_shape
68
+
69
+
70
+ __all__ = ["ReLU", "ReLU6", "HardSwish", "Sigmoid", "Softmax", "BatchNorm2d"]
espdlx/layers/base.py ADDED
@@ -0,0 +1,11 @@
1
+ """Shared base for espdlx constrained blocks."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import torch
6
+ from torch import nn
7
+
8
+
9
+ class Layer(nn.Module):
10
+ def validate_shapes(self, in_shape: tuple) -> tuple:
11
+ raise NotImplementedError