Skip to content

Commit 0a0839b

Browse files
committed
Corrected 1x1 kernel edge case, updated testbenches for additional layer filenames
1 parent 4515800 commit 0a0839b

7 files changed

Lines changed: 48 additions & 22 deletions

File tree

rtl/blocks/conv_layer.sv

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -57,8 +57,8 @@ module conv_layer #(
5757

5858
,localparam int unsigned PaddedWidth = LineWidthPx + (2 * Padding)
5959
,localparam int unsigned PaddedHeight = LineCountPx + (2 * Padding)
60-
,localparam int XBits = (LineWidthPx <= 1) ? 1 : $clog2(PaddedWidth + 1)
61-
,localparam int YBits = (LineCountPx <= 1) ? 1 : $clog2(PaddedHeight + 1)
60+
,localparam int XBits = $clog2(PaddedWidth + 1)
61+
,localparam int YBits = $clog2(PaddedHeight + 1)
6262

6363
,localparam int unsigned WeightIndex = InChannels * KernelArea * WeightBits
6464
,parameter logic signed [OutChannels*InChannels*KernelWidth*KernelWidth*WeightBits-1:0] Weights = '0

sim/functional_models/cnn_model.py

Lines changed: 15 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -203,15 +203,28 @@ def step(self, x: List[int]):
203203

204204
# Return flattened results (e.g. [class_id1, class_id2, ...])
205205
# For classifier, res item is (class_id, logits)
206-
return [item[0] for item in current_burst]
206+
final_res = [item[0] for item in current_burst]
207+
if final_res:
208+
print("DEBUG CNNModel step returned:", final_res)
209+
return final_res
207210

208211
def consume(self):
209212
"""Standard ModelRunner interface for top-level DUT."""
210213
from util.bitwise import unpack_terms
211214
packed = int(self._dut.data_i.value.integer)
212215
is_unsigned = self._InBits > 2
213216
raw_val = unpack_terms(packed, self._InBits, self._InChannels, signed=not is_unsigned)
214-
return self.step(raw_val)
217+
218+
if not hasattr(self, '_consume_count'):
219+
self._consume_count = 0
220+
self._consume_count += 1
221+
222+
res = self.step(raw_val)
223+
if res is not None:
224+
print(f"DEBUG: CNNModel.consume() produced {res} on call {self._consume_count}")
225+
elif self._consume_count == 784:
226+
print(f"DEBUG: CNNModel.consume() produced None on call 784!!!")
227+
return res
215228

216229
def produce(self, expected):
217230
"""Standard ModelRunner interface for top-level DUT."""

sim/integration_testing/cnn/test_cnn.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -221,7 +221,12 @@ def _write_vh(path: str, name: str, total_bits: int, raw_val: int) -> None:
221221
# 4. Run simulator
222222
os.environ["CHECK_INFERENCE"] = os.environ.get("CHECK_INFERENCE", "1")
223223
os.environ["SAMPLE_IDX"] = os.environ.get("SAMPLE_IDX", "10")
224-
224+
225+
# Remove the single-layer fallback keys added by inject_weights_and_biases;
226+
# cnn.sv only has FileName_0…FileName_N, not plain FileName / FileName_hi.
227+
params.pop("FileName", None)
228+
params.pop("FileName_hi", None)
229+
225230
runner(
226231
simulator=simulator,
227232
timescale="1ps/1ps",

sim/integration_testing/cnn_framed/tb_cnn_framed.sv

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -7,10 +7,10 @@ module cnn_framed #(
77
parameter int unsigned BusBits = 8
88
,parameter FileName_0 = "nn/data/roms/hex/zeros.hex"
99
,parameter FileName_1 = "nn/data/roms/hex/zeros.hex"
10-
,parameter FileName_1_hi = "nn/data/roms/hex/zeros.hex"
1110
,parameter FileName_2 = "nn/data/roms/hex/zeros.hex"
12-
,parameter FileName_2_hi = "nn/data/roms/hex/zeros.hex"
1311
,parameter FileName_3 = "nn/data/roms/hex/zeros.hex"
12+
,parameter FileName_4 = "nn/data/roms/hex/zeros.hex"
13+
,parameter FileName_5 = "nn/data/roms/hex/zeros.hex"
1414

1515
,parameter int unsigned WidthIn = 320
1616
,parameter int unsigned HeightIn = 240
@@ -61,10 +61,10 @@ module cnn_framed #(
6161
.BusBits (BusBits)
6262
,.FileName_0 (FileName_0)
6363
,.FileName_1 (FileName_1)
64-
,.FileName_1_hi(FileName_1_hi)
6564
,.FileName_2 (FileName_2)
66-
,.FileName_2_hi(FileName_2_hi)
6765
,.FileName_3 (FileName_3)
66+
,.FileName_4 (FileName_4)
67+
,.FileName_5 (FileName_5)
6868
) dut (
6969
.clk_i (clk_i)
7070
,.rst_i (rst_i)

sim/integration_testing/cnn_framed/test_cnn_framed.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -252,6 +252,9 @@ def _write_vh(path: str, name: str, total_bits: int, raw_val: int) -> None:
252252
os.environ["CHECK_INFERENCE"] = os.environ.get("CHECK_INFERENCE", "1")
253253
os.environ["SAMPLE_IDX"] = os.environ.get("SAMPLE_IDX", "10")
254254

255+
params.pop("FileName", None)
256+
params.pop("FileName_hi", None)
257+
255258
runner(
256259
simulator=simulator,
257260
timescale="1ps/1ps",

sim/integration_testing/cnn_uart/test_cnn_uart.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -253,6 +253,9 @@ def _write_vh(path: str, name: str, total_bits: int, raw_val: int) -> None:
253253

254254
os.environ["SAMPLE_IDX"] = os.environ.get("SAMPLE_IDX", "10")
255255

256+
params.pop("FileName", None)
257+
params.pop("FileName_hi", None)
258+
256259
runner(
257260
simulator=simulator,
258261
timescale="1ps/1ps",

sim/util/components.py

Lines changed: 15 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -100,22 +100,24 @@ async def _run_input(self):
100100
if self._rst_i.value.is_resolvable and int(self._rst_i.value) == 1:
101101
continue
102102

103-
# Evaluate Input Handshake
104103
v = self._valid_i.value == 1 if self._valid_i is not None else True
105104
r = self._ready_o.value == 1 if self._ready_o is not None else True
106105

107-
if not (v and r):
108-
continue
109-
110-
expected = self._model.consume()
111-
112-
# Restore backward compatibility for single items vs lists
113-
if expected is not None:
114-
if isinstance(expected, (list, tuple)):
115-
for item in expected:
116-
self._events.put(item)
117-
else:
118-
self._events.put(expected)
106+
if not hasattr(self, '_run_input_calls'):
107+
self._run_input_calls = 0
108+
109+
if v and r:
110+
self._run_input_calls += 1
111+
if self._run_input_calls % 100 == 0 or self._run_input_calls == 1 or self._run_input_calls >= 780:
112+
print(f"DEBUG: ModelRunner v and r is True on call {self._run_input_calls}")
113+
114+
expected = self._model.consume()
115+
if expected is not None:
116+
if isinstance(expected, (list, tuple)):
117+
for item in expected:
118+
self._events.put(item)
119+
else:
120+
self._events.put(expected)
119121

120122
async def _run_output(self):
121123
from decimal import Decimal

0 commit comments

Comments
 (0)