CIRCT 24.0.0git
Loading...
Searching...
No Matches
common.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
5from __future__ import annotations
6from math import ceil
7
8from pycde.common import Clock, Input, InputChannel, Output, OutputChannel, Reset
9from pycde.constructs import (AssignableSignal, ControlReg, Counter, Mux,
10 NamedWire, Reg, Wire)
11from pycde import esi
12from pycde.module import Module, generator, modparams
13from pycde.signals import BitsSignal, ChannelSignal, StructSignal
14from pycde.support import clog2
15from pycde.system import System
16from pycde.types import (Array, Bits, Bundle, BundledChannel, Channel,
17 ChannelDirection, StructType, Type, UInt, Window)
18
19from ..components import ChannelArbiter, MaxOutstandingLimiter
20
21from typing import Callable, Dict, List, Optional, Tuple
22import typing
23
24MagicNumber = 0x207D98E5_E5100E51 # random + ESI__ESI
25VersionNumber = 0 # Version 0: format subject to change
26
27IndirectionMagicNumber = 0x312bf0cc_E5100E51 # random + ESI__ESI
28IndirectionVersionNumber = 0 # Version 0: format subject to change
29
30# Magic value which, when written by the host to header slot 7, requests a
31# design reset. Keep in sync with 'ResetMagicNumber' in the runtime
32# (cpp/include/esi/Accelerator.h). This magic number guards against "write
33# spraying" which other devices have been know to do on boot.
34ResetMagicNumber = 0x00000E510000B007
35# Number of cycles to wait after a reset is requested before asserting it. This
36# gives in-flight transactions time to drain.
37ResetCycles = 8192
38
39MMIOWordBytes = esi.MMIODataType.width // 8
40
41
42class ESI_Manifest_ROM(Module):
43 """Module which will be created later by CIRCT which will contain the
44 compressed manifest."""
45
46 module_name = "__ESI_Manifest_ROM"
47
48 clk = Clock()
49 address = Input(Bits(29))
50 # Data is two cycles delayed after address changes.
51 data = Output(Bits(64))
52
53
55 """Wrap the manifest ROM with ESI bundle."""
56
57 clk = Clock()
58 read = Input(esi.MMIO.read.type)
59
60 @generator
61 def build(self):
62 data, data_valid = Wire(Bits(64)), Wire(Bits(1))
63 data_chan, data_ready = Channel(Bits(64)).wrap(data, data_valid)
64 address_chan = self.read.unpack(data=data_chan)['offset']
65 address, address_valid = address_chan.unwrap(data_ready)
66 address_words = address.as_bits(32)[3:] # Lop off the lower three bits.
67
68 rom = ESI_Manifest_ROM(clk=self.clk, address=address_words)
69 data.assign(rom.data)
70 data_valid.assign(address_valid.reg(self.clk, name="data_valid", cycles=2))
71
72
73@modparams
74def HeaderMMIO(manifest_loc: int) -> Module:
75
76 class HeaderMMIO(Module):
77 """Construct the ESI header MMIO adhering to the MMIO layout specified in
78 the ChannelMMIO service implementation."""
79
80 clk = Clock()
81 rst = Reset()
82 read = Input(esi.MMIO.read_write.type)
83 # Asserted for one cycle when the host writes the reset magic number to
84 # header slot 7. Propagates up to the BSP which performs the actual reset.
85 reset_request = Output(Bits(1))
86
87 @generator
88 def build(ports):
89 clk = ports.clk
90 rst = ports.rst
91 data_chan_wire = Wire(Channel(esi.MMIODataType))
92 input_bundles = ports.read.unpack(data=data_chan_wire)
93 cmd_chan = input_bundles['cmd']
94
95 # Two-stage half-throughput pipeline: stage 1 captures the incoming
96 # command, stage 2 holds the looked-up response. Each stage carries its
97 # own occupancy bit.
98 cmd_ready = Wire(Bits(1))
99 s1_to_s2_xact = Wire(Bits(1))
100 cmd_raw, cmd_valid = cmd_chan.unwrap(cmd_ready)
101
102 # Stage 1: command capture register and occupancy bit.
103 s1_load = cmd_valid & cmd_ready
104 cmd = cmd_raw.reg(clk, rst, ce=s1_load, name="cmd")
105 s1_valid = ControlReg(clk,
106 rst,
107 asserts=[s1_load],
108 resets=[s1_to_s2_xact],
109 name="s1_valid")
110 # Accept a new command when stage 1 is empty.
111 cmd_ready.assign(~s1_valid)
112
113 address_words = cmd.offset.as_bits()[3:] # Lop off the lower three bits.
114 slot = address_words[:3]
115
116 cycles = Counter(64)(clk=ports.clk,
117 rst=ports.rst,
118 clear=Bits(1)(0),
119 increment=Bits(1)(1),
120 instance_name="cycle_counter")
121
122 # Layout the header as an array.
123 core_freq = System.current().core_freq
124 if core_freq is None:
125 core_freq = 0
126 header = Array(Bits(64), 8)([
127 0, # Generally a good idea to not use address 0.
128 MagicNumber, # ESI magic number.
129 VersionNumber, # ESI version number.
130 manifest_loc, # Absolute address of the manifest ROM.
131 0, # Reserved for future use.
132 cycles.out.as_bits(), # Cycle counter.
133 core_freq, # Core frequency, if known.
134 0, # Slot 7: write the reset magic number here to request a reset.
135 ])
136 header.name = "header"
137
138 # Stage 2: registered response value and its occupancy bit.
139 s2_valid = Wire(Bits(1))
140 data_chan_ready = Wire(Bits(1))
141 s2_xact = s2_valid & data_chan_ready
142 # Stage 1 advances into stage 2 only when stage 2 is empty.
143 s1_to_s2_xact.assign(s1_valid & ~s2_valid)
144
145 header_out = header[slot].reg(clk=clk,
146 rst=rst,
147 ce=s1_to_s2_xact,
148 name="header_out")
149 s2_valid.assign(
150 ControlReg(clk,
151 rst,
152 asserts=[s1_to_s2_xact],
153 resets=[s2_xact],
154 name="header_out_valid"))
155 # Wrap the response.
156 data_chan, data_chan_ready_sig = Channel(esi.MMIODataType).wrap(
157 header_out, s2_valid)
158 data_chan_wire.assign(data_chan)
159 data_chan_ready.assign(data_chan_ready_sig)
160
161 # Detect a write of the reset magic number to slot 7. Register the request
162 # so it is a clean one-cycle pulse, asserted as the command advances into
163 # the response stage. 'DesignResetController' latches it, so a single-cycle
164 # pulse is sufficient to trigger the reset.
165 reset_detect = (cmd.write & (slot == Bits(3)(7)) &
166 (cmd.data == Bits(64)(ResetMagicNumber)))
167 ports.reset_request = reset_detect & s1_to_s2_xact
168
169 return HeaderMMIO
170
171
172@modparams
174 data_type: Type, num_outs: int,
175 next_sel_width: int) -> type["ChannelDemuxNImpl"]:
176 """N-way channel demultiplexer for valid/ready signaling. Contains
177 valid/ready registers on the output channels. The selection signal is now
178 embedded in the input channel payload as a struct {sel, data}. Input
179 signals ready when the selected output register is empty."""
180
181 assert num_outs >= 1, "num_outs must be at least 1."
182
183 class ChannelDemuxNImpl(Module):
184 clk = Clock()
185 rst = Reset()
186
187 # Input channel now carries selection along with data.
188 InPayloadType = StructType([
189 ("sel", Bits(clog2(num_outs))),
190 ("next_sel", Bits(next_sel_width)),
191 ("data", data_type),
192 ])
193 inp = Input(Channel(InPayloadType))
194 OutPayloadType = StructType([
195 ("next_sel", Bits(next_sel_width)),
196 ("data", data_type),
197 ])
198 # Outputs are channels of OutPayloadType, which includes both 'next_sel' and 'data' fields.
199 for i in range(num_outs):
200 locals()[f"output_{i}"] = Output(Channel(OutPayloadType))
201
202 @generator
203 def generate(ports) -> None:
204 # Half-stage demux: one register per output channel. Input is ready
205 # when the currently selected output register is empty (not valid).
206 clk = ports.clk
207 rst = ports.rst
208 sel_width = clog2(num_outs)
209
210 # Unwrap input with backpressure from selected output register.
211 input_ready = Wire(Bits(1), name="input_ready")
212 in_payload, in_valid = ports.inp.unwrap(input_ready)
213 in_sel = in_payload.sel
214 in_next_sel = in_payload.next_sel
215 in_data = in_payload.data
216
217 # Track per-output valid regs and build a purely combinational
218 # expression 'selected_valid_expr' = OR_i((sel==i)&valid_i). Avoid
219 # assigning to a Wire multiple times.
220 valid_regs: List[BitsSignal] = []
221 selected_valid_expr = Bits(1)(0)
222
223 for i in range(num_outs):
224 # Write when input transaction targets this output and output not holding data yet.
225 will_write = Wire(Bits(1), name=f"will_write_{i}")
226 write_cond = (in_valid & input_ready & (in_sel == Bits(sel_width)(i)))
227 will_write.assign(write_cond)
228
229 # Data and next_sel registers.
230 out_msg_reg = ChannelDemuxNImpl.OutPayloadType({
231 "next_sel": in_next_sel,
232 "data": in_data
233 }).reg(clk=clk, rst=rst, ce=will_write, name=f"out{i}_msg_reg")
234
235 # Valid register cleared on successful downstream consume.
236 consume = Wire(Bits(1), name=f"consume_{i}")
237 valid_reg = ControlReg(
238 clk=clk,
239 rst=rst,
240 asserts=[will_write],
241 resets=[consume],
242 name=f"out{i}_valid_reg",
243 )
244 valid_regs.append(valid_reg)
245
246 # Channel wrapper.
247 ch_sig, ch_ready = Channel(ChannelDemuxNImpl.OutPayloadType).wrap(
248 out_msg_reg, valid_reg)
249 setattr(ports, f"output_{i}", ch_sig)
250 consume.assign(valid_reg & ch_ready)
251
252 # Accumulate selected_valid expression.
253 selected_valid_expr = selected_valid_expr | (
254 (in_sel == Bits(sel_width)(i)) & valid_reg)
255
256 # Input ready only when selected output has no valid data latched.
257 input_ready.assign(selected_valid_expr ^ Bits(1)(1))
258
259 def get_out(self, index: int) -> ChannelSignal:
260 return getattr(self, f"output_{index}")
261
262 return ChannelDemuxNImpl
263
264
265@modparams
267 data_type: Type, num_outs: int,
268 branching_factor_log2: int) -> type["ChannelDemuxTree"]:
269 """Pipelined N-way channel demultiplexer for valid/ready signaling. This
270 implementation uses a tree structure of
271 ChannelDemuxN_HalfStage_ReadyBlocking modules to reduce fanout pressure.
272 Supports maximum half-throughput to save complexity and area.
273 """
274
275 root_sel_width = clog2(num_outs)
276 # Simplify algorithm by making sure num_outs is a power of two.
277 num_outs = 2**root_sel_width
278 sel_width = branching_factor_log2
279 fanout = 2**sel_width
280
281 class ChannelDemuxTree(Module):
282 clk = Clock()
283 rst = Reset()
284 # Input now embeds selection bits alongside data.
285 InPayloadType = StructType([
286 ("sel", Bits(clog2(num_outs))),
287 ("data", data_type),
288 ])
289 inp = Input(Channel(InPayloadType))
290
291 # Outputs (data only).
292 for i in range(num_outs):
293 locals()[f"output_{i}"] = Output(Channel(data_type))
294
295 @generator
296 def build(ports) -> None:
297 assert branching_factor_log2 > 0
298 if num_outs == 1:
299 # Strip selection bits and return single channel.
300 setattr(ports, "output_0", ports.inp.transform(lambda p: p.data))
301 return
302
303 def payload_type(sel_width: int, next_sel_width: int) -> Type:
304 return StructType([
305 ("sel", Bits(sel_width)),
306 ("next_sel", Bits(next_sel_width)),
307 ("data", data_type),
308 ])
309
310 def next_sel_width_calc(curr_sel_width) -> int:
311 return max(curr_sel_width - sel_width, 0)
312
313 def payload_next(curr_msg: StructSignal) -> StructSignal:
314 """Given current level payload, produce next level payload by
315 stripping off the top selection bits."""
316
317 next_sel_width = next_sel_width_calc(curr_msg.next_sel.type.width)
318 curr_sel_width = curr_msg.next_sel.type.width
319 new_sel_width = min(curr_sel_width, sel_width)
320 return payload_type(
321 new_sel_width,
322 next_sel_width,
323 )({
324 # Use the MSB bits of next_sel as the next level selection.
325 "sel": (curr_msg.next_sel[next_sel_width:]
326 if curr_sel_width > 0 else Bits(0)(0)),
327 "next_sel": (curr_msg.next_sel[:next_sel_width]
328 if next_sel_width > 0 else Bits(0)(0)),
329 "data": curr_msg.data,
330 })
331
332 current_channels: List[ChannelSignal] = [
333 ports.inp.transform(lambda m: payload_type(0, root_sel_width)({
334 "sel": Bits(0)(0),
335 "next_sel": m.sel,
336 "data": m.data,
337 }))
338 ]
339
340 curr_sel_width = root_sel_width
341 level = 0
342 while len(current_channels) < num_outs:
343 next_level: List[ChannelSignal] = []
344 level_num_outs = min(2**curr_sel_width, fanout)
345 for i, c in enumerate(current_channels):
347 data_type,
348 num_outs=level_num_outs,
349 next_sel_width=next_sel_width_calc(curr_sel_width),
350 )(
351 clk=ports.clk,
352 rst=ports.rst,
353 inp=c.transform(payload_next),
354 instance_name=f"demux_l{level}_i{i}",
355 )
356 for j in range(level_num_outs):
357 next_level.append(dmux.get_out(j))
358 current_channels = next_level
359 curr_sel_width -= sel_width
360 level += 1
361
362 for i in range(num_outs):
363 # Strip off next_sel bits for final output.
364 setattr(
365 ports,
366 f"output_{i}",
367 current_channels[i].transform(lambda p: p.data),
368 )
369
370 def get_out(self, index: int) -> ChannelSignal:
371 return getattr(self, f"output_{index}")
372
373 return ChannelDemuxTree
374
375
377 regions: Tuple[Tuple[int, Optional[int]],
378 ...]) -> type["MMIOPrefixRouterImpl"]:
379 """Build a pipelined address-prefix tree which routes MMIO commands.
380
381 Each ``(base, size)`` must describe a disjoint, power-of-two-aligned block.
382 The last region may pass ``size=None`` to claim every address at or above its
383 base; the manifest uses this since its size isn't known during generation.
384 Internal nodes test one address bit and register both outgoing paths. Leaves
385 replace the global address with the block-local low bits, so no subtractor is
386 needed (the open-ended region is the one exception).
387
388 For four 0x100-byte regions at 0x0/0x100/0x200/0x300 plus an open-ended
389 region at 0x400: the open-ended region is peeled off first, then each level
390 splits the remaining candidates in half, so depth is ``O(log N)`` rather than
391 one level per region::
392
393 cmd ──▶ addr[31:10]≠0 ──1──▶ @0x400 (open ended, offset -= 0x400)
394 │
395 0
396 â–¼
397 addr[9]
398 ╱ ╲
399 0 1
400 â–¼ â–¼
401 addr[8] addr[8]
402 ╱ ╲ ╱ ╲
403 0 1 0 1
404 â–¼ â–¼ â–¼ â–¼
405 @0x0 @0x100 @0x200 @0x300
406
407 The four leaves all take their client-local offset as ``addr[7:0]``.
408 """
409
410 assert len(regions) > 1, "MMIO routing requires at least two regions"
411
412 # An open-ended final region matches anything with a bit set above the sized
413 # regions, so it needs no size and can never overflow its window.
414 open_ended_base: Optional[int] = None
415 sized_regions = regions
416 if regions[-1][1] is None:
417 open_ended_base = regions[-1][0]
418 assert open_ended_base > 0 and \
419 open_ended_base & (open_ended_base - 1) == 0, \
420 "an open-ended MMIO region must start at a power-of-two address"
421 sized_regions = regions[:-1]
422
423 for base, size in sized_regions:
424 assert size > 0 and size & (size - 1) == 0, \
425 "MMIO region sizes must be powers of two"
426 assert base % size == 0, "MMIO regions must be aligned to their size"
427 assert all(
428 a_base + a_size <= b_base
429 for (a_base, a_size), (b_base, _) in zip(sized_regions, regions[1:])), \
430 "MMIO regions must be sorted and disjoint"
431
432 entries = [(idx, base, size) for idx, (base, size) in enumerate(sized_regions)
433 ]
434
435 def build_tree(node_entries):
436 if len(node_entries) == 1:
437 return node_entries[0][0]
438
439 lowest_fixed_bit = max(size.bit_length() - 1 for _, _, size in node_entries)
440 candidates = []
441 for bit in range(31, lowest_fixed_bit - 1, -1):
442 zeros = [entry for entry in node_entries if not (entry[1] >> bit) & 1]
443 if zeros and len(zeros) != len(node_entries):
444 candidates.append(
445 (abs(len(node_entries) - 2 * len(zeros)), -bit, bit, zeros))
446 assert candidates, "unable to distinguish disjoint MMIO regions"
447 _, _, bit, zeros = min(candidates)
448 zero_indices = {entry[0] for entry in zeros}
449 ones = [entry for entry in node_entries if entry[0] not in zero_indices]
450 return bit, build_tree(zeros), build_tree(ones)
451
452 tree = build_tree(entries)
453
454 class MMIOPrefixRouterImpl(Module):
455 clk = Clock()
456 rst = Reset()
457 inp = Input(Channel(esi.MMIOReadWriteCmdType))
458 for idx in range(len(regions)):
459 locals()[f"output_{idx}"] = Output(Channel(esi.MMIOReadWriteCmdType))
460
461 @generator
462 def build(ports):
463 """Route through one registered address-bit branch per tree level."""
464
465 def new_demux(command_channel: ChannelSignal, select, name: str):
466 Demux = ChannelDemuxN_HalfStage_ReadyBlocking(esi.MMIOReadWriteCmdType,
467 num_outs=2,
468 next_sel_width=0)
469 demux_input = command_channel.transform(
470 lambda cmd, _sel=select, _type=Demux.InPayloadType: _type({
471 "sel": _sel(cmd),
472 "next_sel": Bits(0)(0),
473 "data": cmd
474 }))
475 return Demux(clk=ports.clk,
476 rst=ports.rst,
477 inp=demux_input,
478 instance_name=name)
479
480 def route(node, command_channel: ChannelSignal, path: str):
481 if isinstance(node, int):
482 _, size = regions[node]
483 local_width = size.bit_length() - 1
484
485 def localize(cmd):
486 local_offset = cmd.offset.as_bits()[:local_width].pad_or_truncate(
487 32).as_uint()
488 return esi.MMIOReadWriteCmdType({
489 "write": cmd.write,
490 "offset": local_offset,
491 "data": cmd.data
492 })
493
494 setattr(ports, f"output_{node}", command_channel.transform(localize))
495 return
496
497 bit, zero_node, one_node = node
498 demux = new_demux(command_channel,
499 lambda cmd, _bit=bit: cmd.offset.as_bits()[_bit],
500 f"prefix_{path}_bit{bit}")
501 route(zero_node,
502 demux.get_out(0).transform(lambda p: p.data), path + "0")
503 route(one_node,
504 demux.get_out(1).transform(lambda p: p.data), path + "1")
505
506 client_channel = ports.inp
507 if open_ended_base is not None:
508 # Peel off the open-ended region first: any address bit above the sized
509 # regions selects it. Its base isn't a bit slice away from the local
510 # offset (the region is unbounded), so subtract it explicitly.
511 split_bit = open_ended_base.bit_length() - 1
512 split = new_demux(
513 ports.inp, lambda cmd: cmd.offset.as_bits()[split_bit:].or_reduce(),
514 f"prefix_above_bit{split_bit}")
515
516 def localize_open_ended(payload):
517 cmd = payload.data
518 local_offset = (cmd.offset - UInt(32)(open_ended_base)).as_uint(32)
519 return esi.MMIOReadWriteCmdType({
520 "write": cmd.write,
521 "offset": local_offset,
522 "data": cmd.data
523 })
524
525 setattr(ports, f"output_{len(regions) - 1}",
526 split.get_out(1).transform(localize_open_ended))
527 client_channel = split.get_out(0).transform(lambda p: p.data)
528
529 route(tree, client_channel, "")
530
531 def get_out(self, index: int) -> ChannelSignal:
532 return getattr(self, f"output_{index}")
533
534 return MMIOPrefixRouterImpl
535
536
537@modparams
538def DesignResetController(
539 delay_cycles: int) -> type["DesignResetControllerImpl"]:
540 """Counts `delay_cycles` clock cycles after a reset request is observed, then
541 asserts `design_reset` for one cycle. This module must be driven by the
542 *external* reset only (not the reset it generates) so that the countdown is
543 not disturbed by the reset it produces.
544
545 `reset_pending` is asserted from the moment a reset is requested until it
546 fires. It is intended to be used to quiesce the design (e.g. stop accepting
547 new transactions) so that nothing is in flight when the reset is asserted."""
548
549 if delay_cycles < 1:
550 raise ValueError("'delay_cycles' must be at least 1.")
551
552 counter_width = max(clog2(delay_cycles), 1)
553
554 class DesignResetControllerImpl(Module):
555 clk = Clock()
556 rst = Reset()
557 reset_request = Input(Bits(1))
558 design_reset = Output(Bits(1))
559 # High from the cycle a reset is requested until it fires. Use this to stop
560 # accepting new work so in-flight transactions can drain before the reset.
561 reset_pending = Output(Bits(1))
562
563 @generator
564 def build(ports):
565 fire = Wire(Bits(1))
566 # Latch that a reset has been requested until we fire the reset.
567 pending = ControlReg(clk=ports.clk,
568 rst=ports.rst,
569 asserts=[ports.reset_request],
570 resets=[fire],
571 name="reset_pending")
572 # Count cycles while a reset is pending.
573 count = Counter(counter_width)(clk=ports.clk,
574 rst=ports.rst,
575 clear=fire | ~pending,
576 increment=pending,
577 instance_name="reset_delay_counter")
578 fire.assign(pending &
579 (count.out == UInt(counter_width)(delay_cycles - 1)))
580 ports.design_reset = fire
581 ports.reset_pending = pending
582
583 return DesignResetControllerImpl
584
585
586class ChannelMMIO(esi.ServiceImplementation):
587 """MMIO service implementation with MMIO bundle interfaces. Should be
588 relatively easy to adapt to physical interfaces by wrapping the wires to
589 channels then bundles. Allows the implementation to be shared and (hopefully)
590 platform independent.
591
592 Whether or not to support unaligned accesses is up to the clients. The header
593 and manifest do not support unaligned accesses and throw away the lower three
594 bits.
595
596 Only allows one outstanding request at a time. This is enforced in hardware
597 by a `MaxOutstandingLimiter` on the command channel, which stalls incoming
598 commands until the previous response has been consumed. If a client fails to
599 return a response, the MMIO service will hang. TODO: add some kind of
600 timeout.
601
602 Implementation-defined MMIO layout:
603 - 0x0: 0 constant
604 - 0x8: Magic number (0x207D98E5_E5100E51)
605 - 0x12: ESI version number (0)
606 - 0x18: Location of the manifest ROM (absolute address)
607
608 - 0x100: Start of MMIO space for requests. Mapping is contained in the
609 manifest so can be dynamically queried.
610
611 - Directly above the last client allocation (rounded up to a power of two):
612 Start of the manifest ROM. It matches any address with a bit set above
613 the client space, so it is never truncated by its window.
614 - addr(Manifest ROM) + 0: Size of compressed manifest
615 - addr(Manifest ROM) + 8: Start of compressed manifest
616
617 This layout _should_ be pretty standard, but different BSPs may have various
618 different restrictions. Any BSP which uses this service implementation will
619 have this layout, possibly with an offset or address window.
620 """
621
622 clk = Clock()
623 rst = Input(Bits(1))
624
625 cmd = Input(esi.MMIO.read_write.type)
626
627 # Asserted for one cycle when the host requests a design reset via an MMIO
628 # write to the header. Propagates up to the BSP which performs the reset.
629 reset_request = Output(Bits(1))
630
631 # Amount of MMIO space a client gets when it doesn't request a size. Must be
632 # at least large enough for the header block, which uses eight 64-bit slots.
633 DefaultRegionSpace = 0x100
634
635 # Start at this address for assigning MMIO addresses to service requests.
636 initial_offset: int = DefaultRegionSpace
637
638 @generator
639 def generate(ports, bundles: esi._ServiceGeneratorBundles):
640 table, manifest_loc = ChannelMMIO.build_table(bundles)
641 ChannelMMIO.build_read(ports, manifest_loc, table)
642 return True
643
644 @staticmethod
646 bundles) -> Tuple[Dict[int, Tuple[Optional[int], AssignableSignal]], int]:
647 """Build a table of read and write addresses to BundleSignals."""
648
649 def align(value: int, alignment: int) -> int:
650 return (value + alignment - 1) // alignment * alignment
651
652 offset = ChannelMMIO.initial_offset
653 table: Dict[int, Tuple[Optional[int], AssignableSignal]] = {}
654 for bundle in bundles.to_client_reqs:
655 requested_size = bundle.options.get("size",
656 ChannelMMIO.DefaultRegionSpace)
657 if isinstance(requested_size,
658 bool) or not isinstance(requested_size, int):
659 raise ValueError("MMIO request option 'size' must be an integer")
660 if requested_size <= 0:
661 raise ValueError("MMIO request option 'size' must be positive")
662 # Round the allocation up to a power of two and align the base to it so
663 # that the address decode is a pure prefix match and the client-local
664 # offset is a bit slice.
665 size = 1 << (max(requested_size, MMIOWordBytes) - 1).bit_length()
666 offset = align(offset, size)
667 next_offset = offset + size
668 if next_offset >= 1 << 32:
669 raise ValueError("MMIO address allocation exceeds the 32-bit space")
670
671 if bundle.port == 'read':
672 table[offset] = size, bundle
673 bundle.add_record(details={
674 "offset": offset,
675 "size": size,
676 "type": "ro"
677 })
678 elif bundle.port == 'read_write':
679 table[offset] = size, bundle
680 bundle.add_record(details={
681 "offset": offset,
682 "size": size,
683 "type": "rw"
684 })
685 else:
686 raise ValueError(f"Unrecognized MMIO port name: {bundle.port}")
687 offset = next_offset
688
689 # The manifest sits directly above the clients and claims every address
690 # with a bit set above them. Its size isn't known until CIRCT builds the
691 # manifest ROM, so leaving the region open-ended keeps it from overflowing
692 # while letting the BAR be only as large as the manifest actually needs.
693 manifest_loc = 1 << max(offset - 1, 0).bit_length()
694 if manifest_loc >= 1 << 32:
695 raise ValueError(
696 "MMIO address allocation leaves no room for the manifest")
697 return table, manifest_loc
698
699 @staticmethod
700 def build_read(ports, manifest_loc: int,
701 table: Dict[int, Tuple[Optional[int], AssignableSignal]]):
702 """Builds the read side of the MMIO service."""
703
704 # Instantiate the header and manifest ROM. Fill in the read_table with
705 # bundle wires to be assigned identically to the other MMIO clients.
706 header_bundle_wire = Wire(esi.MMIO.read_write.type)
707 table[0] = ChannelMMIO.DefaultRegionSpace, header_bundle_wire
708 header = HeaderMMIO(manifest_loc)(clk=ports.clk,
709 rst=ports.rst,
710 read=header_bundle_wire)
711
712 mani_bundle_wire = Wire(esi.MMIO.read.type)
713 # 'None' size: the manifest claims every address above the clients.
714 table[manifest_loc] = None, mani_bundle_wire
715 ESI_Manifest_ROM_Wrapper(clk=ports.clk, read=mani_bundle_wire)
716
717 # Unpack the cmd bundle.
718 data_resp_channel = Wire(Channel(esi.MMIODataType))
719 counted_output = Wire(Channel(esi.MMIODataType))
720 cmd_channel = ports.cmd.unpack(data=counted_output)["cmd"]
721 counted_output.assign(data_resp_channel)
722
723 # Enforce the single-outstanding-transaction invariant in hardware: hold
724 # off accepting a new command until the response to the previous command
725 # has been consumed by the host. Snoop the response wire for the
726 # completion pulse.
727 resp_xact, _ = counted_output.snoop_xact()
728 cmd_limiter = MaxOutstandingLimiter(cmd_channel.type.inner_type,
729 max_outstanding=1)(
730 clk=ports.clk,
731 rst=ports.rst,
732 in_=cmd_channel,
733 complete=resp_xact,
734 instance_name="cmd_rate_limiter",
735 )
736 cmd_channel = cmd_limiter.out
737
738 sorted_table = sorted(table.items())
739 command_router = MMIOPrefixRouter(
740 tuple((base, size) for base, (size, _) in sorted_table))(
741 clk=ports.clk,
742 rst=ports.rst,
743 inp=cmd_channel,
744 instance_name="command_router",
745 )
746
747 client_data_channels = []
748 for idx, (_, (_, bundle_wire)) in enumerate(sorted_table):
749 client_cmd_channel = command_router.get_out(idx)
750 bundle_type = bundle_wire.type
751 if bundle_type == esi.MMIO.read.type:
752 offset = client_cmd_channel.transform(lambda cmd: cmd.offset)
753 bundle, bundle_froms = esi.MMIO.read.type.pack(offset=offset)
754 elif bundle_type == esi.MMIO.read_write.type:
755 bundle, bundle_froms = esi.MMIO.read_write.type.pack(
756 cmd=client_cmd_channel)
757 else:
758 assert False, "Unrecognized bundle type."
759 bundle_wire.assign(bundle)
760 client_data_channels.append(bundle_froms["data"])
761 # `cmd_rate_limiter` above caps the design at one outstanding MMIO command,
762 # and `client_cmd_demux` routes that one command to exactly one client. So
763 # provided each client only asserts its response `valid` in reply to a
764 # command it was given -- see `ChannelMergeOneValid` for why that second
765 # half matters, and note it is required by `ChannelMux` too -- at most one
766 # client response is ever valid, and arbitration is unnecessary.
767 resp_channel = esi.ChannelMergeOneValid(client_data_channels, ports.clk,
768 ports.rst)
769 data_resp_channel.assign(resp_channel)
770
771 # The header surfaces a reset request when the host writes the reset magic
772 # number to slot 7. Propagate it up to the caller (the BSP).
773 ports.reset_request = header.reset_request
774
775
776class TelemetryMMIO(esi.ServiceImplementation):
777 """An ESI service implementation which provides telemetry data through an MMIO
778 region. Each client request is assigned a register in the MMIO space. When a
779 read request is received for the assigned address, it gets routed to the
780 assigned client. When a write request is received, it is discarded. The
781 assignment table is stored in the manifest.
782
783 **REQUIREMENTS.** Both are needed to make the response merge safe:
784
785 1. The `MMIO` service implementation this connects to must not issue a read
786 command while a previous read's response is still outstanding.
787 2. Every telemetry client must assert its `data` channel's `valid` only in
788 response to a `get`. `Telemetry.report_signal` does this; a client which
789 holds `valid` high permanently -- legal ESI, and the natural way to
790 express an always-available counter -- does not.
791
792 Given both, each command is demuxed to exactly one client and at most one
793 client is offering a response at a time, so the responses can be merged with
794 `ChannelMergeOneValid` instead of arbitrated -- which keeps the response path
795 from building a combinational cone across every telemetry client.
796
797 Nothing here enforces either, and violating either loses responses. Note (2)
798 is not specific to this merge: `ChannelMux2` is fixed-priority, so under an
799 arbiter a permanently-valid client starves every client behind it."""
800
801 clk = Clock()
802 rst = Reset()
803
804 @generator
805 def generate(ports, bundles: esi._ServiceGeneratorBundles) -> bool:
806 if len(bundles.to_client_reqs) == 0:
807 # No clients to connect to, so we don't need to do anything.
808 return True
809
810 # Assign each telemetry client a register offset in MMIO space.
811 offset = 0
812 table: Dict[int, AssignableSignal] = {}
813 for bundle in bundles.to_client_reqs:
814 # Only support 'report' port for telemetry.
815 if bundle.port == 'report':
816 table[offset] = bundle
817 bundle.add_record(details={"offset": offset, "type": "mmio"})
818 offset += MMIOWordBytes
819 else:
820 raise ValueError(f"Unrecognized port name: {bundle.port}")
821
822 # Request exactly the space the register table occupies rather than relying
823 # on the MMIO service's default allocation.
824 mmio_cmd = esi.MMIO.read_write(esi.AppID("__telemetry_mmio"),
825 options={"size": offset})
826
827 # Unpack the cmd bundle.
828 data_resp_channel = Wire(Channel(esi.MMIODataType), "telemetry_data_resp")
829 counted_output = Wire(Channel(esi.MMIODataType), "telemetry_counted_output")
830 cmd_channel = mmio_cmd.unpack(data=counted_output)["cmd"]
831 counted_output.assign(data_resp_channel)
832
833 # Decode the address to select the client.
834 cmd_ready_wire = Wire(Bits(1), "telemetry_cmd_ready")
835 cmd, cmd_valid = cmd_channel.unwrap(cmd_ready_wire)
836 client_addr_chan, client_addr_ready = Channel(Bits(0)).wrap(
837 Bits(0)(0), cmd_valid)
838 cmd_ready_wire.assign(client_addr_ready)
839
840 # Build the demux/mux and assign the results of each appropriately.
841 read_clients_clog2 = clog2(len(table))
842 chan_sel = cmd.offset.as_bits()[3:read_clients_clog2 + 3]
843 client_cmd_channels = esi.ChannelDemux(
844 sel=chan_sel,
845 input=client_addr_chan,
846 num_outs=len(table),
847 instance_name="telemetry_client_cmd_demux")
848 client_data_channels = []
849 for (idx, offset) in enumerate(sorted(table.keys())):
850 bundle_wire = table[offset]
851 bundle_type = bundle_wire.type
852 # For telemetry, the client expects a 'get' channel and returns 'data'.
853 offset_chan = client_cmd_channels[idx]
854 bundle, bundle_froms = bundle_type.pack(get=offset_chan)
855
856 bundle_wire.assign(bundle)
857 client_data_channels.append(
858 bundle_froms["data"].transform(lambda m: m.as_bits(64)))
859 # The demux above routes each command to exactly one client, so -- given
860 # both requirements in this class' docs -- at most one client is offering a
861 # response at a time and no arbitration is needed to merge them.
862 resp_channel = esi.ChannelMergeOneValid(
863 client_data_channels,
864 ports.clk,
865 ports.rst,
866 instance_name="telemetry_resp_merge")
867 data_resp_channel.assign(resp_channel)
868 return True
869
870
871class MMIOIndirection(Module):
872 """Some platforms do not support MMIO space greater than a certain size (e.g.
873 Vitis 2022's limit is 4k). This module implements a level of indirection to
874 provide access to a full 32-bit address space.
875
876 MMIO addresses:
877 - 0x0: 0 constant
878 - 0x8: 64 bit ESI magic number for Indirect MMIO (0x312bf0cc_E5100E51)
879 - 0x10: Version number for Indirect MMIO (0)
880 - 0x18: Location of read/write in the virtual MMIO space.
881 - 0x20: A read from this location will initiate a read in the virtual MMIO
882 space specified by the address stored in 0x18 and return the result.
883 A write to this location will initiate a write into the virtual MMIO
884 space to the virtual address specified in 0x18.
885 """
886 clk = Clock()
887 rst = Reset()
888
889 upstream = Input(esi.MMIO.read_write.type)
890 downstream = Output(esi.MMIO.read_write.type)
891
892 @generator
893 def build(ports):
894 # This implementation assumes there is only one outstanding upstream MMIO
895 # transaction in flight at once. TODO: enforce this or make it more robust.
896
897 reg_bits = 8
898 location_reg = UInt(reg_bits)(0x18)
899 indirect_mmio_reg = UInt(reg_bits)(0x20)
900 virt_address = Wire(UInt(32))
901
902 # Set up the upstream MMIO interface. Capture last upstream command in a
903 # mailbox which never empties to give access to the last command for all
904 # time.
905 upstream_resp_chan_wire = Wire(Channel(esi.MMIODataType))
906 upstream_cmd_chan = ports.upstream.unpack(
907 data=upstream_resp_chan_wire)["cmd"]
908 _, _, upstream_cmd_data = upstream_cmd_chan.snoop()
909
910 # Set up a channel demux to separate the MMIO commands which get processed
911 # locally with ones which should be transformed and fowarded downstream.
912 phys_loc = upstream_cmd_data.offset.as_uint(reg_bits)
913 fwd_upstream = NamedWire(phys_loc == indirect_mmio_reg, "fwd_upstream")
914 local_reg_cmd_chan, downstream_cmd_channel = esi.ChannelDemux(
915 upstream_cmd_chan, fwd_upstream, 2, "upstream_demux")
916
917 # Set up the downstream MMIO interface.
918 downstream_cmd_channel = downstream_cmd_channel.transform(
919 lambda cmd: esi.MMIOReadWriteCmdType({
920 "write": cmd.write,
921 "offset": virt_address,
922 "data": cmd.data
923 }))
924 ports.downstream, froms = esi.MMIO.read_write.type.pack(
925 cmd=downstream_cmd_channel)
926 downstream_data_chan = froms["data"]
927
928 # Process local regs.
929 (local_reg_cmd_valid, local_reg_cmd_ready,
930 local_reg_cmd) = local_reg_cmd_chan.snoop()
931 write_virt_address = (local_reg_cmd_valid & local_reg_cmd_ready &
932 local_reg_cmd.write & (phys_loc == location_reg))
933 virt_address.assign(
934 local_reg_cmd.data.as_uint(32).reg(
935 name="virt_address",
936 clk=ports.clk,
937 ce=write_virt_address,
938 ))
939
940 # Build the pysical MMIO register space.
941 local_reg_resp_array = Array(Bits(64), 4)([
942 0x0, # 0x0
943 IndirectionMagicNumber, # 0x8
944 IndirectionVersionNumber, # 0x10
945 virt_address.as_bits(64), # 0x18
946 ])
947 local_reg_resp_chan = local_reg_cmd_chan.transform(
948 lambda cmd: local_reg_resp_array[cmd.offset.as_uint(2)])
949
950 # Mux together the local register responses and the downstream data to
951 # create the upstream response.
952 upstream_resp = esi.ChannelMux([local_reg_resp_chan, downstream_data_chan])
953 upstream_resp_chan_wire.assign(upstream_resp)
954
955
956@modparams
957def SliceReadGearbox(input_bitwidth: int,
958 output_bitwidth: int) -> type["SliceReadGearboxImpl"]:
959 """Narrow one engine word to a single-message client element no wider than the
960 word (``OUT <= IN``). The element sits in the word's low bits, so the datapath
961 is a slice; ``valid_bytes`` is unused (a single element is never a partial
962 word). Wider single elements use `ConcatReadGearbox`; packed list reads use
963 `DepackReadGearbox`/`ShiftReadGearbox`."""
964
965 if input_bitwidth <= 0 or input_bitwidth % 8 != 0:
966 raise ValueError("engine word width must be a positive multiple of 8 bits")
967 if not 0 < output_bitwidth <= input_bitwidth:
968 raise ValueError("SliceReadGearbox requires 0 < output <= input")
969
970 in_bytes = input_bitwidth // 8
971 vb_width = clog2(in_bytes)
972
973 class SliceReadGearboxImpl(Module):
974 clk = Clock()
975 rst = Reset()
976 in_ = InputChannel(
977 StructType([
978 ("tag", esi.HostMem.TagType),
979 ("data", Bits(input_bitwidth)),
980 ("valid_bytes", UInt(vb_width)),
981 ("last", Bits(1)),
982 ]))
983 out = OutputChannel(
984 StructType([
985 ("tag", esi.HostMem.TagType),
986 ("data", Bits(output_bitwidth)),
987 ("last", Bits(1)),
988 ]))
989
990 @generator
991 def build(ports):
992 up_ready = Wire(Bits(1), name="up_ready")
993 up, up_valid = ports.in_.unwrap(up_ready)
994 client_channel, client_ready = SliceReadGearboxImpl.out.type.wrap(
995 {
996 "tag": up.tag,
997 "data": up.data[:output_bitwidth],
998 "last": up.last,
999 }, up_valid)
1000 up_ready.assign(client_ready)
1001 ports.out = client_channel
1002
1003 return SliceReadGearboxImpl
1004
1005
1006@modparams
1007def ConcatReadGearbox(input_bitwidth: int,
1008 output_bitwidth: int) -> type["ConcatReadGearboxImpl"]:
1009 """Concatenate ``ceil(OUT/IN)`` consecutive engine words into one client
1010 element wider than the word (``OUT > IN``). Serves single-message reads (any
1011 ``OUT > IN``; the low ``OUT`` bits of the concatenation are the element) and
1012 contiguous list reads whose element is a whole number of output_bitwidth
1013 (``OUT % IN == 0``, so elements never straddle). ``valid_bytes`` is unused:
1014 such lists have no partial words and a single element is one flit. Straddling
1015 lists use `ShiftReadGearbox`."""
1016
1017 if input_bitwidth <= 0 or input_bitwidth % 8 != 0:
1018 raise ValueError("engine word width must be a positive multiple of 8 bits")
1019 if output_bitwidth <= input_bitwidth:
1020 raise ValueError("ConcatReadGearbox requires output > input")
1021
1022 in_bytes = input_bitwidth // 8
1023 vb_width = clog2(in_bytes)
1024
1025 class ConcatReadGearboxImpl(Module):
1026 clk = Clock()
1027 rst = Reset()
1028 in_ = InputChannel(
1029 StructType([
1030 ("tag", esi.HostMem.TagType),
1031 ("data", Bits(input_bitwidth)),
1032 ("valid_bytes", UInt(vb_width)),
1033 ("last", Bits(1)),
1034 ]))
1035 out = OutputChannel(
1036 StructType([
1037 ("tag", esi.HostMem.TagType),
1038 ("data", Bits(output_bitwidth)),
1039 ("last", Bits(1)),
1040 ]))
1041
1042 @generator
1043 def build(ports):
1044 ready_for_upstream = Wire(Bits(1), name="ready_for_upstream")
1045 # Register the input for fmax; the ESI channel buffer keeps the handshake
1046 # elastic.
1047 in_reg = ports.in_.buffer(ports.clk, ports.rst, stages=1)
1048 up, upstream_valid = in_reg.unwrap(ready_for_upstream)
1049 upstream_data = up.data
1050 upstream_last = up.last
1051 upstream_xact = ready_for_upstream & upstream_valid
1052
1053 # Registers accumulate `chunks` upstream words into one client element;
1054 # the output is their concatenation. For a list, elements stream back to
1055 # back and 'last' rides the final word of the burst's final element.
1056 chunks = ceil(output_bitwidth / input_bitwidth)
1057 counter_width = clog2(chunks)
1058 reg_ces = [Wire(Bits(1)) for _ in range(chunks)]
1059 regs = [
1060 upstream_data.reg(ports.clk,
1061 ports.rst,
1062 ce=reg_ces[idx],
1063 name=f"chunk_reg_{idx}") for idx in range(chunks)
1064 ]
1065 client_data_bits = BitsSignal.concat(reversed(regs))[:output_bitwidth]
1066
1067 # Pair-index counter: the word accepted this cycle is written to
1068 # chunk_reg[counter]. 'Counter' clears in preference to incrementing, so
1069 # mask the clear with the accept -- a consume and an accept on the same
1070 # cycle means the accepted word is chunk 0 of the *next* element, so the
1071 # index must land on 1, not 0. 'chunks' need not be a power of two, so
1072 # wrap explicitly rather than relying on the counter's natural rollover.
1073 counter = Wire(UInt(counter_width), name="chunk_counter")
1074 client_xact = Wire(Bits(1))
1075 set_client_valid = counter == UInt(counter_width)(chunks - 1)
1076 counter.assign(
1077 Counter(counter_width)(clk=ports.clk,
1078 rst=ports.rst,
1079 clear=(upstream_xact & set_client_valid) |
1080 (client_xact & ~upstream_xact),
1081 increment=upstream_xact,
1082 instance_name="chunk_counter").out)
1083 client_valid = ControlReg(ports.clk, ports.rst,
1084 [set_client_valid & upstream_xact],
1085 [client_xact])
1086 for idx, reg_ce in enumerate(reg_ces):
1087 reg_ce.assign(upstream_xact & (counter == UInt(counter_width)(idx)))
1088 # 'last' of the final engine word that completes this client flit.
1089 client_last = upstream_last.reg(ports.clk,
1090 ports.rst,
1091 ce=upstream_xact,
1092 name="last_reg")
1093 tag_reg = up.tag.reg(ports.clk,
1094 ports.rst,
1095 ce=upstream_xact,
1096 name="tag_reg")
1097
1098 client_channel, client_ready = ConcatReadGearboxImpl.out.type.wrap(
1099 {
1100 "tag": tag_reg,
1101 "data": client_data_bits,
1102 "last": client_last,
1103 }, client_valid)
1104 client_xact.assign(client_valid & client_ready)
1105 ready_for_upstream.assign(~client_valid | client_ready)
1106 ports.out = client_channel
1107
1108 return ConcatReadGearboxImpl
1109
1110
1111@modparams
1112def DepackReadGearbox(input_bitwidth: int,
1113 output_bitwidth: int) -> type["DepackReadGearboxImpl"]:
1114 """Unpack a byte-aligned element that divides the engine word
1115 (``OUT % 8 == 0`` and ``IN % OUT == 0``) from a contiguous list response. Each
1116 word holds ``IN/OUT`` gap-free elements that never straddle, so a counter
1117 drives a parts:1 element mux -- no shifter (e.g. 32b/64b, 64b/256b).
1118 ``valid_bytes`` locates the last element in the burst's (possibly partial)
1119 final word. Straddling relationships use `ShiftReadGearbox`."""
1120
1121 if input_bitwidth % 8 != 0:
1122 raise ValueError("engine word width must be a multiple of 8 bits")
1123 if output_bitwidth == 0 or output_bitwidth % 8 != 0 \
1124 or input_bitwidth % output_bitwidth != 0:
1125 raise ValueError(
1126 "DepackReadGearbox requires a byte-aligned element that divides the "
1127 "engine word")
1128
1129 in_bytes = input_bitwidth // 8
1130 # 'valid_bytes' is the real byte count minus 1; a word always has >= 1 byte.
1131 vb_width = clog2(in_bytes)
1132 count_width = clog2(in_bytes + 1)
1133 parts = input_bitwidth // output_bitwidth
1134 elem_bytes = output_bitwidth // 8
1135
1136 class DepackReadGearboxImpl(Module):
1137 clk = Clock()
1138 rst = Reset()
1139 in_ = InputChannel(
1140 StructType([
1141 ("tag", esi.HostMem.TagType),
1142 ("data", Bits(input_bitwidth)),
1143 ("valid_bytes", UInt(vb_width)),
1144 ("last", Bits(1)),
1145 ]))
1146 out = OutputChannel(
1147 StructType([
1148 ("tag", esi.HostMem.TagType),
1149 ("data", Bits(output_bitwidth)),
1150 ("last", Bits(1)),
1151 ]))
1152
1153 @generator
1154 def build(ports):
1155 client_ready = Wire(Bits(1), name="client_ready")
1156 up_ready = Wire(Bits(1), name="up_ready")
1157 # Register the input for fmax; the ESI channel buffer keeps the handshake
1158 # elastic and decouples the ready path.
1159 in_reg = ports.in_.buffer(ports.clk, ports.rst, stages=1)
1160 up, up_valid = in_reg.unwrap(up_ready)
1161 client_xact = up_valid & client_ready
1162
1163 if parts == 1:
1164 # One element per word; nothing to select.
1165 last_in_word = Bits(1)(1)
1166 client_data = up.data
1167 else:
1168 idx_width = clog2(parts)
1169 idx = Reg(UInt(idx_width),
1170 clk=ports.clk,
1171 rst=ports.rst,
1172 rst_value=0,
1173 ce=client_xact,
1174 name="idx")
1175 # (idx + 1) * elem_bytes == real valid bytes marks the word's last
1176 # element (the final word may hold fewer than `parts`); add 1 back to
1177 # the biased 'valid_bytes' to recover the real count.
1178 real_valid_bytes = (up.valid_bytes + UInt(1)(1)).as_uint(count_width)
1179 consumed = ((idx + UInt(1)(1)) *
1180 UInt(count_width)(elem_bytes)).as_uint(count_width)
1181 last_in_word = consumed == real_valid_bytes
1182 # parts:1 element-select mux -- the entire datapath, no shifter.
1183 word_parts = Array(Bits(output_bitwidth), parts)([
1184 up.data[k * output_bitwidth:(k + 1) * output_bitwidth]
1185 for k in range(parts)
1186 ])
1187 client_data = word_parts[idx]
1188 idx.assign(
1189 Mux(last_in_word, (idx + UInt(1)(1)).as_uint(idx_width),
1190 UInt(idx_width)(0)))
1191
1192 # Consume the buffered word as its last element leaves.
1193 up_ready.assign(client_xact & last_in_word)
1194 client_channel, client_ready_sig = DepackReadGearboxImpl.out.type.wrap(
1195 {
1196 "tag": up.tag,
1197 "data": client_data,
1198 "last": (up.last & last_in_word).as_bits(),
1199 }, up_valid)
1200 client_ready.assign(client_ready_sig)
1201 ports.out = client_channel
1202
1203 return DepackReadGearboxImpl
1204
1205
1206@modparams
1207def ShiftReadGearbox(input_bitwidth: int,
1208 output_bitwidth: int) -> type["ShiftReadGearboxImpl"]:
1209 """Universal fallback: unpack a contiguous, byte-packed element stream (a
1210 `read_list` response) for ANY ``(input_bitwidth, output_bitwidth)`` pair.
1211
1212 Elements are packed at their natural byte stride ``stride = ceil(OUT/8)``
1213 bytes, so element k begins at wire bit ``k*stride*8`` and, in general,
1214 straddles engine-word boundaries at an arbitrary bit offset. A byte-addressed
1215 shift-register accumulator realigns each element across words. This is correct
1216 for every width relationship; `SliceReadGearbox`, `ConcatReadGearbox` and
1217 `DepackReadGearbox` are optimizations that avoid this barrel shifter for the
1218 regular (non-straddling) cases.
1219
1220 Each input word carries ``valid_bytes`` (how many of its bytes are real) and
1221 ``last``. Both are framed to one whole `read_list` request rather than to the
1222 transport: `HostMemReadReqSplitter` drops the per-chunk framing of the reads
1223 it issues and re-derives these from the request's total length, so only the
1224 request's final word is ever partial. That length is ``num_elements *
1225 stride``, so tracking real bytes lets the gearbox emit exactly the right
1226 elements and place the list-terminating ``last`` on the final one -- no
1227 padding element is ever emitted."""
1228
1229 if input_bitwidth % 8 != 0:
1230 raise ValueError("engine word width must be a multiple of 8 bits")
1231 if output_bitwidth <= 0:
1232 raise ValueError("client element width must be positive")
1233 in_bytes = input_bitwidth // 8
1234 stride_bytes = (output_bitwidth + 7) // 8
1235 stride_bits = stride_bytes * 8
1236 # Hold at most one partial element plus one freshly accepted word.
1237 buf_bytes = stride_bytes + in_bytes
1238 buf_bits = buf_bytes * 8
1239 # 'valid_bytes' is the real byte count minus 1; a word always has >= 1 byte.
1240 vb_width = clog2(in_bytes)
1241 cnt_width = clog2(buf_bytes + 1)
1242 # The append offset is only ever in [0, stride_bytes] (has_room), so the shift
1243 # index needs fewer bits than the full count -- see `build`.
1244 offset_width = clog2(stride_bytes + 1)
1245
1246 class ShiftReadGearboxImpl(Module):
1247 clk = Clock()
1248 rst = Reset()
1249 in_ = InputChannel(
1250 StructType([
1251 ("tag", esi.HostMem.TagType),
1252 ("data", Bits(input_bitwidth)),
1253 ("valid_bytes", UInt(vb_width)),
1254 ("last", Bits(1)),
1255 ]))
1256 out = OutputChannel(
1257 StructType([
1258 ("tag", esi.HostMem.TagType),
1259 ("data", Bits(output_bitwidth)),
1260 ("last", Bits(1)),
1261 ]))
1262
1263 @generator
1264 def build(ports):
1265 client_ready = Wire(Bits(1), name="client_ready")
1266 up_ready = Wire(Bits(1), name="up_ready")
1267 # Register the input for fmax; the ESI channel buffer keeps the handshake
1268 # elastic and decouples the ready path.
1269 in_reg = ports.in_.buffer(ports.clk, ports.rst, stages=1)
1270 up, up_valid = in_reg.unwrap(up_ready)
1271
1272 from pycde.circt.dialects import comb
1273
1274 # Byte-addressed accumulator: `buffer` holds `count` valid bytes packed
1275 # from bit 0 up; the element being emitted is buffer[0:OUT].
1276 buffer = Reg(Bits(buf_bits),
1277 clk=ports.clk,
1278 rst=ports.rst,
1279 rst_value=0,
1280 name="buffer")
1281 count = Reg(UInt(cnt_width),
1282 clk=ports.clk,
1283 rst=ports.rst,
1284 rst_value=0,
1285 name="count")
1286 saw_last = Wire(Bits(1), name="saw_last")
1287
1288 # Accept a whole engine word only when there's room, and never while
1289 # draining a finished burst -- otherwise the next burst's bytes would mix
1290 # into this one's buffer.
1291 has_room = count <= UInt(cnt_width)(buf_bytes - in_bytes)
1292 up_ready.assign(has_room & ~saw_last)
1293 up_xact = up_ready & up_valid
1294
1295 # Emit an element once a full stride slot is buffered.
1296 client_valid = count >= UInt(cnt_width)(stride_bytes)
1297 client_xact = client_valid & client_ready
1298
1299 # The burst's final word sets `saw_last`; the emit that drains the buffer
1300 # to empty terminates the list. These never coincide: emitting needs a
1301 # slot buffered by a prior cycle's accept, so `after_emit == 0` on an
1302 # accept cycle is impossible.
1303 added = Mux(up_xact,
1304 UInt(cnt_width)(0), (up.valid_bytes.as_uint(cnt_width) +
1305 UInt(1)(1)).as_uint(cnt_width))
1306 after_add = (count + added).as_uint(cnt_width)
1307 after_emit = (after_add -
1308 UInt(cnt_width)(stride_bytes)).as_uint(cnt_width)
1309 set_saw_last = (up_xact & up.last).as_bits()
1310 is_final_slot = (after_emit == UInt(cnt_width)(0))
1311 burst_ending = saw_last | set_saw_last
1312 client_last = client_valid & burst_ending & is_final_slot
1313 final_emit = client_xact & client_last
1314
1315 # Append the accepted word at bit offset count*8 (dynamic left shift);
1316 # then, if we emit this cycle, drop the consumed slot (constant right
1317 # shift by the stride). The append offset is <= stride_bytes when
1318 # accepting, so bound the shift index to that range: its high bits are
1319 # constant 0, which lets constant-propagation prune the upper barrel-
1320 # shifter stages (synthesis won't infer this bound from the count reg).
1321 append_off = count.as_bits()[0:offset_width]
1322 shamt = BitsSignal.concat([append_off,
1323 Bits(3)(0)]).pad_or_truncate(buf_bits)
1324 word_ext = up.data.pad_or_truncate(buf_bits)
1325 shifted_word = BitsSignal(
1326 comb.ShlOp(word_ext.value, shamt.value).result, Bits(buf_bits))
1327 appended = buffer | Mux(up_xact, Bits(buf_bits)(0), shifted_word)
1328 drained = appended[stride_bits:buf_bits].pad_or_truncate(buf_bits)
1329 buffer.assign(
1330 Mux(final_emit, Mux(client_xact, appended, drained),
1331 Bits(buf_bits)(0)))
1332
1333 # count += accepted real bytes (biased 'valid_bytes' + 1); -= stride on
1334 # emit.
1335 count.assign(Mux(client_xact, after_add, after_emit))
1336
1337 saw_last.assign(
1338 ControlReg(ports.clk, ports.rst, [set_saw_last], [final_emit]))
1339
1340 tag_reg = up.tag.reg(ports.clk, ports.rst, ce=up_xact, name="tag_reg")
1341 client_channel, client_ready_sig = ShiftReadGearboxImpl.out.type.wrap(
1342 {
1343 "tag": tag_reg,
1344 "data": buffer[0:output_bitwidth],
1345 "last": client_last,
1346 }, client_valid)
1347 client_ready.assign(client_ready_sig)
1348 ports.out = client_channel
1349
1350 return ShiftReadGearboxImpl
1351
1352
1353def select_read_gearbox(is_list: bool, input_bitwidth: int,
1354 output_bitwidth: int):
1355 """Pick the read-gearbox module for a client of the given kind and width
1356 relationship. Every gearbox shares the {tag, data, valid_bytes, last} input
1357 (from `HostMemReadReqSplitter`) and the {tag, data, last} output, so callers
1358 wire them identically. `ShiftReadGearbox` is the correct-for-everything
1359 fallback; the others avoid its barrel shifter for regular relationships."""
1360 if not is_list:
1361 # A single element starts at bit 0 and never straddles at a bit offset.
1362 if output_bitwidth <= input_bitwidth:
1363 return SliceReadGearbox(input_bitwidth, output_bitwidth)
1364 return ConcatReadGearbox(input_bitwidth, output_bitwidth)
1365 if output_bitwidth > input_bitwidth:
1366 # Super-word list: a whole-word-multiple element never straddles.
1367 if output_bitwidth % input_bitwidth == 0:
1368 return ConcatReadGearbox(input_bitwidth, output_bitwidth)
1369 return ShiftReadGearbox(input_bitwidth, output_bitwidth)
1370 # Sub-word list: a byte-aligned element that divides the word never straddles.
1371 if input_bitwidth % output_bitwidth == 0 and output_bitwidth % 8 == 0:
1372 return DepackReadGearbox(input_bitwidth, output_bitwidth)
1373 return ShiftReadGearbox(input_bitwidth, output_bitwidth)
1374
1375
1376# Maximum size, in bytes, of a single upstream read request. Reads larger than
1377# this are split by the requester into multiple requests. The default is a
1378# conservative PCIe-derived cap (Max_Read_Request_Size tops out at 4096 bytes,
1379# but root ports often negotiate less); it mirrors kPcieMaxReadRequestBytes in
1380# the Cosim backend (cpp/lib/backends/Cosim.cpp).
1381DEFAULT_MAX_READ_REQUEST_BYTES = 64 * 4 # 64 double words
1382
1383# Maximum size, in bytes, of a single upstream write transaction; an element
1384# whose write payload is wider is split into multiple <= this-size transactions.
1385# The default is a conservative PCIe-derived Max-Payload-Size cap.
1386DEFAULT_MAX_WRITE_PAYLOAD_BYTES = 256
1387
1388
1389@modparams
1390def HostMemReadReqSplitter(req_channel_type: Channel,
1391 resp_channel_type: Channel, max_chunk_bytes: int):
1392 """Split oversized host memory read requests into request-sized chunks before
1393 arbitration and reassemble the per-chunk responses into a single logical
1394 burst.
1395
1396 A burst read (`read_list`) can request many more bytes than a single upstream
1397 read request can carry. This module breaks such a request into
1398 `max_chunk_bytes`-sized (word-aligned) chunks addressed sequentially from the
1399 base. Splitting here -- *before* the requests
1400 are arbitrated onto the shared upstream read channel -- lets each client's
1401 chunks interleave with other clients' requests, so one large burst does not
1402 monopolize host memory bandwidth.
1403
1404 On the response path the per-chunk end-of-list markers are dropped and a
1405 single burst-final `last` is re-derived from the total transfer length, so the
1406 gearbox and client see one contiguous response stream identical to an unsplit
1407 read.
1408
1409 Only one logical request is in flight at a time (matching the read processor's
1410 one-outstanding-transaction-per-client model): a new request is not accepted
1411 until the current burst's chunks have all been issued and its responses have
1412 fully drained. This will be a performance limiter.
1413 TODO: make this able to issue >1 one read at a time.
1414
1415 req_channel_type: channel of the upstream read request {address, length
1416 (bytes), tag}.
1417 resp_channel_type: channel of the upstream response {tag, data, last}.
1418 max_chunk_bytes: largest per-chunk byte count; must be > 0 and a multiple of
1419 the response word size.
1420 """
1421 assert max_chunk_bytes > 0
1422
1423 req_struct = req_channel_type.inner_type
1424 resp_struct = resp_channel_type.inner_type
1425 req_fields = dict(req_struct.fields)
1426 addr_width = req_fields["address"].bitwidth
1427 length_width = req_fields["length"].bitwidth
1428 tag_type = req_fields["tag"]
1429 word_bytes = dict(resp_struct.fields)["data"].bitwidth // 8
1430 word_shift = clog2(word_bytes)
1431 words_width = length_width - word_shift
1432 # The response is augmented with a per-word 'valid_bytes': the number of real
1433 # bytes in the (possibly partial) final word, biased by -1. A burst word
1434 # always has >= 1 real byte, so encoding count-1 fits in one fewer bit.
1435 vb_width = clog2(word_bytes)
1436 resp_fields = dict(resp_struct.fields)
1437 resp_out_struct = StructType([
1438 ("tag", resp_fields["tag"]),
1439 ("data", resp_fields["data"]),
1440 ("valid_bytes", UInt(vb_width)),
1441 ("last", Bits(1)),
1442 ])
1443 resp_out_channel_type = Channel(resp_out_struct)
1444
1445 class HostMemReadReqSplitterImpl(Module):
1446 clk = Clock()
1447 rst = Reset()
1448 req_in = Input(req_channel_type)
1449 req_out = Output(req_channel_type)
1450 resp_in = Input(resp_channel_type)
1451 resp_out = Output(resp_out_channel_type)
1452
1453 @generator
1454 def build(ports):
1455 clk = ports.clk
1456 rst = ports.rst
1457
1458 # Burst state shared by the request-splitting and response-reassembly
1459 # FSMs. One logical request is processed at a time.
1460 emit_busy = Wire(Bits(1), name="emit_busy") # issuing chunk requests
1461 resp_busy = Wire(Bits(1), name="resp_busy") # responses still draining
1462 cur_addr = Wire(UInt(addr_width), name="cur_addr")
1463 remaining = Wire(UInt(length_width), name="remaining") # req bytes left
1464 tag_reg = Wire(tag_type, name="tag_reg")
1465 words_left = Wire(UInt(words_width), name="words_left") # resp words left
1466
1467 idle = (~emit_busy) & (~resp_busy)
1468
1469 # --- Request intake and splitting ---
1470 req_ready = Wire(Bits(1))
1471 req_payload, req_valid = ports.req_in.unwrap(req_ready)
1472 accept = idle & req_valid
1473 req_ready.assign(accept)
1474
1475 max_chunk = UInt(length_width)(max_chunk_bytes)
1476 chunk_len = Mux(remaining > max_chunk, remaining, max_chunk)
1477 last_chunk = remaining <= max_chunk
1478
1479 # Round the emitted read length up to a whole word. The reader response
1480 # is word-granular and 'valid_bytes' still carries the real trailing byte
1481 # count, so total_words and the reassembled element count are unchanged;
1482 # this just keeps every read word-aligned for single-flit HostMem
1483 # transports that reject sub-word read lengths.
1484 if word_shift == 0:
1485 chunk_len_out = chunk_len
1486 else:
1487 chunk_words = (chunk_len + UInt(length_width)(word_bytes - 1)
1488 ).as_bits()[word_shift:].as_uint(words_width)
1489 chunk_len_out = BitsSignal.concat(
1490 [chunk_words.as_bits(), Bits(word_shift)(0)]).as_uint(length_width)
1491
1492 req_out_ch, req_out_ready = req_channel_type.wrap(
1493 req_struct({
1494 "address": cur_addr,
1495 "length": chunk_len_out,
1496 "tag": tag_reg,
1497 }), emit_busy)
1498 ports.req_out = req_out_ch
1499 chunk_xact = emit_busy & req_out_ready
1500
1501 emit_busy.assign(
1502 ControlReg(clk,
1503 rst, [accept], [chunk_xact & last_chunk],
1504 name="emit_busy_reg"))
1505
1506 # cur_addr: load base on accept, advance by the chunk on each issue.
1507 cur_addr_incr = (cur_addr +
1508 chunk_len.as_uint(addr_width)).as_uint(addr_width)
1509 cur_addr.assign(
1510 Mux(accept, Mux(chunk_xact, cur_addr, cur_addr_incr),
1511 req_payload.address).reg(clk,
1512 rst,
1513 rst_value=0,
1514 ce=accept | chunk_xact,
1515 name="cur_addr_reg"))
1516
1517 # remaining: load length on accept, subtract each issued chunk.
1518 remaining_dec = (remaining - chunk_len).as_uint(length_width)
1519 remaining.assign(
1520 Mux(accept, Mux(chunk_xact, remaining, remaining_dec),
1521 req_payload.length).reg(clk,
1522 rst,
1523 rst_value=0,
1524 ce=accept | chunk_xact,
1525 name="remaining_reg"))
1526
1527 tag_reg.assign(req_payload.tag.reg(clk, rst, ce=accept, name="tag_reg_r"))
1528
1529 # --- Response reassembly: re-derive the burst-final 'last' and the byte
1530 # count of the (possibly partial) final word. Elements need not tile
1531 # evenly into words, so count words with ceil(length / word_bytes). ---
1532 total_words = ((req_payload.length + UInt(length_width)(word_bytes - 1)
1533 ).as_bits()[word_shift:]).as_uint(words_width)
1534 # Bytes valid in the final word = length - (total_words - 1) * word_bytes.
1535 words_before_last = (total_words -
1536 UInt(words_width)(1)).as_uint(words_width)
1537 bytes_before_last = BitsSignal.concat(
1538 [words_before_last.as_bits(),
1539 Bits(word_shift)(0)]).as_uint(length_width)
1540 final_valid_bytes = (req_payload.length - bytes_before_last -
1541 UInt(length_width)(1)).as_uint(vb_width).reg(
1542 clk, rst, ce=accept, name="final_valid_bytes")
1543 resp_ready = Wire(Bits(1))
1544 resp_payload, resp_valid = ports.resp_in.unwrap(resp_ready)
1545 is_final_word = words_left == UInt(words_width)(1)
1546 resp_out_ch, resp_out_ready = resp_out_channel_type.wrap(
1547 resp_out_struct({
1548 "tag":
1549 resp_payload.tag,
1550 "data":
1551 resp_payload.data,
1552 "valid_bytes":
1553 Mux(is_final_word,
1554 UInt(vb_width)(word_bytes - 1), final_valid_bytes),
1555 "last":
1556 is_final_word,
1557 }), resp_valid)
1558 ports.resp_out = resp_out_ch
1559 resp_ready.assign(resp_out_ready)
1560 resp_xact = resp_valid & resp_out_ready
1561
1562 # words_left: load total on accept, decrement per received word.
1563 words_dec = (words_left - UInt(words_width)(1)).as_uint(words_width)
1564 words_left.assign(
1565 Mux(accept, Mux(resp_xact, words_left, words_dec),
1566 total_words).reg(clk,
1567 rst,
1568 rst_value=0,
1569 ce=accept | resp_xact,
1570 name="words_left_reg"))
1571
1572 resp_busy.assign(
1573 ControlReg(clk,
1574 rst, [accept], [resp_xact & is_final_word],
1575 name="resp_busy_reg"))
1576
1577 return HostMemReadReqSplitterImpl
1578
1579
1581 read_width: int,
1582 hostmem_module,
1583 reqs: List[esi._OutputBundleSetter],
1584 max_read_request_bytes: int = DEFAULT_MAX_READ_REQUEST_BYTES):
1585 """Construct a host memory read request module to orchestrate the the read
1586 connections. Responsible for both gearboxing the data, multiplexing the
1587 requests, reassembling out-of-order responses and routing the responses to the
1588 correct clients.
1589
1590 Generate this module dynamically to allow for multiple read clients of
1591 multiple types to be directly accomodated."""
1592
1593 class HostmemReadProcessorImpl(Module):
1594 clk = Clock()
1595 rst = Reset()
1596
1597 # Add an output port for each read client.
1598 reqPortMap: Dict[esi._OutputBundleSetter, str] = {}
1599 for req in reqs:
1600 name = "client_" + req.client_name_str
1601 locals()[name] = Output(req.type)
1602 reqPortMap[req] = name
1603
1604 # And then the port which goes to the host.
1605 upstream = Output(hostmem_module.read.type)
1606
1607 @generator
1608 def build(ports):
1609 """Build the read side of the HostMem service."""
1610
1611 # If there's no read clients, just return a no-op read bundle.
1612 if len(reqs) == 0:
1613 upstream_req_channel, _ = Channel(hostmem_module.UpstreamReadReq).wrap(
1614 {
1615 "tag": 0,
1616 "length": 0,
1617 "address": 0
1618 }, 0)
1619 upstream_read_bundle, _ = hostmem_module.read.type.pack(
1620 req=upstream_req_channel)
1621 ports.upstream = upstream_read_bundle
1622 return
1623
1624 # Since we use the tag to identify the client, we can't have more than 256
1625 # read clients. Supporting more than 256 clients would require
1626 # tag-rewriting, which we'll probably have to implement at some point.
1627 # TODO: Implement tag-rewriting.
1628 assert len(reqs) <= 256, "More than 256 read clients not supported."
1629
1630 # Pack the upstream bundle and leave the request as a wire.
1631 upstream_req_channel = Wire(Channel(hostmem_module.UpstreamReadReq))
1632 upstream_read_bundle, froms = hostmem_module.read.type.pack(
1633 req=upstream_req_channel)
1634 ports.upstream = upstream_read_bundle
1635 upstream_resp_channel = froms["resp"]
1636
1637 # Demux the upstream response frames {tag, data, last} to each client by
1638 # tag. Each client's stream then flows through a `HostMemReadReqSplitter`
1639 # (which annotates per-word 'valid_bytes' and the burst-final 'last') into
1640 # the leaf gearbox chosen by `select_read_gearbox`.
1641 demux = esi.TaggedDemux(len(reqs), upstream_resp_channel.type)(
1642 clk=ports.clk, rst=ports.rst, in_=upstream_resp_channel)
1643
1644 word_bytes = read_width // 8
1645 tagged_client_reqs = []
1646 for idx, client in enumerate(reqs):
1647 # Find the response channel in the request bundle.
1648 resp_type = [
1649 c.channel for c in client.type.channels if c.name == 'resp'
1650 ][0]
1651 demuxed_upstream_channel = demux.get_out(idx)
1652
1653 # TODO: Should responses come back out-of-order (interleaved tags),
1654 # re-order them here so the gearbox doesn't get confused. (Longer term.)
1655 # For now, only support one outstanding transaction at a time. This has
1656 # the additional benefit of letting the upstream tag be the client
1657 # identifier. TODO: Implement the gating logic here.
1658 client_type = resp_type.inner_type
1659 is_list = isinstance(client_type, Window)
1660
1661 # A read_list response is a parallel window over
1662 # struct{tag, data: list<element>} (num_items=1), lowering to
1663 # struct{tag, data: element, data_size, last}; a single read carries the
1664 # element directly. Pull the element width out of whichever shape.
1665 if is_list:
1666 lowered = client_type.lowered_type
1667 lowered_fields = dict(lowered.fields)
1668 element_type = lowered_fields["data"]
1669 element_bits = element_type.bitwidth
1670 data_size_type = lowered_fields["data_size"]
1671 if element_bits == 0:
1672 raise ValueError("read_list element type cannot be zero-width.")
1673 else:
1674 if client_type.data.bitwidth == 0:
1675 raise ValueError("Client data type cannot be zero-width. Use a "
1676 "single-bit type if no data is needed.")
1677 element_bits = client_type.data.bitwidth
1678 # Elements are packed contiguously in host memory at their natural byte
1679 # size, independent of the engine word width.
1680 elem_stride_bytes = (element_bits + 7) // 8
1681
1682 # Both single-message and list reads flow demux -> splitter -> gearbox
1683 # with a uniform {tag, data, valid_bytes, last} interface. The splitter
1684 # chunks oversized requests (so even a wide single element is
1685 # request-chunked) and annotates each word with 'valid_bytes' plus the
1686 # burst-final 'last'; `select_read_gearbox` picks the leaf gearbox for
1687 # this (is_list, read_width, element_bits). 'splitter_resp' breaks the
1688 # request/response construction cycle (the client request is derived
1689 # from the gearbox's response bundle).
1690 max_chunk_bytes = (max_read_request_bytes // word_bytes) * word_bytes
1691 gearbox_mod = select_read_gearbox(is_list, read_width, element_bits)
1692 splitter_resp = Wire(gearbox_mod.in_.type)
1693 gearbox = gearbox_mod(clk=ports.clk, rst=ports.rst, in_=splitter_resp)
1694
1695 if is_list:
1696 # Propagate 'last', then re-wrap the element as the response window.
1697 client_resp_channel = gearbox.out.transform(
1698 lambda m, lowered=lowered, element_type=element_type,
1699 data_size_type=data_size_type, client_type=client_type:
1700 client_type.wrap(
1701 lowered({
1702 "tag": m.tag,
1703 "data": m.data.bitcast(element_type),
1704 "data_size": data_size_type(0),
1705 "last": m.last,
1706 })))
1707 client_bundle, froms = client.type.pack(resp=client_resp_channel)
1708 client_req = froms["req"]
1709 logical_req = client_req.transform(
1710 lambda r, idx=idx, elem_stride_bytes=elem_stride_bytes:
1711 hostmem_module.UpstreamReadReq({
1712 "address":
1713 r.address,
1714 "length": (r.length * UInt(64)
1715 (elem_stride_bytes)).as_uint(32),
1716 "tag":
1717 idx,
1718 }))
1719 else:
1720 # Single-message read: one element; discard the 'last' burst marker.
1721 client_resp_channel = gearbox.out.transform(
1722 lambda m, client_type=client_type: client_type({
1723 "tag": m.tag,
1724 "data": m.data.bitcast(client_type.data)
1725 }))
1726 client_bundle, froms = client.type.pack(resp=client_resp_channel)
1727 client_req = froms["req"]
1728 logical_req = client_req.transform(
1729 lambda r, idx=idx, elem_stride_bytes=elem_stride_bytes:
1730 hostmem_module.UpstreamReadReq({
1731 "address": r.address,
1732 "length": UInt(32)(elem_stride_bytes),
1733 # TODO: Change this once we support tag-rewriting.
1734 "tag": idx,
1735 }))
1736
1737 splitter = HostMemReadReqSplitter(
1738 logical_req.type, demuxed_upstream_channel.type,
1739 max_chunk_bytes)(clk=ports.clk,
1740 rst=ports.rst,
1741 req_in=logical_req,
1742 resp_in=demuxed_upstream_channel)
1743 splitter_resp.assign(splitter.resp_out)
1744 tagged_client_req = splitter.req_out
1745
1746 tagged_client_reqs.append(tagged_client_req)
1747
1748 # Set the port for the client request.
1749 setattr(ports, HostmemReadProcessorImpl.reqPortMap[client],
1750 client_bundle)
1751
1752 # Assign the multiplexed read request to the upstream request. Use the
1753 # list-aware, pipelined ChannelArbiter (vs. the combinational ChannelMux)
1754 # for a registered N:1 mux that closes timing at high client fan-in. Read
1755 # requests are single-flit, so list-awareness is a no-op here.
1756 # `mux_pipeline_levels=2` retimes the wide payload selection mux, whose
1757 # depth otherwise grows as log2(num_clients); the added latency is
1758 # absorbed by the arbiter's output FIFO / credit counter.
1759 # TODO: Don't release a request until the client is ready to accept
1760 # the response otherwise the system could deadlock.
1761 muxed_client_reqs = ChannelArbiter(tagged_client_reqs,
1762 ports.clk,
1763 ports.rst,
1764 mux_pipeline_levels=2,
1765 pipelined_scheduler=True,
1766 telemetry=False)
1767 upstream_req_channel.assign(muxed_client_reqs)
1768 HostmemReadProcessorImpl.reqPortMap.clear()
1769
1770 return HostmemReadProcessorImpl
1771
1772
1773@modparams
1774def TaggedWriteGearbox(input_bitwidth: int, output_bitwidth: int,
1775 max_burst_bytes: int) -> type["TaggedWriteGearboxImpl"]:
1776 """Build a gearbox to convert the client data to upstream write chunks.
1777 Assumes a struct {address, tag, data} and only gearboxes the data. Tag is
1778 stored separately and the struct is re-assembled later on.
1779
1780 'max_burst_bytes' caps a single contiguous upstream write transaction (a
1781 max-payload-size analog): when an element spans more than 'max_burst_bytes',
1782 its engine words are split into multiple <= 'max_burst_bytes' transactions by
1783 emitting the framing 'last' at each boundary. 0 disables the cap."""
1784
1785 if output_bitwidth % 8 != 0:
1786 raise ValueError("Output bitwidth must be a multiple of 8.")
1787 input_pad_bits = 0
1788 if input_bitwidth % 8 != 0:
1789 input_pad_bits = 8 - (input_bitwidth % 8)
1790 input_padded_bitwidth = input_bitwidth + input_pad_bits
1791
1792 # Number of engine words per capped transaction (0 = uncapped).
1793 max_burst_words = (max_burst_bytes //
1794 (output_bitwidth // 8)) if max_burst_bytes else 0
1795 if max_burst_words:
1796 assert (max_burst_words & (max_burst_words - 1)) == 0, \
1797 "max_burst_bytes / (output_bitwidth // 8) must be a power of two"
1798
1799 class TaggedWriteGearboxImpl(Module):
1800 clk = Clock()
1801 rst = Reset()
1802 in_ = InputChannel(
1803 StructType([
1804 ("address", UInt(64)),
1805 ("tag", esi.HostMem.TagType),
1806 ("data", Bits(input_bitwidth)),
1807 ]))
1808 out = OutputChannel(
1809 StructType([
1810 ("address", UInt(64)),
1811 ("tag", esi.HostMem.TagType),
1812 ("data", Bits(output_bitwidth)),
1813 ("valid_bytes", Bits(8)),
1814 ("last", Bits(1)),
1815 ]))
1816
1817 num_chunks = ceil(input_padded_bitwidth / output_bitwidth)
1818
1819 @generator
1820 def build(ports):
1821 upstream_ready = Wire(Bits(1))
1822 ready_for_client = Wire(Bits(1))
1823 client_tag_and_data, client_valid = ports.in_.unwrap(ready_for_client)
1824 client_data = client_tag_and_data.data
1825 if input_pad_bits > 0:
1826 client_data = client_data.pad_or_truncate(input_padded_bitwidth)
1827 client_xact = ready_for_client & client_valid
1828 input_bitwidth_bytes = input_padded_bitwidth // 8
1829 output_bitwidth_bytes = output_bitwidth // 8
1830
1831 # Determine if gearboxing is necessary and whether it needs to be
1832 # gearboxed up or just sliced down.
1833 if output_bitwidth == input_padded_bitwidth:
1834 upstream_data_bits = client_data
1835 upstream_valid = client_valid
1836 ready_for_client.assign(upstream_ready)
1837 tag = client_tag_and_data.tag
1838 address = client_tag_and_data.address
1839 valid_bytes = Bits(8)(input_bitwidth_bytes)
1840 last = Bits(1)(1)
1841 elif output_bitwidth > input_padded_bitwidth:
1842 upstream_data_bits = client_data.as_bits(output_bitwidth)
1843 upstream_valid = client_valid
1844 ready_for_client.assign(upstream_ready)
1845 tag = client_tag_and_data.tag
1846 address = client_tag_and_data.address
1847 valid_bytes = Bits(8)(input_bitwidth_bytes)
1848 last = Bits(1)(1)
1849 else:
1850 # Create registers equal to the number of upstream transactions needed
1851 # to complete the transmission.
1852 num_chunks = TaggedWriteGearboxImpl.num_chunks
1853 num_chunks_idx_bitwidth = clog2(num_chunks)
1854 if input_padded_bitwidth % output_bitwidth == 0:
1855 padding_numbits = 0
1856 else:
1857 padding_numbits = output_bitwidth - (input_padded_bitwidth %
1858 output_bitwidth)
1859 client_data_padded = BitsSignal.concat(
1860 [Bits(padding_numbits)(0), client_data])
1861 chunks = [
1862 client_data_padded[i * output_bitwidth:(i + 1) * output_bitwidth]
1863 for i in range(num_chunks)
1864 ]
1865 chunk_regs = Array(Bits(output_bitwidth), num_chunks)([
1866 c.reg(ports.clk, ce=client_xact, name=f"chunk_{idx}")
1867 for idx, c in enumerate(chunks)
1868 ])
1869 increment = Wire(Bits(1))
1870 clear = Wire(Bits(1))
1871 counter = Counter(num_chunks_idx_bitwidth)(clk=ports.clk,
1872 rst=ports.rst,
1873 increment=increment,
1874 clear=clear)
1875 upstream_data_bits = chunk_regs[counter.out]
1876 upstream_valid = ControlReg(ports.clk, ports.rst, [client_xact],
1877 [clear])
1878 upstream_xact = upstream_valid & upstream_ready
1879 clear.assign(upstream_xact & (counter.out == (num_chunks - 1)))
1880 increment.assign(upstream_xact)
1881 ready_for_client.assign(~upstream_valid)
1882 address_padding_bits = clog2(output_bitwidth_bytes)
1883 counter_bytes = BitsSignal.concat(
1884 [counter.out.as_bits(),
1885 Bits(address_padding_bits)(0)]).as_uint()
1886
1887 # Construct the output channel. Shared logic across all three cases.
1888 tag_reg = client_tag_and_data.tag.reg(ports.clk,
1889 ce=client_xact,
1890 name="tag_reg")
1891 addr_reg = client_tag_and_data.address.reg(ports.clk,
1892 ce=client_xact,
1893 name="address_reg")
1894 address = (addr_reg + counter_bytes).as_uint(64)
1895 tag = tag_reg
1896 elem_end = counter.out == (num_chunks - 1)
1897 valid_bytes = Mux(elem_end,
1898 Bits(8)(output_bitwidth_bytes),
1899 Bits(8)((output_bitwidth - padding_numbits) // 8))
1900 if max_burst_words and num_chunks > max_burst_words:
1901 # Max-payload-size cap: end the upstream write transaction at the
1902 # element end OR every max_burst_words engine words, whichever comes
1903 # first, so a wide element's write is split into <= max_burst_bytes
1904 # transactions. Each word keeps its own sequential address; only the
1905 # transaction-framing 'last' changes.
1906 burst_shift = clog2(max_burst_words)
1907 burst_end = counter.out.as_bits()[:burst_shift].and_reduce()
1908 last = elem_end | burst_end
1909 else:
1910 last = elem_end
1911
1912 upstream_channel, upstrm_ready_sig = TaggedWriteGearboxImpl.out.type.wrap(
1913 {
1914 "address": address,
1915 "tag": tag,
1916 "data": upstream_data_bits,
1917 "valid_bytes": valid_bytes,
1918 "last": last,
1919 }, upstream_valid)
1920 upstream_ready.assign(upstrm_ready_sig)
1921 ports.out = upstream_channel
1922
1923 return TaggedWriteGearboxImpl
1924
1925
1926@modparams
1927def EmitEveryN(message_type: Type, N: int) -> type['EmitEveryNImpl']:
1928 """Emit (forward) one message for every N input messages. The emitted message
1929 is the last one of the N received. N must be >= 1."""
1930
1931 if N < 1:
1932 raise ValueError("N must be >= 1")
1933
1934 class EmitEveryNImpl(Module):
1935 clk = Clock()
1936 rst = Reset()
1937 in_ = InputChannel(message_type)
1938 out = OutputChannel(message_type)
1939
1940 @generator
1941 def build(ports):
1942 ready_for_in = Wire(Bits(1))
1943 in_data, in_valid = ports.in_.unwrap(ready_for_in)
1944 xact = in_valid & ready_for_in
1945
1946 # Fast path: N == 1 -> pass-through.
1947 if N == 1:
1948 out_chan, out_ready = EmitEveryNImpl.out.type.wrap(in_data, in_valid)
1949 ready_for_in.assign(out_ready)
1950 ports.out = out_chan
1951 return
1952
1953 counter_width = clog2(N)
1954 counter_clear = Wire(Bits(1))
1955 counter = Counter(counter_width)(clk=ports.clk,
1956 rst=ports.rst,
1957 increment=xact,
1958 clear=counter_clear)
1959
1960 # Capture last message of the group.
1961 last_msg = in_data.reg(ports.clk, ports.rst, ce=xact, name="last_msg")
1962 # Clear the counter.
1963 hit_last = (counter.out == UInt(counter_width)(N - 1)) & xact
1964 counter_clear.assign(hit_last)
1965
1966 emit_accepted = Wire(Bits(1))
1967 out_valid = ControlReg(ports.clk, ports.rst, [hit_last], [emit_accepted])
1968
1969 out_chan, out_ready = EmitEveryNImpl.out.type.wrap(last_msg, out_valid)
1970 # Stall input while waiting for downstream to accept the aggregated output.
1971 ready_for_in.assign(~(out_valid & ~out_ready))
1972 emit_accepted.assign(out_valid & out_ready) # Output consumed downstream.
1973
1974 ports.out = out_chan
1975
1976 return EmitEveryNImpl
1977
1978
1980 write_width: int,
1981 hostmem_module,
1982 reqs: List[esi._OutputBundleSetter],
1983 max_write_payload_bytes: int = DEFAULT_MAX_WRITE_PAYLOAD_BYTES
1984) -> type["HostMemWriteProcessorImpl"]:
1985 """Construct a host memory write request module to orchestrate the the write
1986 connections. Responsible for both gearboxing the data, multiplexing the
1987 requests, reassembling out-of-order responses and routing the responses to the
1988 correct clients.
1989
1990 Generate this module dynamically to allow for multiple write clients of
1991 multiple types to be directly accomodated."""
1992
1993 class HostMemWriteProcessorImpl(Module):
1994
1995 clk = Clock()
1996 rst = Reset()
1997
1998 # Add an output port for each read client.
1999 reqPortMap: Dict[esi._OutputBundleSetter, str] = {}
2000 for req in reqs:
2001 name = "client_" + req.client_name_str
2002 locals()[name] = Output(req.type)
2003 reqPortMap[req] = name
2004
2005 # And then the port which goes to the host.
2006 upstream = Output(hostmem_module.write.type)
2007
2008 @generator
2009 def build(ports):
2010 clk = ports.clk
2011 rst = ports.rst
2012
2013 # Width of the frame's 'data_size' field: log2 of the number of bytes per
2014 # engine word. It holds (valid_bytes - 1) for the final (possibly partial)
2015 # word of a write.
2016 size_width = clog2(write_width // 8)
2017
2018 # If there's no write clients, just create a no-op write bundle
2019 if len(reqs) == 0:
2020 req, _ = Channel(hostmem_module.UpstreamWriteReq).wrap(
2021 {
2022 "address": 0,
2023 "tag": 0,
2024 "data": 0,
2025 "data_size": 0,
2026 "last": 0,
2027 }, 0)
2028 write_bundle, _ = hostmem_module.write.type.pack(req=req)
2029 ports.upstream = write_bundle
2030 return
2031
2032 assert len(reqs) <= 256, "More than 256 write clients not supported."
2033
2034 upstream_req_channel = Wire(Channel(hostmem_module.UpstreamWriteReq))
2035 upstream_write_bundle, froms = hostmem_module.write.type.pack(
2036 req=upstream_req_channel)
2037 ports.upstream = upstream_write_bundle
2038 upstream_ack_tag = froms["ackTag"]
2039
2040 demuxed_acks = esi.TaggedDemux(len(reqs), upstream_ack_tag.type)(
2041 clk=ports.clk, rst=ports.rst, in_=upstream_ack_tag)
2042
2043 # TODO: re-write the tags and store the client and client tag.
2044
2045 # Build the write request channels and ack wires.
2046 write_channels: List[ChannelSignal] = []
2047 for idx, req in enumerate(reqs):
2048 # Get the request channel and its data type.
2049 reqch = [c.channel for c in req.type.channels if c.name == 'req'][0]
2050 client_type = reqch.inner_type
2051 input_flit_ack = Wire(upstream_ack_tag.type)
2052
2053 if isinstance(client_type, Window):
2054 # Windowed (list) write: the client streams a list of elements to be
2055 # written to sequential addresses from a base. Lowered frame:
2056 # struct{address, tag, data: elem[num_items], data_size, last}. One
2057 # element per frame (num_items=1 here).
2058 bundle_sig, wfroms = req.type.pack(ackTag=input_flit_ack)
2059 windowed_req = wfroms["req"]
2060 lowered = client_type.lowered_type
2061 array_type = dict(lowered.fields)["data"]
2062 element_bits = array_type.element_type.bitwidth
2063 # Elements are packed contiguously in host memory at their natural
2064 # byte size, independent of the engine word width (matches the
2065 # read_list path). Each per-element write is byte-enabled via
2066 # data_size, so a sub-word element writes only its own bytes.
2067 elem_stride = (element_bits + 7) // 8
2068
2069 gearbox_mod = TaggedWriteGearbox(element_bits, write_width,
2070 max_write_payload_bytes)
2071 gearbox_in_type = gearbox_mod.in_.type.inner_type
2072
2073 # Unwrap the window frames; compute a base+offset address from a
2074 # per-burst element counter (reset after each burst's final element).
2075 ready_for_frame = Wire(Bits(1))
2076 frame_win, frame_valid = windowed_req.unwrap(ready_for_frame)
2077 frame = frame_win.unwrap()
2078 frame_xact = frame_valid & ready_for_frame
2079 elem_clear = Wire(Bits(1))
2080 elem_counter = Counter(64)(clk=ports.clk,
2081 rst=ports.rst,
2082 clear=elem_clear,
2083 increment=frame_xact)
2084 elem_clear.assign(frame_xact & frame["last"])
2085 elem_addr = (frame["address"] +
2086 elem_counter.out * UInt(64)(elem_stride)).as_uint(64)
2087 gearbox_in_chan, gearbox_in_ready = Channel(gearbox_in_type).wrap(
2088 gearbox_in_type({
2089 "tag": frame["tag"],
2090 "address": elem_addr,
2091 "data": frame["data"][0].bitcast(gearbox_in_type.data),
2092 }), frame_valid)
2093 ready_for_frame.assign(gearbox_in_ready)
2094 gearbox = gearbox_mod(clk=ports.clk,
2095 rst=ports.rst,
2096 in_=gearbox_in_chan)
2097 else:
2098 # Single-message write.
2099 write_req_bundle_type = esi.HostMem.write_req_bundle_type(
2100 client_type.data)
2101 bundle_sig, sfroms = write_req_bundle_type.pack(ackTag=input_flit_ack)
2102 gearbox_mod = TaggedWriteGearbox(client_type.data.bitwidth,
2103 write_width, max_write_payload_bytes)
2104 gearbox_in_type = gearbox_mod.in_.type.inner_type
2105 bitcast_client_req = sfroms["req"].transform(
2106 lambda m, git=gearbox_in_type: git({
2107 "tag": m.tag,
2108 "address": m.address,
2109 "data": m.data.bitcast(git.data)
2110 }))
2111 gearbox = gearbox_mod(clk=ports.clk,
2112 rst=ports.rst,
2113 in_=bitcast_client_req)
2114
2115 write_channels.append(
2116 gearbox.out.transform(
2117 lambda m, idx=idx: hostmem_module.UpstreamWriteReq({
2118 "address":
2119 m.address,
2120 "tag":
2121 idx,
2122 "data":
2123 m.data,
2124 "data_size": (m.valid_bytes.as_uint() - UInt(8)
2125 (1)).as_bits()[:size_width],
2126 "last":
2127 m.last,
2128 })))
2129
2130 # Count the number of acks received from hostmem for this client
2131 # and only send one back to the client per input.
2132 ack_every_n = EmitEveryN(upstream_ack_tag.type, gearbox_mod.num_chunks)(
2133 clk=clk, rst=rst, in_=demuxed_acks.get_out(idx))
2134 input_flit_ack.assign(ack_every_n.out)
2135
2136 # Set the port for the client request.
2137 setattr(ports, HostMemWriteProcessorImpl.reqPortMap[req], bundle_sig)
2138
2139 # Multiplex the write requests onto the single upstream channel with the
2140 # list-aware, pipelined ChannelArbiter (matching the read side). A real
2141 # windowed write (multi-word client flits) engages the arbiter's list-
2142 # awareness -- via the frame's 'last' -- to keep a client's words
2143 # contiguous; single-word (<= engine width) clients emit one message per
2144 # word, for which single-flit arbitration is correct.
2145 # `mux_pipeline_levels=2` retimes the (wide -- a full engine word plus
2146 # address) payload selection mux; the added latency is absorbed by the
2147 # arbiter's output FIFO / credit counter.
2148 muxed_write_channel = ChannelArbiter(write_channels,
2149 ports.clk,
2150 ports.rst,
2151 mux_pipeline_levels=2,
2152 pipelined_scheduler=True,
2153 telemetry=False)
2154 upstream_req_channel.assign(muxed_write_channel)
2155
2156 return HostMemWriteProcessorImpl
2157
2158
2159@modparams
2160def ChannelHostMem(
2161 read_width: int,
2162 write_width: int,
2163 max_read_request_bytes: int = DEFAULT_MAX_READ_REQUEST_BYTES,
2164 max_write_payload_bytes: int = DEFAULT_MAX_WRITE_PAYLOAD_BYTES
2165) -> typing.Type['ChannelHostMemImpl']:
2166
2167 class ChannelHostMemImpl(esi.ServiceImplementation):
2168 """Builds a HostMem service which multiplexes multiple HostMem clients into
2169 two (read and write) bundles of the given data width."""
2170
2171 clk = Clock()
2172 rst = Reset()
2173
2174 UpstreamReadReq = StructType([
2175 ("address", UInt(64)),
2176 ("length", UInt(32)), # In bytes.
2177 ("tag", UInt(8)),
2178 ])
2179 read = Output(
2180 Bundle([
2181 BundledChannel("req", ChannelDirection.TO, UpstreamReadReq),
2182 BundledChannel(
2183 "resp", ChannelDirection.FROM,
2184 StructType([
2185 ("tag", esi.HostMem.TagType),
2186 ("data", Bits(read_width)),
2187 ("last", Bits(1)),
2188 ])),
2189 ]))
2190
2191 if write_width % 8 != 0:
2192 raise ValueError("Write width must be a multiple of 8.")
2193 UpstreamWriteReq = StructType([
2194 ("address", UInt(64)),
2195 ("tag", UInt(8)),
2196 ("data", Bits(write_width)),
2197 ("data_size", Bits(clog2(write_width // 8))),
2198 ("last", Bits(1)),
2199 ])
2200 write = Output(
2201 Bundle([
2202 BundledChannel("req", ChannelDirection.TO, UpstreamWriteReq),
2203 BundledChannel("ackTag", ChannelDirection.FROM, UInt(8)),
2204 ]))
2205
2206 @generator
2207 def generate(ports, bundles: esi._ServiceGeneratorBundles):
2208 # Split the read side out into a separate module. Must assign the output
2209 # ports to the clients since we can't service a request in a different
2210 # module.
2211 read_reqs = [
2212 req for req in bundles.to_client_reqs
2213 if req.port in ('read', 'read_list')
2214 ]
2215 read_proc_module = HostmemReadProcessor(read_width, ChannelHostMemImpl,
2216 read_reqs, max_read_request_bytes)
2217 read_proc = read_proc_module(clk=ports.clk, rst=ports.rst)
2218 ports.read = read_proc.upstream
2219 for req in read_reqs:
2220 req.assign(getattr(read_proc, read_proc_module.reqPortMap[req]))
2221
2222 # The write side.
2223 write_reqs = [
2224 req for req in bundles.to_client_reqs if req.port == 'write'
2225 ]
2226 write_proc_module = HostMemWriteProcessor(write_width, ChannelHostMemImpl,
2227 write_reqs,
2228 max_write_payload_bytes)
2229 write_proc = write_proc_module(clk=ports.clk, rst=ports.rst)
2230 ports.write = write_proc.upstream
2231 for req in write_reqs:
2232 req.assign(getattr(write_proc, write_proc_module.reqPortMap[req]))
2233
2234 return ChannelHostMemImpl
2235
2236
2237@modparams
2238def DummyToHostEngine(client_type: Type) -> type['DummyToHostEngineImpl']:
2239 """Create a fake DMA engine which just throws everything away."""
2240
2241 class DummyToHostEngineImpl(esi.EngineModule):
2242
2243 @property
2244 def TypeName(self):
2245 return "DummyToHostEngine"
2246
2247 clk = Clock()
2248 rst = Reset()
2249 input_channel = InputChannel(client_type)
2250
2251 @generator
2252 def build(ports):
2253 pass
2254
2255 return DummyToHostEngineImpl
2256
2257
2258@modparams
2259def DummyFromHostEngine(client_type: Type) -> type['DummyFromHostEngineImpl']:
2260 """Create a fake DMA engine which just never produces messages."""
2261
2262 class DummyFromHostEngineImpl(esi.EngineModule):
2263
2264 @property
2265 def TypeName(self):
2266 return "DummyFromHostEngine"
2267
2268 clk = Clock()
2269 rst = Reset()
2270 output_channel = OutputChannel(client_type)
2271
2272 @generator
2273 def build(ports):
2274 valid = Bits(1)(0)
2275 data = Bits(client_type.bitwidth)(0).bitcast(client_type)
2276 channel, ready = Channel(client_type).wrap(data, valid)
2277 ports.output_channel = channel
2278
2279 return DummyFromHostEngineImpl
2280
2281
2282def _resolve_engine_pair(path: str) -> Tuple[Callable, Callable]:
2283 """Resolve a dotted Python import path to a
2284 `(to_host_engine_gen, from_host_engine_gen)` tuple, used to override the
2285 default engine pair for a specific service request.
2286
2287 The path may point at either:
2288 - a module-level 2-tuple attribute, e.g.
2289 `"mypkg.mymod.MyEnginePair"` where `MyEnginePair` is
2290 `(MyToHost, MyFromHost)`; or
2291 - a zero-arg factory callable returning such a tuple.
2292 """
2293 import importlib
2294 if not isinstance(path, str):
2295 raise TypeError(
2296 "Engine override path must be a dotted 'pkg.mod.attr' string; "
2297 f"got {type(path).__name__}")
2298 module_path, _, attr_path = path.rpartition(".")
2299 if not module_path or not attr_path:
2300 raise ValueError(
2301 "Engine override path must be a dotted 'pkg.mod.attr' string; "
2302 f"got {path!r}")
2303 obj = importlib.import_module(module_path)
2304 for part in attr_path.split("."):
2305 obj = getattr(obj, part)
2306 if callable(obj):
2307 obj = obj()
2308 if not (isinstance(obj, tuple) and len(obj) == 2):
2309 raise TypeError(
2310 f"Engine override {path!r} must resolve to a 2-tuple "
2311 f"(to_host_engine_gen, from_host_engine_gen); got {type(obj).__name__}")
2312 if not (callable(obj[0]) and callable(obj[1])):
2313 raise TypeError(
2314 f"Engine override {path!r} must resolve to a 2-tuple of callables; got "
2315 f"({type(obj[0]).__name__}, {type(obj[1]).__name__})")
2316 return obj
2317
2318
2319def ChannelEngineService(
2320 to_host_engine_gen: Callable,
2321 from_host_engine_gen: Callable) -> type['ChannelEngineService']:
2322 """Returns a channel service implementation which calls
2323 to_host_engine_gen(<client_type>) or from_host_engine_gen(<client_type>) to
2324 generate the to_host and from_host engines for each channel. Does not support
2325 engines which can service multiple clients at once.
2326
2327 Individual service requests may override the default engine pair by passing
2328 `options={"engine": "pkg.mod.attr"}` at the service-request call site (e.g.
2329 `HostComms.some_bundle(AppID(...), options={"engine": "..."})`). The path
2330 is resolved by `_resolve_engine_pair` and must yield a
2331 `(to_host_engine_gen, from_host_engine_gen)` tuple with the same call shape
2332 as the defaults; the override applies to every channel of that request's
2333 bundle.
2334 """
2335
2336 class ChannelEngineService(esi.ServiceImplementation):
2337 """Service implementation which services the clients via a per-channel DMA
2338 engine."""
2339
2340 clk = Clock()
2341 rst = Reset()
2342
2343 @generator
2344 def build(ports, bundles: esi._ServiceGeneratorBundles):
2345 clk = ports.clk
2346 rst = ports.rst
2347
2348 def build_engine_appid(client_appid: List[esi.AppID],
2349 channel_name: str) -> str:
2350 appid_strings = [str(appid) for appid in client_appid]
2351 return f"{'_'.join(appid_strings)}.{channel_name}"
2352
2353 def build_engine(bc: BundledChannel,
2354 bundle_to_host_gen: Callable,
2355 bundle_from_host_gen: Callable,
2356 input_channel=None) -> Type:
2357 idbase = build_engine_appid(bundle.client_name, bc.name)
2358 eng_appid = esi.AppID(idbase)
2359 # DMA engines require at least 1 byte of data; substitute Bits(8)
2360 # for zero-width (void) channel types so the engine never sees a
2361 # zero-length transfer.
2362 engine_client_type = bc.channel.inner_type
2363 is_void = (engine_client_type.bitwidth == 0)
2364 if is_void:
2365 engine_client_type = Bits(8)
2366 if bc.direction == ChannelDirection.FROM:
2367 engine_mod = bundle_to_host_gen(engine_client_type)
2368 else:
2369 engine_mod = bundle_from_host_gen(engine_client_type)
2370 eng_inputs = {
2371 "clk": ports.clk,
2372 "rst": ports.rst,
2373 }
2374 eng_details: Dict[str, object] = {"engine_inst": eng_appid}
2375 if input_channel is not None:
2376 # For void channels, widen the 0-bit input to the 8-bit
2377 # placeholder the engine expects.
2378 if is_void:
2379 input_channel = input_channel.transform(lambda _: Bits(8)(0))
2380 if (engine_mod.input_channel.type.signaling
2381 != input_channel.type.signaling):
2382 input_channel = input_channel.buffer(
2383 clk,
2384 rst,
2385 stages=1,
2386 output_signaling=engine_mod.input_channel.type.signaling)
2387 eng_inputs["input_channel"] = input_channel
2388 if hasattr(engine_mod, "mmio"):
2389 mmio_appid = esi.AppID(idbase + ".mmio")
2390 eng_inputs["mmio"] = esi.MMIO.read_write(mmio_appid)
2391 eng_details["mmio"] = mmio_appid
2392 if hasattr(engine_mod, "hostmem_write"):
2393 eng_inputs["hostmem_write"] = esi.HostMem.write_from_bundle(
2394 esi.AppID(idbase + ".hostmem_write"),
2395 engine_mod.hostmem_write.type,
2396 options=bundle.options)
2397 if hasattr(engine_mod, "hostmem_read"):
2398 eng_inputs["hostmem_read"] = esi.HostMem.read_from_bundle(
2399 esi.AppID(idbase + ".hostmem_read"),
2400 engine_mod.hostmem_read.type,
2401 options=bundle.options)
2402 engine = engine_mod(appid=eng_appid, **eng_inputs)
2403 engine_rec = bundles.emit_engine(engine, details=eng_details)
2404 engine_rec.add_record(bundle, {bc.name: {}})
2405 return engine
2406
2407 for bundle in bundles.to_client_reqs:
2408 # Per-request engine override: if the client's service request carries
2409 # an `"engine"` option, use that engine pair instead of the defaults
2410 # for every channel of this bundle. This is purely a hardware-side
2411 # substitution.
2412 engine_override = bundle.options.get("engine")
2413 if engine_override is None:
2414 bundle_to_host_gen = to_host_engine_gen
2415 bundle_from_host_gen = from_host_engine_gen
2416 else:
2417 bundle_to_host_gen, bundle_from_host_gen = _resolve_engine_pair(
2418 engine_override)
2419
2420 bundle_type = bundle.type
2421 to_channels = {}
2422 # Create a DMA engine for each channel headed TO the client (from the host).
2423 for bc in bundle_type.channels:
2424 if bc.direction == ChannelDirection.TO:
2425 engine = build_engine(bc, bundle_to_host_gen, bundle_from_host_gen)
2426 out_chan = engine.output_channel
2427 # For void channels, narrow the 8-bit placeholder back to 0-bit.
2428 if bc.channel.inner_type.bitwidth == 0:
2429 out_chan = out_chan.transform(lambda _: Bits(0)(0))
2430 to_channels[bc.name] = out_chan
2431
2432 client_bundle_sig, froms = bundle_type.pack(**to_channels)
2433 bundle.assign(client_bundle_sig)
2434
2435 # Create a DMA engine for each channel headed FROM the client (to the host).
2436 for bc in bundle_type.channels:
2437 if bc.direction == ChannelDirection.FROM:
2438 build_engine(bc, bundle_to_host_gen, bundle_from_host_gen,
2439 froms[bc.name])
2440
2441 return ChannelEngineService
return wrap(CMemoryType::get(unwrap(ctx), baseType, numElements))
build_read(ports, int manifest_loc, Dict[int, Tuple[Optional[int], AssignableSignal]] table)
Definition common.py:701
Tuple[Dict[int, Tuple[Optional[int], AssignableSignal]], int] build_table(bundles)
Definition common.py:646
generate(ports, esi._ServiceGeneratorBundles bundles)
Definition common.py:639
bool generate(ports, esi._ServiceGeneratorBundles bundles)
Definition common.py:805
HostmemReadProcessor(int read_width, hostmem_module, List[esi._OutputBundleSetter] reqs, int max_read_request_bytes=DEFAULT_MAX_READ_REQUEST_BYTES)
Definition common.py:1584
type["ChannelDemuxNImpl"] ChannelDemuxN_HalfStage_ReadyBlocking(Type data_type, int num_outs, int next_sel_width)
Definition common.py:175
type["ChannelDemuxTree"] ChannelDemuxTree_HalfStage_ReadyBlocking(Type data_type, int num_outs, int branching_factor_log2)
Definition common.py:268
type["ShiftReadGearboxImpl"] ShiftReadGearbox(int input_bitwidth, int output_bitwidth)
Definition common.py:1208
select_read_gearbox(bool is_list, int input_bitwidth, int output_bitwidth)
Definition common.py:1354
Tuple[Callable, Callable] _resolve_engine_pair(str path)
Definition common.py:2282
type["MMIOPrefixRouterImpl"] MMIOPrefixRouter(Tuple[Tuple[int, Optional[int]],...] regions)
Definition common.py:378
type["SliceReadGearboxImpl"] SliceReadGearbox(int input_bitwidth, int output_bitwidth)
Definition common.py:958
Module HeaderMMIO(int manifest_loc)
Definition common.py:74
type["ConcatReadGearboxImpl"] ConcatReadGearbox(int input_bitwidth, int output_bitwidth)
Definition common.py:1008
type[ 'DummyToHostEngineImpl'] DummyToHostEngine(Type client_type)
Definition common.py:2238
type[ 'DummyFromHostEngineImpl'] DummyFromHostEngine(Type client_type)
Definition common.py:2259
type["TaggedWriteGearboxImpl"] TaggedWriteGearbox(int input_bitwidth, int output_bitwidth, int max_burst_bytes)
Definition common.py:1775
type["DepackReadGearboxImpl"] DepackReadGearbox(int input_bitwidth, int output_bitwidth)
Definition common.py:1113
type[ 'EmitEveryNImpl'] EmitEveryN(Type message_type, int N)
Definition common.py:1927
type["HostMemWriteProcessorImpl"] HostMemWriteProcessor(int write_width, hostmem_module, List[esi._OutputBundleSetter] reqs, int max_write_payload_bytes=DEFAULT_MAX_WRITE_PAYLOAD_BYTES)
Definition common.py:1984
HostMemReadReqSplitter(Channel req_channel_type, Channel resp_channel_type, int max_chunk_bytes)
Definition common.py:1391