CIRCT 24.0.0git
Loading...
Searching...
No Matches
channel_arbiter.py
Go to the documentation of this file.
1# Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
2# See https://llvm.org/LICENSE.txt for license information.
3# SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
4
5# Hardware for the ChannelArbiter cosim integration tests. Builds two
6# independent DUTs under a single top:
7#
8# * ChannelArbiterTest ("arbiter_test"): NUM_INPUTS host-driven `from_host`
9# UInt(32) channels multiplexed into a single `to_host` channel. The host
10# checks that every value is delivered exactly once (no loss / duplication).
11# Instantiated twice: a power-of-two ("balanced") count and a
12# non-power-of-two ("unbalanced") count.
13#
14# * ChannelArbiterListTest ("list_test"): two in-hardware list-window
15# producers (distinct `src` ids and message lengths) contend for the
16# arbiter. A hardware checker verifies that, once a message is granted, its
17# flits arrive contiguously (no interleaving) up to the flit tagged `last`,
18# and reports (sticky-error, src, count) per completed message on a
19# `to_host` channel.
20#
21# * ChannelArbiterTokenTest ("token_test"): TOKEN_NUM_INPUTS in-hardware
22# zero-width (`i0`) token producers, each emitting exactly TOKENS_PER_INPUT
23# tokens, contend for the arbiter. A hardware counter reports a running
24# 1-based ordinal per delivered token, so the host can confirm every token
25# is delivered exactly once (conservation) and every producer is served
26# (no starvation) -- exercising the credit-counter-as-buffer i0 path.
27
28import sys
29
30from pycde import (AppID, Clock, Input, Module, Output, Reset, System,
31 generator)
32from pycde.constructs import Counter, Mux, Wire
33from pycde.types import Bits, Channel, List as ListType, StructType, UInt, Window
34from pycde import esi
35
36from esiaccel.bsp import get_bsp
37from esiaccel.components import ChannelArbiter
38
39NUM_INPUTS = 4 # power-of-two ("balanced") input count.
40ODD_NUM_INPUTS = 3 # non-power-of-two ("unbalanced") input count.
41PIPE_NUM_INPUTS = 6 # input count for the pipelined-mux-tree variant.
42TOKEN_NUM_INPUTS = 5 # input count for the zero-width (i0) token variant.
43# Sustained-throughput probe: enough inputs that a per-sweep scheduler bubble
44# would be clearly visible (it would cap throughput at n/(n+1) == 0.8) while
45# keeping simulation time short.
46THROUGHPUT_NUM_INPUTS = 4
47THROUGHPUT_WINDOW = 1000 # measurement window, in cycles.
48TOKENS_PER_INPUT = 8 # tokens each producer emits in the token test.
49# A fan-in wide enough to exercise large-N arbitration. The BSP instantiates
50# arbiters with ~31 inputs (one per host-memory write client), a regime none of
51# the small counts above reach.
52WIDE_NUM_INPUTS = 13
53# Above `channel_arbiter._WIDE_FANIN_THRESHOLD`; all counts above are below it.
54WIDE_FANIN_NUM_INPUTS = 17
55
56# A list-window payload: a struct with a `src` tag and a variable-length list.
57# `Window.default_of` adds a per-flit `last` field to the lowered frame struct,
58# which is exactly what `ChannelArbiter` uses to keep messages contiguous.
59ListInto = StructType({'src': UInt(8), 'items': ListType(UInt(16))})
60Flit = Window.default_of(ListInto)
61FlitLowered = Flit.lowered_type # struct<src: ui8, items: ui16, last: i1>
62
63
64def HostMux(num_inputs: int,
65 mux_pipeline_levels=None,
66 pipelined_scheduler=False,
67 wide_fanin=None,
68 throttle=False):
69 """A host-driven single-flit multiplexer: `num_inputs` `from_host` UInt(32)
70 channels muxed into a single `to_host` channel. `mux_pipeline_levels` pipelines
71 the selection mux tree; `pipelined_scheduler` selects the decoupled
72 grant-queue arbitration. `throttle` drains the output one beat in four
73 through a minimal output FIFO, so the credit counter runs down to zero."""
74
75 class HostMux(Module):
76 clk = Clock()
77 rst = Reset()
78
79 @generator
80 def build(ports):
81 ins = [
82 esi.ChannelService.from_host(AppID(f"in_{i}"), UInt(32))
83 for i in range(num_inputs)
84 ]
85 out = ChannelArbiter(ins,
86 ports.clk,
87 ports.rst,
88 mux_pipeline_levels=mux_pipeline_levels,
89 pipelined_scheduler=pipelined_scheduler,
90 wide_fanin=wide_fanin,
91 output_fifo_depth=2 if throttle else None,
92 telemetry=False)
93 if throttle:
94 phase = Counter(2)(clk=ports.clk,
95 rst=ports.rst,
96 clear=Bits(1)(0),
97 increment=Bits(1)(1)).out
98 en = phase == UInt(2)(0)
99 ready = Wire(Bits(1))
100 data, valid = out.unwrap(ready)
101 out, out_ready = Channel(UInt(32)).wrap(data, valid & en)
102 ready.assign(out_ready & en)
103 esi.ChannelService.to_host(AppID("out"), out)
104
105 HostMux.__name__ = (f"HostMux_{num_inputs}_{mux_pipeline_levels}_"
106 f"{pipelined_scheduler}_{wide_fanin}_{throttle}")
107 return HostMux
108
109
110def ListProducer(src_id: int, length: int):
111 """A module which continuously emits back-to-back list messages of `length`
112 flits tagged with `src_id`. `items` counts 0..length-1 and `last` is set on
113 the final flit. Always valid, so two of these contend for the arbiter."""
114
115 class ListProducer(Module):
116 clk = Clock()
117 rst = Reset()
118 out = Output(Channel(Flit))
119
120 @generator
121 def build(ports):
122 i = Wire(UInt(8))
123 last = i == UInt(8)(length - 1)
124 st = FlitLowered({
125 'src': UInt(8)(src_id),
126 'items': i.as_uint(16),
127 'last': last,
128 })
129 chan, ready = Channel(Flit).wrap(Flit.wrap(st), Bits(1)(1))
130 ports.out = chan
131 # valid is constant 1, so a transaction happens whenever `ready`.
132 nxt = Mux(last, (i + UInt(8)(1)).as_uint(8), UInt(8)(0))
133 i.assign(nxt.reg(ports.clk, ports.rst, ce=ready, rst_value=0))
134
135 ListProducer.__name__ = f"ListProducer_src{src_id}_len{length}"
136 return ListProducer
137
138
139class ListChecker(Module):
140 """Consumes the muxed list-window stream and verifies message contiguity.
141
142 Emits one UInt(32) report per completed message:
143 bit 24 : sticky interleave-error flag (should stay 0)
144 bits 23:16 : src of the completed message
145 bits 15:0 : running completed-message count
146 """
147
148 clk = Clock()
149 rst = Reset()
150 in_ = Input(Channel(Flit))
151 report = Output(Channel(UInt(32)))
152
153 @generator
154 def build(ports):
155 active = Wire(Bits(1))
156 err = Wire(Bits(1))
157
158 in_ready = Wire(Bits(1))
159 win, valid = ports.in_.unwrap(in_ready)
160 st = win.unwrap()
161 src = st['src']
162 last = st['last']
163
164 # A message boundary: the current flit completes a message.
165 is_completing = valid & last
166 # Emit a report exactly when a message completes; back-pressure the input
167 # on that flit until the report is accepted.
168 report_valid = is_completing
169
170 # src of the in-progress message, latched at its first flit.
171 start_any = valid & in_ready & ~active
172 cur_src = src.reg(ports.clk, ports.rst, ce=start_any, rst_value=0)
173
174 # Interleave error: a flit whose src differs from the owner mid-message.
175 interleave_err = (valid & in_ready) & active & (src != cur_src)
176 err.assign((err | interleave_err).reg(ports.clk, ports.rst, rst_value=0))
177
178 # `active` tracks whether we are mid-message (past the first flit, before
179 # `last`).
180 xact = valid & in_ready
181 begin_multi = xact & ~active & ~last
182 end_msg = xact & last
183 active_next = Mux(end_msg, Mux(begin_multi, active, Bits(1)(1)), Bits(1)(0))
184 active.assign(active_next.reg(ports.clk, ports.rst, rst_value=0))
185
186 # Completed-message counter.
187 msg_count = Counter(16)(clk=ports.clk,
188 rst=ports.rst,
189 clear=Bits(1)(0),
190 increment=end_msg)
191
192 report_data = ((err.as_uint(32) * UInt(32)(0x1000000)).as_uint(32) +
193 (cur_src.as_uint(32) * UInt(32)(0x10000)).as_uint(32) +
194 msg_count.out.as_uint(32)).as_uint(32)
195 report_chan, report_ready = Channel(UInt(32)).wrap(report_data,
196 report_valid)
197 ports.report = report_chan
198
199 # Accept every non-completing flit; on a completing flit, only accept when
200 # the report is accepted so no completion is dropped.
201 in_ready.assign(Mux(is_completing, Bits(1)(1), report_ready))
202
203
204def ChannelArbiterListTestMod(pipelined_scheduler: bool, wide_fanin=None):
205 """Contending list producers -> arbiter -> contiguity checker. Message
206 atomicity is the property most at risk from any arbitration change, so it is
207 covered for both arbitration modes."""
208
209 class ChannelArbiterListTest(Module):
210 clk = Clock()
211 rst = Reset()
212
213 @generator
214 def build(ports):
215 # More producers than the grant-queue depth, so the scheduled variant
216 # exercises a queue that actually fills.
217 prods = [
218 ListProducer(src, 3 + (src % 3))(clk=ports.clk, rst=ports.rst)
219 for src in range(1, 7)
220 ] if pipelined_scheduler else [
221 ListProducer(1, 3)(clk=ports.clk, rst=ports.rst),
222 ListProducer(2, 4)(clk=ports.clk, rst=ports.rst),
223 ]
224 muxed = ChannelArbiter([p.out for p in prods],
225 ports.clk,
226 ports.rst,
227 pipelined_scheduler=pipelined_scheduler,
228 wide_fanin=wide_fanin,
229 telemetry=False)
230 chk = ListChecker(clk=ports.clk, rst=ports.rst, in_=muxed)
231 esi.ChannelService.to_host(AppID("report"), chk.report)
232
233 ChannelArbiterListTest.__name__ = (
234 f"ChannelArbiterListTest_{pipelined_scheduler}_{wide_fanin}")
235 return ChannelArbiterListTest
236
237
238def TokenProducer(count: int):
239 """Emits exactly `count` zero-width (`i0`) tokens then idles: `valid` stays
240 high until `count` tokens have been accepted. There is no payload -- only the
241 valid/ready handshake carries information."""
242
243 class TokenProducer(Module):
244 clk = Clock()
245 rst = Reset()
246 out = Output(Channel(Bits(0)))
247
248 @generator
249 def build(ports):
250 xact = Wire(Bits(1))
251 sent = Counter(16)(clk=ports.clk,
252 rst=ports.rst,
253 clear=Bits(1)(0),
254 increment=xact)
255 valid = sent.out < UInt(16)(count)
256 chan, ready = Channel(Bits(0)).wrap(Bits(0)(0), valid)
257 ports.out = chan
258 xact.assign(valid & ready)
259
260 TokenProducer.__name__ = f"TokenProducer_{count}"
261 return TokenProducer
262
263
264class TokenChecker(Module):
265 """Consumes the muxed zero-width token stream and emits one UInt(32) report
266 per delivered token carrying its 1-based ordinal (1, 2, 3, ...). Backpressure
267 from the report channel is fed to the arbiter, exercising its credit buffer.
268 """
269
270 clk = Clock()
271 rst = Reset()
272 in_ = Input(Channel(Bits(0)))
273 report = Output(Channel(UInt(32)))
274
275 @generator
276 def build(ports):
277 in_ready = Wire(Bits(1))
278 _tok, valid = ports.in_.unwrap(in_ready) # zero-width payload: ignore data.
279
280 # Running count of delivered tokens; this token's ordinal is count + 1.
281 count = Counter(32)(clk=ports.clk,
282 rst=ports.rst,
283 clear=Bits(1)(0),
284 increment=valid & in_ready)
285 report_data = (count.out + UInt(32)(1)).as_uint(32)
286 report_chan, report_ready = Channel(UInt(32)).wrap(report_data, valid)
287 ports.report = report_chan
288 # Accept a token exactly when its report is consumed by the host.
289 in_ready.assign(report_ready)
290
291
293 """`TOKEN_NUM_INPUTS` bounded zero-width token producers -> arbiter -> token
294 counter. Each producer emits exactly `TOKENS_PER_INPUT` tokens, so the host
295 must see exactly TOKEN_NUM_INPUTS * TOKENS_PER_INPUT tokens: no loss or
296 duplication (conservation), and every producer served (no starvation, since
297 the total can only be reached if each producer's tokens all get through)."""
298
299 clk = Clock()
300 rst = Reset()
301
302 @generator
303 def build(ports):
304 prods = [
305 TokenProducer(TOKENS_PER_INPUT)(clk=ports.clk, rst=ports.rst)
306 for _ in range(TOKEN_NUM_INPUTS)
307 ]
308 muxed = ChannelArbiter([p.out for p in prods],
309 ports.clk,
310 ports.rst,
311 telemetry=False)
312 chk = TokenChecker(clk=ports.clk, rst=ports.rst, in_=muxed)
313 esi.ChannelService.to_host(AppID("token_report"), chk.report)
314
315
317 """Never idles: `valid` is tied high, so the only thing limiting the
318 arbiter's delivery rate is the arbiter itself."""
319
320 class AlwaysValidProducer(Module):
321 clk = Clock()
322 rst = Reset()
323 out = Output(Channel(UInt(32)))
324
325 @generator
326 def build(ports):
327 chan, _ready = Channel(UInt(32)).wrap(UInt(32)(tag), Bits(1)(1))
328 ports.out = chan
329
330 AlwaysValidProducer.__name__ = f"AlwaysValidProducer_{tag}"
331 return AlwaysValidProducer
332
333
334class ThroughputProbe(Module):
335 """Drains the arbiter at full rate (`ready` tied high, so the host can never
336 backpressure it) and counts delivered beats over a fixed cycle window. Holds
337 the tally on its report channel once the window closes.
338
339 Draining in hardware is the point: the host-driven tests are rate-limited by
340 the cosim DPI, so they cannot observe sustained throughput at all."""
341
342 clk = Clock()
343 rst = Reset()
344 in_ = Input(Channel(UInt(32)))
345 report = Output(Channel(UInt(32)))
346
347 @generator
348 def build(ports):
349 _data, valid = ports.in_.unwrap(Bits(1)(1)) # never backpressure.
350 cycles = Counter(32)(clk=ports.clk,
351 rst=ports.rst,
352 clear=Bits(1)(0),
353 increment=Bits(1)(1))
354 running = cycles.out < UInt(32)(THROUGHPUT_WINDOW)
355 beats = Counter(32)(clk=ports.clk,
356 rst=ports.rst,
357 clear=Bits(1)(0),
358 increment=valid & running)
359 # Report only once the window has closed; the value is then stable, so the
360 # host can read it whenever it gets around to it.
361 chan, _ready = Channel(UInt(32)).wrap(beats.out, ~running)
362 ports.report = chan
363
364
365def ChannelArbiterThroughputTestMod(pipelined_scheduler: bool):
366 """`THROUGHPUT_NUM_INPUTS` never-idle producers -> arbiter -> throughput
367 probe. Pins the arbiter's sustained delivery rate: a scheduler which needs a
368 refill/turnaround cycle between grants shows up here as a throughput well
369 below one beat per cycle, while correctness tests stay green."""
370
371 class ChannelArbiterThroughputTest(Module):
372 clk = Clock()
373 rst = Reset()
374
375 @generator
376 def build(ports):
377 prods = [
378 AlwaysValidProducer(i)(clk=ports.clk, rst=ports.rst)
379 for i in range(THROUGHPUT_NUM_INPUTS)
380 ]
381 muxed = ChannelArbiter([p.out for p in prods],
382 ports.clk,
383 ports.rst,
384 pipelined_scheduler=pipelined_scheduler,
385 telemetry=False)
386 probe = ThroughputProbe(clk=ports.clk, rst=ports.rst, in_=muxed)
387 esi.ChannelService.to_host(AppID("throughput_report"), probe.report)
388
389 ChannelArbiterThroughputTest.__name__ = (
390 f"ChannelArbiterThroughputTest_{pipelined_scheduler}")
391 return ChannelArbiterThroughputTest
392
393
394class Top(Module):
395 clk = Clock()
396 rst = Reset()
397
398 @generator
399 def construct(ports):
400 HostMux(NUM_INPUTS)(clk=ports.clk,
401 rst=ports.rst,
402 appid=AppID("arbiter_test"))
403 HostMux(ODD_NUM_INPUTS)(clk=ports.clk,
404 rst=ports.rst,
405 appid=AppID("arbiter_test_odd"))
406 HostMux(PIPE_NUM_INPUTS,
407 mux_pipeline_levels=1)(clk=ports.clk,
408 rst=ports.rst,
409 appid=AppID("arbiter_test_pipe"))
410 HostMux(WIDE_NUM_INPUTS)(clk=ports.clk,
411 rst=ports.rst,
412 appid=AppID("arbiter_test_wide"))
413 HostMux(WIDE_NUM_INPUTS,
414 pipelined_scheduler=True)(clk=ports.clk,
415 rst=ports.rst,
416 appid=AppID("arbiter_test_sched"))
417 HostMux(ODD_NUM_INPUTS,
418 pipelined_scheduler=True)(clk=ports.clk,
419 rst=ports.rst,
420 appid=AppID("arbiter_test_sched_odd"))
421 HostMux(WIDE_FANIN_NUM_INPUTS,
422 throttle=True)(clk=ports.clk,
423 rst=ports.rst,
424 appid=AppID("arbiter_test_widefanin"))
425 HostMux(WIDE_FANIN_NUM_INPUTS, pipelined_scheduler=True,
426 throttle=True)(clk=ports.clk,
427 rst=ports.rst,
428 appid=AppID("arbiter_test_widefanin_sched"))
429 HostMux(NUM_INPUTS, wide_fanin=True)(clk=ports.clk,
430 rst=ports.rst,
431 appid=AppID("arbiter_test_forced_on"))
432 HostMux(WIDE_FANIN_NUM_INPUTS,
433 wide_fanin=False)(clk=ports.clk,
434 rst=ports.rst,
435 appid=AppID("arbiter_test_forced_off"))
436 ChannelArbiterListTestMod(False)(clk=ports.clk,
437 rst=ports.rst,
438 appid=AppID("list_test"))
439 ChannelArbiterListTestMod(True)(clk=ports.clk,
440 rst=ports.rst,
441 appid=AppID("list_test_sched"))
442 # Multi-flit messages through the `wide_fanin` `msg_end` path.
443 ChannelArbiterListTestMod(False, True)(clk=ports.clk,
444 rst=ports.rst,
445 appid=AppID("list_test_widefanin"))
447 True)(clk=ports.clk,
448 rst=ports.rst,
449 appid=AppID("list_test_widefanin_sched"))
450 ChannelArbiterTokenTest(clk=ports.clk,
451 rst=ports.rst,
452 appid=AppID("token_test"))
453 ChannelArbiterThroughputTestMod(False)(clk=ports.clk,
454 rst=ports.rst,
455 appid=AppID("throughput_test"))
456 ChannelArbiterThroughputTestMod(True)(clk=ports.clk,
457 rst=ports.rst,
458 appid=AppID("throughput_test_sched"))
459
460
461if __name__ == "__main__":
462 bsp = get_bsp(sys.argv[2] if len(sys.argv) > 2 else None)
463 s = System(bsp(Top), name="ChannelArbiterTest", output_directory=sys.argv[1])
464 s.compile()
465 s.package()
return wrap(CMemoryType::get(unwrap(ctx), baseType, numElements))
ListProducer(int src_id, int length)
ChannelArbiterThroughputTestMod(bool pipelined_scheduler)
ChannelArbiterListTestMod(bool pipelined_scheduler, wide_fanin=None)
AlwaysValidProducer(int tag)
TokenProducer(int count)
HostMux(int num_inputs, mux_pipeline_levels=None, pipelined_scheduler=False, wide_fanin=None, throttle=False)