From 3134f0e4c5231d8f17aa56ecc73bdc5a01f61bb4 Mon Sep 17 00:00:00 2001 From: Jeremy Lau <30300826+fdxmw@users.noreply.github.com> Date: Fri, 17 Jul 2026 16:10:50 -0700 Subject: [PATCH 1/3] Miscellaneous `CompiledSimulation` and `Simulation` cleanup, for 8% speedup (pyrtlnet CompiledSimulation startup): * Refactor to remove redundant code. * Use f-strings instead of `.format`. * Use `_limb_name` (formerly `_getarglimb`) consistently. * Define a `dataclass` for `_build_concat`'s `pieces`. * Define `dataclass`es for `Input` and `Output` metadata. * Rename a lot of variables to improve readability. * Add comments, type annotations, and assertions. * Add tests. --- pyrtl/compilesim.py | 1236 +++++++++++++++++++++++--------------- pyrtl/simulation.py | 9 +- tests/test_compilesim.py | 161 ++++- 3 files changed, 903 insertions(+), 503 deletions(-) diff --git a/pyrtl/compilesim.py b/pyrtl/compilesim.py index e7022dfc..6aa690d4 100644 --- a/pyrtl/compilesim.py +++ b/pyrtl/compilesim.py @@ -7,7 +7,8 @@ import subprocess import tempfile import warnings -from collections.abc import Mapping +from collections.abc import Callable, Mapping +from dataclasses import dataclass from os import path from typing import TextIO @@ -31,37 +32,36 @@ class DllMemInspector(Mapping): """Dictionary-like access to a hashmap in a CompiledSimulation.""" def __init__(self, sim, mem): - self._aw = mem.addrwidth + self._addr_width = mem.addrwidth self._limbs = sim._limbs(mem) - self._vn = vn = sim.varname[mem] - self._mem = ctypes.c_void_p.in_dll(sim._dll, vn) + self._var_name = var_name = sim.var_names[mem] + self._mem = ctypes.c_void_p.in_dll(sim._dll, var_name) self._sim = sim # keep reference to avoid freeing dll - def __getitem__(self, ind): - arr = self._sim._mem_lookup(self._mem, ind) - val = 0 - for n in reversed(range(self._limbs)): - val <<= 64 - val |= arr[n] - return val + def __getitem__(self, index): + array = self._sim._mem_lookup(self._mem, index) + value = 0 + for limb in reversed(range(self._limbs)): + value = (value << 64) | array[limb] + return value def __iter__(self): return iter(range(len(self))) def __len__(self): - return 1 << self._aw + return 1 << self._addr_width def __eq__(self, other): if ( isinstance(other, DllMemInspector) and self._sim is other._sim - and self._vn == other._vn + and self._var_name == other._var_name ): return True return all(self[x] == other.get(x, 0) for x in self) def __hash__(self): - return hash(self._sim) ^ hash(self._vn) + return hash(self._sim) ^ hash(self._var_name) class CompiledSimulation: @@ -136,16 +136,16 @@ def __init__( self._remove_untraceable() self.default_value = default_value - self._regmap = {} # Updated below - self._memmap = memory_value_map + self._register_value_map = {} # Updated below + self._memory_value_map = memory_value_map self._uid_counter = 0 - self.varname = {} # mapping from wires and memories to C variables + self.var_names = {} # Map from WireVectors and MemBlocks to C variable names. - for r in self.block.wirevector_subset(Register): - rval = register_value_map.get(r, r.reset_value) - if rval is None: - rval = self.default_value - self._regmap[r] = rval + for reg in self.block.wirevector_subset(Register): + reset_value = register_value_map.get(reg, reg.reset_value) + if reset_value is None: + reset_value = self.default_value + self._register_value_map[reg] = reset_value self.tracer._set_initial_values( default_value, register_value_map, memory_value_map @@ -157,11 +157,11 @@ def __init__( def inspect_mem(self, mem: MemBlock) -> dict[int, int]: return DllMemInspector(self, mem) - def inspect(self, w: str) -> int: - if isinstance(w, WireVector): - w = w.name + def inspect(self, wire_name: str) -> int: + if isinstance(wire_name, WireVector): + wire_name = wire_name.name try: - vals = self.tracer.trace[w] + vals = self.tracer.trace[wire_name] except KeyError: pass else: @@ -204,16 +204,16 @@ def step_multiple( longest = sorted( provided_inputs.items(), key=lambda t: len(t[1]), reverse=True )[0] - msteps = len(longest[1]) + provided_steps = len(longest[1]) if nsteps: - if nsteps > msteps: + if nsteps > provided_steps: msg = ( "nsteps is specified but is greater than the number of values " "supplied for each input" ) raise PyrtlError(msg) else: - nsteps = msteps + nsteps = provided_steps if nsteps < 1: msg = "must simulate at least one step" @@ -233,35 +233,37 @@ def step_multiple( raise PyrtlError(msg) failed = [] - for i in range(nsteps): - self.step({w: int(v[i]) for w, v in provided_inputs.items()}) + for step in range(nsteps): + self.step( + {name: int(value[step]) for name, value in provided_inputs.items()} + ) for expvar in expected_outputs: - expected = expected_outputs[expvar][i] + expected = expected_outputs[expvar][step] if expected == "?": continue expected = int(expected) actual = self.inspect(expvar) if expected != actual: - failed.append((i, expvar, expected, actual)) + failed.append((step, expvar, expected, actual)) if failed and stop_after_first_error: break if failed: if stop_after_first_error: - s = "(stopped after step with first error):" + suffix = "(stopped after step with first error):" else: - s = "on one or more steps:" - print("Unexpected output " + s, file=file) + suffix = "on one or more steps:" + print("Unexpected output " + suffix, file=file) print(f"{'step':>5} {'name':>10} {'expected':>8} {'actual':>8}", file=file) def _sort_tuple(t): # Sort by step and then wire name return (t[0], _trace_sort_key(t[1])) - failed_sorted = sorted(failed, key=_sort_tuple) - for step, name, expected, actual in failed_sorted: + failed = sorted(failed, key=_sort_tuple) + for step, name, expected, actual in failed: print(f"{step:>5} {name:>10} {expected:>8} {actual:>8}", file=file) def run(self, inputs: list[dict[str, int]]): @@ -274,70 +276,79 @@ def run(self, inputs: list[dict[str, int]]): of steps to be executed. """ steps = len(inputs) - # create i/o arrays of the appropriate length - ibuf_type = ctypes.c_uint64 * (steps * self._ibufsz) - obuf_type = ctypes.c_uint64 * (steps * self._obufsz) - ibuf = ibuf_type() - obuf = obuf_type() - # these array will be passed to _crun - self._crun.argtypes = [ctypes.c_uint64, ibuf_type, obuf_type] - - # build the input array - for n, inmap in enumerate(inputs): - for w in inmap: - if isinstance(w, WireVector): - name = w.name - else: - name = w - start, count = self._inputpos[name] - start += n * self._ibufsz - val = inmap[w] - val = infer_val_and_bitwidth(val, bitwidth=self._inputbw[name]).value - # pack input - for pos in range(start, start + count): - ibuf[pos] = val & ((1 << 64) - 1) - val >>= 64 - - # run the simulation - self._crun(steps, ibuf, obuf) - - # save traced wires - for name in self.tracer.trace: - rname = self._probe_mapping.get(name, name) - if rname in self._outputpos: - start, count = self._outputpos[rname] - buf, sz = obuf, self._obufsz - elif rname in self._inputpos: - start, count = self._inputpos[rname] - buf, sz = ibuf, self._ibufsz + # Create arrays for Inputs and Outputs. + inputs_array_type = ctypes.c_uint64 * (steps * self._inputs_array_length) + outputs_array_type = ctypes.c_uint64 * (steps * self._outputs_array_length) + inputs_array = inputs_array_type() + outputs_array = outputs_array_type() + # These arrays will be passed to `_sim_run_all`. + self._sim_run_all.argtypes = [ + ctypes.c_uint64, + inputs_array_type, + outputs_array_type, + ] + + # Build `inputs_array` from `inputs`. `inputs` is a list of `provided_inputs` + # for each step. Each `provided_input` is a map from `WireVector` name to value. + mask = (1 << 64) - 1 + for step, provided_inputs in enumerate(inputs): + for input_name, input_value in provided_inputs.items(): + if isinstance(input_name, WireVector): + input_name = input_name.name + # Figure out where to store `input_value` in `inputs_array` for this + # `step`. + input_metadata = self._inputs_metadata[input_name] + start = input_metadata.start + step * self._inputs_array_length + input_value = infer_val_and_bitwidth( + input_value, bitwidth=input_metadata.bitwidth + ).value + # Pack `input_value` into 64-bit limbs. + for pos in range(start, start + input_metadata.length): + inputs_array[pos] = input_value & mask + input_value = input_value >> 64 + + # Run the simulation with `inputs_array` and `outputs_array`. This invokes the + # compiled C `sim_run_all` function. + self._sim_run_all(steps, inputs_array, outputs_array) + + # Copy values for any traced wires from `outputs_array` and `inputs_array` to + # the SimulationTrace. + for name, values in self.tracer.trace.items(): + actual_name = self._probe_mapping.get(name, name) + if actual_name in self._outputs_metadata: + metadata = self._outputs_metadata[actual_name] + array = outputs_array + array_length = self._outputs_array_length + elif actual_name in self._inputs_metadata: + metadata = self._inputs_metadata[actual_name] + array = inputs_array + array_length = self._inputs_array_length else: msg = "Untraceable wire in tracer" raise PyrtlInternalError(msg) - res = [] + for _step in range(steps): - val = 0 - # unpack output - for pos in reversed(range(start, start + count)): - val <<= 64 - val |= buf[pos] - res.append(val) - start += sz - self.tracer.trace[name].extend(res) - - def _traceable(self, wv): - """Check if wv is able to be traced. + value = 0 + # Unpack output from 64-bit limbs. + start = metadata.start + end = metadata.start + metadata.length + for index in reversed(range(start, end)): + value = (value << 64) | array[index] + values.append(value) + start += array_length + + def _traceable(self, wire: WireVector) -> bool: + """Check if wire is able to be traced. If it is traceable due to a probe, record that probe in _probe_mapping. """ - if isinstance(wv, (Input, Output)): + if isinstance(wire, (Input, Output)): return True - for net in self.block.logic: - if ( - net.op == "w" - and net.args[0].name == wv.name - and isinstance(net.dests[0], Output) + for wire_net in self.block.logic_subset("w"): + if wire_net.args[0].name == wire.name and isinstance( + wire_net.dests[0], Output ): - self._probe_mapping[wv.name] = net.dests[0].name + self._probe_mapping[wire.name] = wire_net.dests[0].name return True return False @@ -347,16 +358,16 @@ def _remove_untraceable(self): Create _probe_mapping for wires only traceable via probes. """ self._probe_mapping = {} - wvs = {wv for wv in self.tracer.wires_to_track if self._traceable(wv)} - self.tracer.wires_to_track = wvs - self.tracer._wires = {wv.name: wv for wv in wvs} - self.tracer.trace.__init__(wvs) + wires = {wire for wire in self.tracer.wires_to_track if self._traceable(wire)} + self.tracer.wires_to_track = wires + self.tracer._wires = {wire.name: wire for wire in wires} + self.tracer.trace.__init__(wires) def _create_dll(self): """Create a dynamically-linked library implementing the simulation logic.""" self._dir = tempfile.mkdtemp() with open(path.join(self._dir, "pyrtlsim.c"), "w") as f: - self._create_code(lambda s: f.write(s + "\n")) + self._create_code(lambda s: f.write(f"{s}\n")) if platform.system() == "Darwin": shared = "-dynamiclib" march = "" @@ -380,312 +391,527 @@ def _create_dll(self): shell=(platform.system() == "Windows"), ) self._dll = ctypes.CDLL(path.join(self._dir, "pyrtlsim.so")) - self._crun = self._dll.sim_run_all - self._crun.restype = None # argtypes set on use + self._sim_run_all = self._dll.sim_run_all + self._sim_run_all.restype = None # argtypes set on use self._initialize_mems = self._dll.initialize_mems self._initialize_mems.restype = None self._mem_lookup = self._dll.lookup self._mem_lookup.restype = ctypes.POINTER(ctypes.c_uint64) - def _limbs(self, w): - """Number of 64-bit words needed to store value of wire.""" - return (w.bitwidth + 63) // 64 - - def _makeini(self, w, v): - """C initializer string for a wire with a given value.""" - pieces = [] - for _ in range(self._limbs(w)): - pieces.append(hex(v & ((1 << 64) - 1))) - v >>= 64 - return ",".join(pieces).join("{}") + def _limbs(self, wire: WireVector) -> int: + """Number of 64-bit words needed to store a WireVector's value.""" + return (wire.bitwidth + 63) // 64 - def _romwidth(self, m): - """Bitwidth of integer type sufficient to hold rom entry. + def _make_initializer(self, wire: WireVector, value: int): + """Return a C initializer string that initializes ``wire``'s limbs to ``value``. - On large memories, returns 64; an array will be needed. + For example, if the value 0x4444_3333_2222_1111_0000 (80 bits) is assigned to a + ``wire`` with two limbs, this returns "{0x4444, 0x3333222211110000}". """ - if m.bitwidth <= 8: + values = [] + mask = (1 << 64) - 1 + for _ in range(self._limbs(wire)): + values.append(hex(value & mask)) + value = value >> 64 + values_str = ",".join(values) + return f"{{{values_str}}}" + + def _rom_width(self, rom: RomBlock) -> int: + """Return the bitwidth of an integer type sufficient to hold an element of + ``rom``. + + Returns 64 for large memories; an array will be needed. + """ + if rom.bitwidth <= 8: return 8 - if m.bitwidth <= 16: + if rom.bitwidth <= 16: return 16 - if m.bitwidth <= 32: + if rom.bitwidth <= 32: return 32 return 64 - def _makemask(self, dest, res, pos): - """Create a bitmask. + def _make_mask( + self, dest: WireVector, value_bitwidth: int | None, dest_limb: int + ) -> str: + """Return an expression that applies a bitmask to the value assigned to + ``dest``, if necessary. - The value being masked is of width `res`. Limb number `pos` of `dest` is being - assigned to. + The value has width ``value_bitwidth``, and we are assigning to limb number + ``dest_limb`` of ``dest``. """ - if (res is None or dest.bitwidth < res) and 0 < (dest.bitwidth - 64 * pos) < 64: + # We must bitmask `value` if: + # 1. `value_needs_mask`: The value has high bits that will not be copied to + # `dest`, AND + # 2. `partial_dest_limb`: We are assigning fewer than 64 bits to the `dest` + # limb. + value_needs_mask = value_bitwidth is None or dest.bitwidth < value_bitwidth + partial_dest_limb = 0 < (dest.bitwidth - 64 * dest_limb) < 64 + if value_needs_mask and partial_dest_limb: return f"&0x{(1 << (dest.bitwidth % 64)) - 1:X}" return "" - def _getarglimb(self, arg, n): - """Get the nth limb of the given wire. + def _limb_name(self, wire: WireVector, wire_limb: int) -> str: + """Return the name of the ``wire_limb``'th limb of ``wire``. - Returns '0' when the wire does not have sufficient limbs. + Returns "0" when ``wire`` does not have sufficient limbs. This is useful when + generating code for comparisons, see ``_build_eq`` and ``_build_cmp``. """ - return f"{self.varname[arg]}[{n}]" if arg.bitwidth > 64 * n else "0" + if wire.bitwidth > 64 * wire_limb: + return f"{self.var_names[wire]}[{wire_limb}]" + return "0" - def _clean_name(self, prefix, obj): - """Create a C variable name with the given prefix based on the name of obj.""" - return "{}{}_{}".format( - prefix, self._uid(), "".join(c for c in obj.name if c.isalnum()) - ) + def _clean_name(self, prefix: str, obj: WireVector | MemBlock) -> str: + """Return a C variable name with the given ``prefix`` based on the name of + ``obj``. + """ + suffix = "".join(char for char in obj.name if char.isalnum()) + return f"{prefix}{self._uid()}_{suffix}" - def _uid(self): - """Get an auto-incrementing number suitable for use as a unique identifier.""" + def _uid(self) -> int: + """Postincrement a counter. The returned number can be used as a unique + identifier. + """ x = self._uid_counter self._uid_counter += 1 return x - def _declare_roms(self, write, roms): - for mem in roms: - self.varname[mem] = vn = self._clean_name("m", mem) - # extract data from mem - romval = [mem._get_read_data(n) for n in range(1 << mem.addrwidth)] + def _declare_roms(self, write: Callable, roms: set[RomBlock]): + for rom in roms: + self.var_names[rom] = var_name = self._clean_name("m", rom) + # Make initialization strings for each of the ROM's elements. + rom_initializers = [] + for addr in range(1 << rom.addrwidth): + rom_initializers.append( + self._make_initializer(rom, rom._get_read_data(addr)) + ) + rom_initializer = ",".join(rom_initializers) write( - f"static const uint{self._romwidth(mem)}_t {vn}[]" - f"[{self._limbs(mem)}] = {{" + f"static const uint{self._rom_width(rom)}_t {var_name}[]" + f"[{self._limbs(rom)}] = {{{rom_initializer}}};" ) - for rv in romval: - write(self._makeini(mem, rv) + ",") - write("};") - def _declare_mems(self, write, mems): + def _declare_mems(self, write: Callable, mems: set[MemBlock]): for mem in mems: - self.varname[mem] = vn = self._clean_name("m", mem) + self.var_names[mem] = var_name = self._clean_name("m", mem) write("EXPORT") - write(f"hashmap_t *{vn};") + write(f"hashmap_t *{var_name};") next_tmp = 0 write("EXPORT") write("void initialize_mems() {") for mem in mems: # Create hashmap - write(f"{self.varname[mem]} = create_hash_map(256, {self._limbs(mem)});") - if mem in self._memmap: - # Insert default values - for k, v in self._memmap[mem].items(): - write(f"val_t t{next_tmp}[] = {self._makeini(mem, v)};") - write(f"insert({self.varname[mem]}, {k}, t{next_tmp});") - next_tmp += 1 + write(f"{self.var_names[mem]} = create_hash_map(256, {self._limbs(mem)});") + values = self._memory_value_map.get(mem) + if values is None: + continue + # Insert default values + for addr, value in values.items(): + write(f"val_t t{next_tmp}[] = {self._make_initializer(mem, value)};") + write(f"insert({self.var_names[mem]}, {addr}, t{next_tmp});") + next_tmp += 1 write("}") - def _declare_wv(self, write, w): - self.varname[w] = vn = self._clean_name("w", w) - if isinstance(w, Const): - write(f"const uint64_t {vn}[{self._limbs(w)}] = {self._makeini(w, w.val)};") - elif isinstance(w, Register): - rval = self._regmap.get(w, w.reset_value) - if rval is None: - rval = self.default_value - write(f"static uint64_t {vn}[{self._limbs(w)}] = {self._makeini(w, rval)};") + def _declare_wirevector(self, write: Callable, wire: WireVector): + self.var_names[wire] = var_name = self._clean_name("w", wire) + wire_name = f"{var_name}[{self._limbs(wire)}]" + if isinstance(wire, Const): + write( + f"const uint64_t {wire_name} = " + f"{self._make_initializer(wire, wire.val)};" + ) + elif isinstance(wire, Register): + reset_value = self._register_value_map[wire] + write( + f"static uint64_t {wire_name} = " + f"{self._make_initializer(wire, reset_value)};" + ) else: - write(f"uint64_t {vn}[{self._limbs(w)}];") + write(f"uint64_t {wire_name};") - def _build_memread(self, write, _op, param, args, dest): - mem = param[1] - for n in range(self._limbs(dest)): + def _build_mem_read( + self, + write: Callable, + op: str, + op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "m": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) + mem = op_param[1] + for dest_limb in range(self._limbs(dest)): + dest_name = self._limb_name(dest, dest_limb) + # This code assumes read addresses are 64-bits or less. + if self._limbs(args[0]) > 1: + msg = "Read addresses longer than 64 bits are not supported." + raise PyrtlInternalError(msg) + read_addr = self._limb_name(args[0], 0) + mask = self._make_mask(dest, mem.bitwidth, dest_limb) if isinstance(mem, RomBlock): - write( - f"{self.varname[dest]}[{n}] = {self.varname[mem]}[" - f"{self.varname[args[0]]}[0]][{n}]" - f"{self._makemask(dest, mem.bitwidth, n)};" - ) + expr = f"{self.var_names[mem]}[{read_addr}][{dest_limb}]{mask};" else: - write( - f"{self.varname[dest]}[{n}] = lookup({self.varname[mem]}, " - f"{self.varname[args[0]]}[0])[{n}]" - f"{self._makemask(dest, mem.bitwidth, n)};" + expr = f"lookup({self.var_names[mem]}, {read_addr})[{dest_limb}]{mask};" + write(f"{dest_name} = {expr}") + + def _write_assignments( + self, + write: Callable, + args: list[WireVector], + dest: WireVector, + arg_index: int, + indent: int = 0, + ) -> str: + """Assign each limb of `args[arg_index]` to each limb of `dest`.""" + for dest_limb in range(self._limbs(dest)): + dest_name = self._limb_name(dest, dest_limb) + arg_name = self._limb_name(args[arg_index], dest_limb) + # Assignment copies all the arg's bits, so we don't need a mask. + if dest.bitwidth != args[arg_index].bitwidth: + msg = ( + f"Assignment dest {dest} and arg {args[arg_index]} bitwidths do " + "not match" ) + raise PyrtlInternalError(msg) + write(f"{' ' * indent}{dest_name} = {arg_name};") - def _build_wire(self, write, _op, _param, args, dest): - for n in range(self._limbs(dest)): - write( - f"{self.varname[dest]}[{n}] = " - f"{self.varname[args[0]]}[{n}]" - f"{self._makemask(dest, args[0].bitwidth, n)};" - ) + def _build_wire( + self, + write: Callable, + op: str, + _op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "w": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) + self._write_assignments(write, args, dest, arg_index=0) - def _build_not(self, write, _op, _param, args, dest): - for n in range(self._limbs(dest)): - write( - f"{self.varname[dest]}[{n}] = " - f"(~{self.varname[args[0]]}[{n}]){self._makemask(dest, None, n)};" - ) + def _build_not( + self, + write: Callable, + op: str, + _op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "~": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) - def _build_bitwise(self, write, op, _param, args, dest): # &, |, ^ only - for n in range(self._limbs(dest)): - arg0 = self._getarglimb(args[0], n) - arg1 = self._getarglimb(args[1], n) - write( - "{dest}[{n}] = ({arg0}{op}{arg1}){mask};".format( - dest=self.varname[dest], - n=n, - arg0=arg0, - arg1=arg1, - op=op, - mask=self._makemask( - dest, max(args[0].bitwidth, args[1].bitwidth), n - ), - ) - ) + for dest_limb in range(self._limbs(dest)): + dest_name = self._limb_name(dest, dest_limb) + arg0_name = self._limb_name(args[0], dest_limb) + mask = self._make_mask(dest, None, dest_limb) + write(f"{dest_name} = (~{arg0_name}){mask};") - def _build_nand(self, write, _op, _param, args, dest): - for n in range(self._limbs(dest)): - arg0 = self._getarglimb(args[0], n) - arg1 = self._getarglimb(args[1], n) - write( - f"{self.varname[dest]}[{n}] = " - f"(~({arg0}&{arg1})){self._makemask(dest, None, n)};" + def _build_bitwise( + self, + write: Callable, + op: str, + _op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "&" and op != "|" and op != "^": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) + + for dest_limb in range(self._limbs(dest)): + dest_name = self._limb_name(dest, dest_limb) + arg0_name = self._limb_name(args[0], dest_limb) + arg1_name = self._limb_name(args[1], dest_limb) + mask = self._make_mask( + dest, max(args[0].bitwidth, args[1].bitwidth), dest_limb ) + write(f"{dest_name} = ({arg0_name}{op}{arg1_name}){mask};") - def _build_eq(self, write, _op, _param, args, dest): - cond = [] - for n in range(max(self._limbs(args[0]), self._limbs(args[1]))): - arg0 = self._getarglimb(args[0], n) - arg1 = self._getarglimb(args[1], n) - cond.append(f"({arg0}=={arg1})") - write( - "{dest}[0] = {cond};".format(dest=self.varname[dest], cond="&&".join(cond)) - ) + def _build_nand( + self, + write: Callable, + op: str, + _op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "n": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) + + for dest_limb in range(self._limbs(dest)): + dest_name = self._limb_name(dest, dest_limb) + arg0_name = self._limb_name(args[0], dest_limb) + arg1_name = self._limb_name(args[1], dest_limb) + mask = self._make_mask(dest, None, dest_limb) + write(f"{dest_name} = (~({arg0_name}&{arg1_name})){mask};") + + def _build_eq( + self, + write: Callable, + op: str, + _op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "=": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) - def _build_cmp(self, write, op, _param, args, dest): # <, > only - cond = None - for n in range(max(self._limbs(args[0]), self._limbs(args[1]))): - arg0 = self._getarglimb(args[0], n) - arg1 = self._getarglimb(args[1], n) - c = f"({arg0}{op}{arg1})" - if cond is None: - cond = c + condition_parts = [] + for arg_limb in range(max(self._limbs(args[0]), self._limbs(args[1]))): + arg0_name = self._limb_name(args[0], arg_limb) + arg1_name = self._limb_name(args[1], arg_limb) + condition_parts.append(f"({arg0_name}=={arg1_name})") + + dest_name = self._limb_name(dest, 0) + condition = "&&".join(condition_parts) + write(f"{dest_name} = {condition};") + + def _build_lt_gt( + self, + write: Callable, + op: str, + _op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "<" and op != ">": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) + + condition = None + for arg_limb in range(max(self._limbs(args[0]), self._limbs(args[1]))): + arg0_name = self._limb_name(args[0], arg_limb) + arg1_name = self._limb_name(args[1], arg_limb) + comparison = f"({arg0_name}{op}{arg1_name})" + if condition is None: + condition = comparison else: - cond = f"({c}||(({arg0}=={arg1})&&{cond}))" - write(f"{self.varname[dest]}[0] = {cond};") + condition = f"({comparison}||(({arg0_name}=={arg1_name})&&{condition}))" - def _build_mux(self, write, _op, _param, args, dest): - write(f"if ({self.varname[args[0]]}[0]) {{") - for n in range(self._limbs(dest)): - write( - f"{self.varname[dest]}[{n}] = " - f"{self.varname[args[2]]}[{n}]" - f"{self._makemask(dest, args[2].bitwidth, n)};" - ) + dest_name = self._limb_name(dest, 0) + write(f"{dest_name} = {condition};") + + def _build_mux( + self, + write: Callable, + op: str, + _op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "x": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) + + arg0_name = self._limb_name(args[0], 0) + write(f"if ({arg0_name}) {{") + self._write_assignments(write, args, dest, arg_index=2, indent=2) write("} else {") - for n in range(self._limbs(dest)): - write( - f"{self.varname[dest]}[{n}] = " - f"{self.varname[args[1]]}[{n}]" - f"{self._makemask(dest, args[1].bitwidth, n)};" - ) + self._write_assignments(write, args, dest, arg_index=1, indent=2) write("}") - def _build_add(self, write, _op, _param, args, dest): + def _build_add( + self, + write: Callable, + op: str, + _op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "+": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) + write("carry = 0;") - for n in range(self._limbs(dest)): - arg0 = self._getarglimb(args[0], n) - arg1 = self._getarglimb(args[1], n) - write(f"tmp = {arg0}+{arg1};") - write( - "{dest}[{n}] = (tmp + carry){mask};".format( - dest=self.varname[dest], - n=n, - mask=self._makemask( - dest, max(args[0].bitwidth, args[1].bitwidth) + 1, n - ), - ) + for dest_limb in range(self._limbs(dest)): + arg0_name = self._limb_name(args[0], dest_limb) + arg1_name = self._limb_name(args[1], dest_limb) + write(f"tmp = {arg0_name}+{arg1_name};") + + dest_name = self._limb_name(dest, dest_limb) + mask = self._make_mask( + dest, max(args[0].bitwidth, args[1].bitwidth) + 1, dest_limb ) - write(f"carry = (tmp < {arg0})|({self.varname[dest]}[{n}] < tmp);") + write(f"{dest_name} = (tmp + carry){mask};") + write(f"carry = (tmp < {arg0_name})|({dest_name} < tmp);") + + def _build_sub( + self, + write: Callable, + op: str, + _op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "-": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) - def _build_sub(self, write, _op, _param, args, dest): write("carry = 0;") - for n in range(self._limbs(dest)): - arg0 = self._getarglimb(args[0], n) - arg1 = self._getarglimb(args[1], n) - write(f"tmp = {arg0}-{arg1};") - write( - f"{self.varname[dest]}[{n}] = (tmp - carry)" - f"{self._makemask(dest, None, n)};" - ) - write(f"carry = (tmp > {arg0})|({self.varname[dest]}[{n}] > tmp);") + for dest_limb in range(self._limbs(dest)): + arg0_name = self._limb_name(args[0], dest_limb) + arg1_name = self._limb_name(args[1], dest_limb) + write(f"tmp = {arg0_name}-{arg1_name};") - def _build_mul(self, write, _op, _param, args, dest): - for n in range(self._limbs(dest)): - write(f"{self.varname[dest]}[{n}] = 0;") - for p0 in range(self._limbs(args[0])): + dest_name = self._limb_name(dest, dest_limb) + mask = self._make_mask(dest, None, dest_limb) + write(f"{dest_name} = (tmp - carry){mask};") + write(f"carry = (tmp > {arg0_name})|({dest_name} > tmp);") + + def _build_mul( + self, + write: Callable, + op: str, + _op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "*": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) + + total_arg_bitwidth = args[0].bitwidth + args[1].bitwidth + + for dest_limb in range(self._limbs(dest)): + dest_name = self._limb_name(dest, dest_limb) + write(f"{dest_name} = 0;") + + # Roughly, this computes (ignoring carries): + # + # for arg0_limb in self._limbs(args[0]): + # for arg1_limb in self._limbs(args[1]): + # dest[arg0_limb + arg1_limb] += args[0][arg0_limb] * args[1][arg1_limb] + for arg0_limb in range(self._limbs(args[0])): write("carry = 0;") - arg0 = self._getarglimb(args[0], p0) - for p1 in range(self._limbs(args[1])): - if self._limbs(dest) <= p0 + p1: - break - arg1 = self._getarglimb(args[1], p1) - write(f"mul128({arg0}, {arg1}, tmplo, tmphi);") - write(f"tmp = {self.varname[dest]}[{p0 + p1}];") - write("tmplo += carry; carry = tmplo < carry; tmplo += tmp;") - write("tmphi += carry + (tmplo < tmp); carry = tmphi;") - write( - "{dest}[{p}] = tmplo{mask};".format( - dest=self.varname[dest], - p=p0 + p1, - mask=self._makemask( - dest, args[0].bitwidth + args[1].bitwidth, p0 + p1 - ), - ) - ) - if self._limbs(dest) > p0 + self._limbs(args[1]): - write( - "{dest}[{p}] = carry{mask};".format( - dest=self.varname[dest], - p=p0 + self._limbs(args[1]), - mask=self._makemask( - dest, - args[0].bitwidth + args[1].bitwidth, - p0 + self._limbs(args[1]), - ), + arg0_name = self._limb_name(args[0], arg0_limb) + + for arg1_limb in range(self._limbs(args[1])): + arg1_name = self._limb_name(args[1], arg1_limb) + write(f"mul128({arg0_name}, {arg1_name}, tmplo, tmphi);") + + dest_limb = arg0_limb + arg1_limb + if dest_limb >= self._limbs(dest): + msg = ( + f"Insufficient dest limbs ({self._limbs(dest)}) for arg limbs " + f"{arg0_limb}, {arg1_limb}" ) - ) + raise PyrtlInternalError(msg) + dest_name = self._limb_name(dest, dest_limb) + write(f"tmplo += carry; carry = tmplo < carry; tmplo += {dest_name};") + write(f"tmphi += carry + (tmplo < {dest_name}); carry = tmphi;") + + mask = self._make_mask(dest, total_arg_bitwidth, dest_limb) + write(f"{dest_name} = tmplo{mask};") + + # Finished multiplying arg1 with one limb of arg0. Write any remaining carry + # bits and move on to the next most significant limb of arg0. + dest_carry_limb = arg0_limb + self._limbs(args[1]) + if dest_carry_limb < self._limbs(dest): + dest_name = self._limb_name(dest, dest_carry_limb) + mask = self._make_mask(dest, total_arg_bitwidth, dest_carry_limb) + write(f"{dest_name} = carry{mask};") + + def _build_concat( + self, + write: Callable, + op: str, + _op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "c": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) + + total_arg_bitwidth = sum(x.bitwidth for x in args) - def _build_concat(self, write, _op, _param, args, dest): - cattotal = sum(x.bitwidth for x in args) - pieces = ( - (self.varname[a], lx, 0, min(64, a.bitwidth - 64 * lx)) - for a in reversed(args) - for lx in range(self._limbs(a)) + @dataclass + class ArgLimbPiece: + """Tracks sources for ``concat``'s ``arg`` bits. + + ``ArgLimbPiece`` represents a consecutive range of an ``arg``'s bits within + an arg limb. + """ + + name: str + """C variable name for this arg limb.""" + + start: int + """Starting offset within this arg limb. Must be in the range [0, 64).""" + + length: int + """The arg's remaining length within this arg limb. Must be in the range [1, + 64). + """ + + arg_limb_pieces = ( + ArgLimbPiece( + name=self._limb_name(arg, arg_limb), + start=0, + length=min(64, arg.bitwidth - 64 * arg_limb), + ) + for arg in reversed(args) + for arg_limb in range(self._limbs(arg)) ) - curr = next(pieces) - for n in range(self._limbs(dest)): - res = [] - dpos = 0 - while True: - arg, alimb, astart, asize = curr - res.append(f"(({arg}[{alimb}]>>{astart})<<{dpos})") - dpos += asize - if dpos > 64: - curr = (arg, alimb, 64 - (dpos - asize), dpos - 64) - break - if dpos >= dest.bitwidth - 64 * n: + + arg_limb_piece = next(arg_limb_pieces) + for dest_limb in range(self._limbs(dest)): + expr_parts = [] + dest_start = 0 + + # Keep copying `arg` bits until this `dest` limb is full, or we're out of + # `args`. + while dest_start < 64: + if arg_limb_piece.start >= 64: + msg = ( + f"arg_limb_piece.start {arg_limb_piece.start} exceeds limb " + "bitwidth" + ) + raise PyrtlInternalError(msg) + + # Copy bits from the current `arg` limb to the current `dest` limb. The + # second shift may throw away some of `arg` limb's bits, but we will + # track that below, and copy any thrown-away `arg` limb bits to the next + # `dest` limb. + expr = shift(arg_limb_piece.name, ">>", arg_limb_piece.start) + expr = shift(expr, "<<", dest_start) + # Concat copies all bits from each arg, so we don't need to mask args. + expr_parts.append(expr) + + dest_start += arg_limb_piece.length + if dest_start > 64: + # We've reached the end of the current `dest` limb, but the current + # `arg` limb still has more bits. The `arg` limb's remaining bits + # must be copied to the next `dest` limb. Update `arg_limb_piece` so + # we remember where to resume copying from this `arg` limb. + arg_limb_piece = ArgLimbPiece( + name=arg_limb_piece.name, + start=64 - (dest_start - arg_limb_piece.length), + length=dest_start - 64, + ) break - curr = next(pieces) - if dpos == 64: + if dest_start >= dest.bitwidth - 64 * dest_limb: + # Done with the last `arg`. break - write( - "{dest}[{n}] = ({res}){mask};".format( - dest=self.varname[dest], - n=n, - res="|".join(res), - mask=self._makemask(dest, cattotal, n), - ) - ) + # We finished the current `arg` limb, but the current `dest` limb still + # has room for more bits. Move on to the next `arg`. + arg_limb_piece = next(arg_limb_pieces) + + dest_name = self._limb_name(dest, dest_limb) + expr = "|".join(expr_parts) + mask = self._make_mask(dest, total_arg_bitwidth, dest_limb) + write(f"{dest_name} = ({expr}){mask};") def _next_limb_start(self, bit_position: int) -> int: """Given a bit position, return the bit position of the first bit in the next 64-bit limb. - This rounds bit_position up to the next even multiple of 64. For example, all - integers in the range [0, 63] return 64, and all integers in the range [64, 127] - return 128. + This rounds ``bit_position`` up to the next even multiple of 64. For example, + all integers in the range [0, 63] return 64, and all integers in the range [64, + 127] return 128. """ return bit_position + (64 - bit_position % 64) @@ -698,10 +924,10 @@ def _split_consecutive_slices( arg limbs into ConsecutiveSlices that don't fetch bits from multiple arg limbs. """ # Get op_params for the `n`th dest limb. - current_param = op_param[64 * n : min(dest.bitwidth, 64 * (n + 1))] + current_op_param = op_param[64 * n : min(dest.bitwidth, 64 * (n + 1))] split_slices = [] - for consecutive_slice in make_consecutive_slices(current_param): + for consecutive_slice in make_consecutive_slices(current_op_param): next_limb_arg_start = self._next_limb_start(consecutive_slice.arg_start) arg_end = consecutive_slice.arg_start + consecutive_slice.length @@ -712,34 +938,45 @@ def _split_consecutive_slices( # consecutive_slice fetches bits from two arg limbs. Split it into two # ConsecutiveSlices that each fetch bits from one arg limb. first_limb_length = next_limb_arg_start - consecutive_slice.arg_start - split_slices.append( - ConsecutiveSlice( - length=first_limb_length, - arg_start=consecutive_slice.arg_start, - dest_start=consecutive_slice.dest_start, - ) - ) - split_slices.append( - ConsecutiveSlice( - length=consecutive_slice.length - first_limb_length, - arg_start=next_limb_arg_start, - dest_start=consecutive_slice.dest_start + first_limb_length, - ) + split_slices.extend( + [ + ConsecutiveSlice( + length=first_limb_length, + arg_start=consecutive_slice.arg_start, + dest_start=consecutive_slice.dest_start, + ), + ConsecutiveSlice( + length=consecutive_slice.length - first_limb_length, + arg_start=next_limb_arg_start, + dest_start=consecutive_slice.dest_start + first_limb_length, + ), + ] ) return split_slices - def _build_slice(self, write, _op, param, args, dest): + def _build_slice( + self, + write: Callable, + op: str, + op_param, + args: list[WireVector], + dest: WireVector, + ): + if op != "s": + msg = f"Unexpected op {op}" + raise PyrtlInternalError(msg) + # Iterate over the slice's dest limbs. - for n in range(self._limbs(dest)): - split_slices = self._split_consecutive_slices(dest, param, n) + for dest_limb in range(self._limbs(dest)): + split_slices = self._split_consecutive_slices(dest, op_param, dest_limb) expr_parts = [] for split_slice in split_slices: # We only need to fetch bits from one arg limb because # _split_consecutive_slices split up any slices that fetch bits from # multiple arg limbs. - expr = f"{self.varname[args[0]]}[{split_slice.arg_start // 64}]" + expr = self._limb_name(args[0], split_slice.arg_start // 64) expr = shift(expr, ">>", split_slice.arg_start % 64) slice_arg_end = split_slice.arg_start + split_slice.length @@ -757,29 +994,27 @@ def _build_slice(self, write, _op, param, args, dest): expr = "|".join(expr_parts) - write(f"{self.varname[dest]}[{n}] = {expr};") + dest_name = self._limb_name(dest, dest_limb) + write(f"{dest_name} = {expr};") def _declare_mem_helpers(self, write): helpers = """ typedef uint64_t val_t; - typedef struct node - { + typedef struct node { uint64_t key; val_t *val; struct node *next; } node_t; - typedef struct hashmap - { + typedef struct hashmap { int size; int val_limbs; val_t *default_value; node_t **list; } hashmap_t; - hashmap_t *create_hash_map(int size, int val_limbs) - { + hashmap_t *create_hash_map(int size, int val_limbs) { int i; hashmap_t *h = (hashmap_t *) malloc(sizeof(hashmap_t)); h->size = size; @@ -793,21 +1028,17 @@ def _declare_mem_helpers(self, write): return h; } - int hash_code(hashmap_t *h, uint64_t key) - { + int hash_code(hashmap_t *h, uint64_t key) { return key % h->size; } - void insert(hashmap_t *h, uint64_t key, val_t val[]) - { + void insert(hashmap_t *h, uint64_t key, val_t val[]) { int pos = hash_code(h, key); struct node *list = h->list[pos]; struct node *new_node = (node_t *) malloc(sizeof(node_t)); struct node *temp = list; - while (temp) - { - if (temp->key == key) - { + while (temp) { + if (temp->key == key) { memcpy(temp->val, val, sizeof(val_t) * h->val_limbs); return; } @@ -821,15 +1052,12 @@ def _declare_mem_helpers(self, write): } EXPORT - val_t* lookup(hashmap_t *h, uint64_t key) - { + val_t* lookup(hashmap_t *h, uint64_t key) { int pos = hash_code(h, key); node_t *list = h->list[pos]; node_t *temp = list; - while (temp) - { - if (temp->key == key) - { + while (temp) { + if (temp->key == key) { return temp->val; } temp = temp->next; @@ -839,78 +1067,126 @@ def _declare_mem_helpers(self, write): """ write(helpers) + @dataclass + class InputMetadata: + """Tracks an ``Input``'s location in ``inputs_array``. + + The ``Input`` is stored at ``inputs_array[start:start + length]``. + """ + + start: int + """Index of ``Input``'s first word in ``inputs_array``. + + ``start`` points to the ``Input``'s least significant bits. + """ + + length: int + """Number of 64-bit limbs used to store this ``Input``.""" + + bitwidth: int + """Total bitwidth of the ``Input``.""" + + @dataclass + class OutputMetadata: + """Tracks an ``Output``'s location in ``outputs_array``. + + The ``Output`` is stored at ``outputs_array[start:start + length]``. + """ + + start: int + """Index of ``Output``'s first word in ``outputs_array``. + + ``start`` points to the ``Output``'s least significant bits. + """ + + length: int + """Number of 64-bit limbs used to store this ``Output``.""" + def _create_code(self, write): write("#include ") write("#include ") write("#include ") - # windows dllexport needed to make symbols visible + # `dllexport` is needed to make symbols visible on Windows. if platform.system() == "Windows": write("#define EXPORT __declspec(dllexport)") else: write("#define EXPORT") - # multiplication macro - # for efficient 64x64 -> 128 bit multiplication without uint128_t - # as -O0 optimization does not handle uint128_t well + # Multiplication macros for efficient 64x64 -> 128 bit multiplication without + # uint128_t. -O1 optimization does not handle uint128_t well. machine_alias = {"amd64": "x86_64", "aarch64": "arm64", "aarch64_be": "arm64"} machine = platform.machine().lower() machine = machine_alias.get(machine, machine) - mulinstr = { + mul_asm = { "x86_64": '"mulq %q3":"=a"(pl),"=d"(ph):"%0"(t0),"r"(t1):"cc"', - "arm64": '"mul %0, %2, %3\\n\\t" \\\n' - '"umulh %1, %2, %3":"=&r"(pl),"=r"(ph):"r"(t0),"r"(t1):"cc"', - "mips64": '"dmultu %2, %3\\n\\t" \\\n' - '"tmflo %0\\n\\t" \\\n' - '"mfhi %1":"=r"(pl),"=r"(ph):"r"(t0),"r"(t1)', + "arm64": ( + '"mul %0, %2, %3\\n\\t" \\\n' + '"umulh %1, %2, %3":"=&r"(pl),"=r"(ph):"r"(t0),"r"(t1):"cc"', + ), + "mips64": ( + '"dmultu %2, %3\\n\\t" \\\n' + '"tmflo %0\\n\\t" \\\n' + '"mfhi %1":"=r"(pl),"=r"(ph):"r"(t0),"r"(t1)', + ), } - if machine in mulinstr: - write(f"#define mul128(t0, t1, pl, ph) __asm__({mulinstr[machine]})") + if machine in mul_asm: + write(f"#define mul128(t0, t1, pl, ph) __asm__({mul_asm[machine]})") - # declare memories - mems = {net.op_param[1] for net in self.block.logic_subset("m@")} - for key in self._memmap: - if key not in mems: + # Declare memories. + self._declare_mem_helpers(write) + + # Find all of the Block's MemBlocks and RomBlocks. + memblocks = set() + romblocks = set() + for memblock in {net.op_param[1] for net in self.block.logic_subset("m@")}: + if isinstance(memblock, RomBlock): + romblocks.add(memblock) + else: + memblocks.add(memblock) + + for memblock in self._memory_value_map: + if memblock not in memblocks: msg = "unrecognized MemBlock in memory_value_map" raise PyrtlError(msg) - if isinstance(key, RomBlock): + if isinstance(memblock, RomBlock): msg = "RomBlock in memory_value_map" raise PyrtlError(msg) - self._declare_mem_helpers(write) - roms = {mem for mem in mems if isinstance(mem, RomBlock)} - self._declare_roms(write, roms) - mems = { - mem - for mem in mems - if isinstance(mem, MemBlock) and not isinstance(mem, RomBlock) - } - self._declare_mems(write, mems) - # single step function - write("static void sim_run_step(uint64_t inputs[], uint64_t outputs[]) {") + self._declare_roms(write, romblocks) + self._declare_mems(write, memblocks) + + # Define `sim_run_step`, which simulates one cycle. `step_inputs` and + # `step_outputs` are arrays of the step's input and output values, respectively. + write( + "static void sim_run_step(" + "uint64_t step_inputs[], uint64_t step_outputs[]) {" + ) write("uint64_t tmp, carry, tmphi, tmplo;") # temporary variables - # declare wire vectors + # Declare all WireVectors. for w in self.block.wirevector_set: - self._declare_wv(write, w) - - # inputs copied in - inputs = list(self.block.wirevector_subset(Input)) - # for each input wire, start and number of elements in input array - self._inputpos = {} - self._inputbw = {} # bitwidth of each input wire - ipos = 0 - for w in inputs: - self._inputpos[w.name] = ipos, self._limbs(w) - self._inputbw[w.name] = w.bitwidth - for n in range(self._limbs(w)): - write(f"{self.varname[w]}[{n}] = inputs[{ipos}];") - ipos += 1 - self._ibufsz = ipos # total length of input array - - # combinational logic + self._declare_wirevector(write, w) + + # Initialize Input WireVectors from `step_inputs`. + # + # `_inputs_metadata` maps each Input to its (start, length, bitwidth) in + # `step_inputs`. This will be used by `run` to pack its `inputs_array`. + self._inputs_metadata = {} + inputs_index = 0 + for input_wire in self.block.wirevector_subset(Input): + self._inputs_metadata[input_wire.name] = self.InputMetadata( + inputs_index, self._limbs(input_wire), input_wire.bitwidth + ) + for input_wire_limb in range(self._limbs(input_wire)): + input_wire_name = self._limb_name(input_wire, input_wire_limb) + write(f"{input_wire_name} = step_inputs[{inputs_index}];") + inputs_index += 1 + self._inputs_array_length = inputs_index # total length of `inputs` array + + # Build combinational logic. op_builders = { - "m": self._build_memread, + "m": self._build_mem_read, "w": self._build_wire, "~": self._build_not, "&": self._build_bitwise, @@ -918,8 +1194,8 @@ def _create_code(self, write): "^": self._build_bitwise, "n": self._build_nand, "=": self._build_eq, - "<": self._build_cmp, - ">": self._build_cmp, + "<": self._build_lt_gt, + ">": self._build_lt_gt, "x": self._build_mux, "+": self._build_add, "-": self._build_sub, @@ -930,67 +1206,91 @@ def _create_code(self, write): for net in self.block: # topological order if net.op in "r@": continue # skip synchronized nets - op, param, args, dest = net.op, net.op_param, net.args, net.dests[0] - write( - "// net {op} : {args} -> {dest}".format( - op=op, - args=", ".join(self.varname[x] for x in args), - dest=self.varname[dest], - ) - ) - op_builders[op](write, op, param, args, dest) - - # memory writes - for net in self.block.logic_subset("@"): - mem = net.op_param[1] - write(f"if ({self.varname[net.args[2]]}[0]) {{") + op = net.op + op_param = net.op_param + args = net.args + dest = net.dests[0] + + args_names = ", ".join(self.var_names[x] for x in args) + dest_name = self.var_names[dest] + write(f"// net {op} : {args_names} -> {dest_name}") + op_builders[op](write, op, op_param, args, dest) + + # Write memories. + for write_net in self.block.logic_subset("@"): + mem = write_net.op_param[1] + write_enabled = self._limb_name(write_net.args[2], 0) + write(f"if ({write_enabled}) {{") + # This code assumes write addresses are 64-bits or less. + if self._limbs(write_net.args[0]) > 1: + msg = "Write addresses longer than 64 bits are not supported." + raise PyrtlInternalError(msg) + write_addr = self._limb_name(write_net.args[0], 0) write( - f"insert({self.varname[mem]}, {self.varname[net.args[0]]}[0], " - f"{self.varname[net.args[1]]});" + f"insert({self.var_names[mem]}, {write_addr}, " + f"{self.var_names[write_net.args[1]]});" ) write("}") - # register updates - regnets = list(self.block.logic_subset("r")) - for x, net in enumerate(regnets): - rin = net.args[0] - write(f"uint64_t regtmp{x}[{self._limbs(rin)}];") - for n in range(self._limbs(rin)): - write(f"regtmp{x}[{n}] = {self.varname[rin]}[{n}];") - # double loop to ensure register-to-register chains update correctly - for x, net in enumerate(regnets): - rout = net.dests[0] - for n in range(self._limbs(rout)): - write(f"{self.varname[rout]}[{n}] = regtmp{x}[{n}];") - - # output copied out + # Update registers. This is done in two passes to ensure register chains update + # correctly. The first pass colllects every `Register.next` value in a temporary + # array. + reg_nets = list(self.block.logic_subset("r")) + for reg_index, reg_net in enumerate(reg_nets): + reg_next = reg_net.args[0] + write(f"uint64_t regtmp{reg_index}[{self._limbs(reg_next)}];") + for reg_next_limb in range(self._limbs(reg_next)): + reg_next_name = self._limb_name(reg_next, reg_next_limb) + write(f"regtmp{reg_index}[{reg_next_limb}] = {reg_next_name};") + # The second pass actually assigns the collected `Register.next` values to + # registers. + for reg_index, reg_net in enumerate(reg_nets): + reg = reg_net.dests[0] + for reg_limb in range(self._limbs(reg)): + reg_name = self._limb_name(reg, reg_limb) + write(f"{reg_name} = regtmp{reg_index}[{reg_limb}];") + + # Copy all Output values to the `step_outputs` array. outputs = list(self.block.wirevector_subset(Output)) - # for each output wire, start and number of elements in output array - self._outputpos = {} - opos = 0 - for w in outputs: - self._outputpos[w.name] = opos, self._limbs(w) - for n in range(self._limbs(w)): - write(f"outputs[{opos}] = {self.varname[w]}[{n}];") - opos += 1 - self._obufsz = opos # total length of output array + # `_outputs_metadata` maps from Output name to its (start, length) in + # `step_outputs`. This will be used by `run` to unpack its `outputs_array`. + self._outputs_metadata = {} + output_index = 0 + for output in outputs: + self._outputs_metadata[output.name] = self.OutputMetadata( + output_index, self._limbs(output) + ) + for output_limb in range(self._limbs(output)): + output_name = self._limb_name(output, output_limb) + write(f"step_outputs[{output_index}] = {output_name};") + output_index += 1 + self._outputs_array_length = output_index # total length of output array write("}") - # entry point + # `sim_run_all` is the compiled simulator's main entry point. `sim_run_all` runs + # multiple simulation steps by calling `sim_run_step` in a loop. Arguments: + # + # `nsteps` is the number of steps to simulate. + # `inputs` is a flattened array of input values for all steps, stored in + # step-major order. Input value `i` for step `s` is located at index: + # `s * self._inputs_array_length + i` + # `outputs` is a flattened array of output values for all steps, stored in + # step-major order. Output value `o` for step `s` is located at index: + # `s * self._outputs_array_length + o` write("EXPORT") write( - "void sim_run_all(" - "uint64_t stepcount, uint64_t inputs[], uint64_t outputs[]) {" + "void sim_run_all(uint64_t nsteps, uint64_t inputs[], uint64_t outputs[]) {" ) - write("uint64_t input_pos = 0, output_pos = 0;") - write("for (uint64_t stepnum = 0; stepnum < stepcount; stepnum++) {") - write("sim_run_step(inputs+input_pos, outputs+output_pos);") - write(f"input_pos += {self._ibufsz};") - write(f"output_pos += {self._obufsz};") + write("uint64_t input_index = 0, output_index = 0;") + write("for (uint64_t step = 0; step < nsteps; step++) {") + write("sim_run_step(inputs + input_index, outputs + output_index);") + write(f"input_index += {self._inputs_array_length};") + write(f"output_index += {self._outputs_array_length};") write("}}") def __del__(self): """Handle removal of the DLL when the simulator is deleted.""" + # TODO: Should this free the memory allocated by `create_hash_map` and `insert`? if self._dll is not None: handle = self._dll._handle if platform.system() == "Windows": diff --git a/pyrtl/simulation.py b/pyrtl/simulation.py index 4d4c37d8..b529ca82 100644 --- a/pyrtl/simulation.py +++ b/pyrtl/simulation.py @@ -1005,11 +1005,10 @@ def _compiled(self) -> str: expr = simple_func[net.op](*argvals) elif net.op == "c": expr_parts = [] - for i in range(len(net.args)): - shiftby = sum(len(j) for j in net.args[i + 1 :]) - expr_parts.append( - shift(self._arg_varname(net.args[i]), "<<", shiftby) - ) + shift_by = 0 + for arg in reversed(net.args): + expr_parts.append(shift(self._arg_varname(arg), "<<", shift_by)) + shift_by += arg.bitwidth expr = " | ".join(expr_parts) elif net.op == "s": arg = self._arg_varname(net.args[0]) diff --git a/tests/test_compilesim.py b/tests/test_compilesim.py index 8f967ff9..de661401 100644 --- a/tests/test_compilesim.py +++ b/tests/test_compilesim.py @@ -123,28 +123,6 @@ def test_bitslice2_and_concat_simulation(self): self.r.next <<= pyrtl.concat(left, right) self.check_trace("o 01377777\n") - def test_sparse_bitslice_simulation(self): - pyrtl.reset_working_block() - - # `bigger_input`'s storage spans three 64-bit limbs. - bigger_input = pyrtl.Input(name="bigger_input", bitwidth=160) - - even_bits = pyrtl.Output(name="even_bits", bitwidth=80) - even_bits <<= bigger_input[::2] - - odd_bits = pyrtl.Output(name="odd_bits", bitwidth=80) - odd_bits <<= bigger_input[1::2] - - sim = self.sim() - # 0xA == 0b1010 - sim.step({"bigger_input": 0xAAAA_AAAA_AAAA_AAAA_AAAA_AAAA_AAAA_AAAA_AAAA_AAAA}) - self.assertEqual(sim.inspect("even_bits"), 0) - self.assertEqual(sim.inspect("odd_bits"), 0xFFFF_FFFF_FFFF_FFFF_FFFF) - # 0x5 == 0b0101 - sim.step({"bigger_input": 0x5555_5555_5555_5555_5555_5555_5555_5555_5555_5555}) - self.assertEqual(sim.inspect("even_bits"), 0xFFFF_FFFF_FFFF_FFFF_FFFF) - self.assertEqual(sim.inspect("odd_bits"), 0) - def test_reg_to_reg_simulation(self): self.r2 = pyrtl.Register(bitwidth=self.bitwidth, name="r2") self.r.next <<= self.r2 @@ -172,17 +150,106 @@ class MultiLimbTestsBase(unittest.TestCase): def setUp(self): pyrtl.reset_working_block() - def test_multiple_limb_bitslice(self): - # `big_input`'s storage spans two 64-bit limbs. - big_input = pyrtl.Input(name="big_input", bitwidth=128) + def test_comparisons(self): + """Test cross-limb comparisons.""" + a = pyrtl.Input(name="a", bitwidth=80) + b = pyrtl.Input(name="b", bitwidth=80) - out = pyrtl.Output(name="out", bitwidth=32) - # This slice must fetch and combine bits from both of `big_input`'s limbs. - out <<= big_input[48:80] + lt = pyrtl.Output(name="lt", bitwidth=1) + gt = pyrtl.Output(name="gt", bitwidth=1) + eq = pyrtl.Output(name="eq", bitwidth=1) + lt <<= a < b + gt <<= a > b + eq <<= a == b sim = self.sim() - sim.step({"big_input": 0xFFFF_EEEE_DDDD_1234_5678_CCCC_BBBB_AAAA}) - self.assertEqual(sim.inspect("out"), 0x1234_5678) + # Comparisons where the high limbs determine the output. + sim.step({"a": 0x0000_AAAA_BBBB_CCCC_DDDD, "b": 0x0001_AAAA_BBBB_CCCC_DDDD}) + self.assertEqual(sim.inspect("lt"), 1) + self.assertEqual(sim.inspect("gt"), 0) + self.assertEqual(sim.inspect("eq"), 0) + + sim.step({"a": 0x0001_AAAA_BBBB_CCCC_DDDD, "b": 0x0000_AAAA_BBBB_CCCC_DDDD}) + self.assertEqual(sim.inspect("lt"), 0) + self.assertEqual(sim.inspect("gt"), 1) + self.assertEqual(sim.inspect("eq"), 0) + + # Comparisons where the high limbs match, so the low limbs must tiebreak. + sim.step({"a": 0xFFFF_0001_AAAA_BBBB_CCCC, "b": 0xFFFF_0000_AAAA_BBBB_CCCC}) + self.assertEqual(sim.inspect("lt"), 0) + self.assertEqual(sim.inspect("gt"), 1) + self.assertEqual(sim.inspect("eq"), 0) + + sim.step({"a": 0xFFFF_0000_AAAA_BBBB_CCCC, "b": 0xFFFF_0001_AAAA_BBBB_CCCC}) + self.assertEqual(sim.inspect("lt"), 1) + self.assertEqual(sim.inspect("gt"), 0) + self.assertEqual(sim.inspect("eq"), 0) + + sim.step({"a": 0xFFFF_0000_AAAA_BBBB_CCCC, "b": 0xFFFF_0000_AAAA_BBBB_CCCC}) + self.assertEqual(sim.inspect("lt"), 0) + self.assertEqual(sim.inspect("gt"), 0) + self.assertEqual(sim.inspect("eq"), 1) + + def test_add_carry(self): + """Test addition that carries into the next limb.""" + a = pyrtl.Input(name="a", bitwidth=128) + b = pyrtl.Input(name="b", bitwidth=128) + + sum = pyrtl.Output(name="sum", bitwidth=129) + sum <<= a + b + + sim = self.sim() + sim.step({"a": 0xFFFF_FFFF_FFFF_FFFF_FFFF_FFFF_FFFF_FFFF, "b": 1}) + self.assertEqual( + sim.inspect("sum"), 0x1_0000_0000_0000_0000_0000_0000_0000_0000 + ) + + sim.step({"a": 1, "b": 0xFFFF_FFFF_FFFF_FFFF_FFFF_FFFF_FFFF_FFFF}) + self.assertEqual( + sim.inspect("sum"), 0x1_0000_0000_0000_0000_0000_0000_0000_0000 + ) + + def test_sub_carry(self): + """Test subtraction that carries into the next limb.""" + a = pyrtl.Input(name="a", bitwidth=128) + b = pyrtl.Input(name="b", bitwidth=128) + + difference = pyrtl.Output(name="difference", bitwidth=129) + difference <<= a - b + + sim = self.sim() + sim.step({"a": 0, "b": 1}) + self.assertEqual( + sim.inspect("difference"), 0x1_FFFF_FFFF_FFFF_FFFF_FFFF_FFFF_FFFF_FFFF + ) + + def test_multiply(self): + """Test two-limb multiplication.""" + a = pyrtl.Input(name="a", bitwidth=80) + b = pyrtl.Input(name="b", bitwidth=80) + + product = pyrtl.Output(name="product", bitwidth=160) + product <<= a * b + + sim = self.sim() + sim.step({"a": 0xAAAA_BBBB_CCCC_DDDD_EEEE, "b": 2**4}) + self.assertEqual(sim.inspect("product"), 0xAAAA_BBBB_CCCC_DDDD_EEEE_0) + + sim.step({"a": 2**4, "b": 0xAAAA_BBBB_CCCC_DDDD_EEEE}) + self.assertEqual(sim.inspect("product"), 0xAAAA_BBBB_CCCC_DDDD_EEEE_0) + + def test_multiply_carry(self): + """Test multiplication with multiple carries.""" + a = pyrtl.Input(name="a", bitwidth=144) + b = pyrtl.Input(name="b", bitwidth=144) + + product = pyrtl.Output(name="product", bitwidth=288) + product <<= a * b + + sim = self.sim() + value = 0xFFFF_FFFF_FFFF_FFFF_FFFF_FFFF_FFFF_FFFF_FFFF + sim.step({"a": value, "b": value}) + self.assertEqual(sim.inspect("product"), value * value) def test_concat_fully_aligned(self): """Test concat where 64-bit args align with the dest limbs.""" @@ -246,6 +313,40 @@ def test_concat_split_offset_start(self): sim.step({"a": 0xAAAA_BBBB_CCCC, "b": 0xDDDD, "c": 0xEEEE}) self.assertEqual(sim.inspect("concat"), 0xAAAA_BBBB_CCCC_DDDD_EEEE) + def test_multiple_limb_bitslice(self): + """Test slices that span multiple limbs.""" + # `big_input`'s storage spans two 64-bit limbs. + big_input = pyrtl.Input(name="big_input", bitwidth=128) + + out = pyrtl.Output(name="out", bitwidth=32) + # This slice must fetch and combine bits from both of `big_input`'s limbs. + out <<= big_input[48:80] + + sim = self.sim() + sim.step({"big_input": 0xFFFF_EEEE_DDDD_1234_5678_CCCC_BBBB_AAAA}) + self.assertEqual(sim.inspect("out"), 0x1234_5678) + + def test_sparse_bitslice(self): + """Test bitslices that skip bits across multiple limbs.""" + # `bigger_input`'s storage spans three 64-bit limbs. + bigger_input = pyrtl.Input(name="bigger_input", bitwidth=160) + + even_bits = pyrtl.Output(name="even_bits", bitwidth=80) + even_bits <<= bigger_input[::2] + + odd_bits = pyrtl.Output(name="odd_bits", bitwidth=80) + odd_bits <<= bigger_input[1::2] + + sim = self.sim() + # 0xA == 0b1010 + sim.step({"bigger_input": 0xAAAA_AAAA_AAAA_AAAA_AAAA_AAAA_AAAA_AAAA_AAAA_AAAA}) + self.assertEqual(sim.inspect("even_bits"), 0) + self.assertEqual(sim.inspect("odd_bits"), 0xFFFF_FFFF_FFFF_FFFF_FFFF) + # 0x5 == 0b0101 + sim.step({"bigger_input": 0x5555_5555_5555_5555_5555_5555_5555_5555_5555_5555}) + self.assertEqual(sim.inspect("even_bits"), 0xFFFF_FFFF_FFFF_FFFF_FFFF) + self.assertEqual(sim.inspect("odd_bits"), 0) + class PrintTraceBase(unittest.TestCase): # note: doesn't include tests for compact=True because all the tests test that From 59e8ca496b5f19505f3678208eb1d99a80a3b377 Mon Sep 17 00:00:00 2001 From: Jeremy Lau <30300826+fdxmw@users.noreply.github.com> Date: Tue, 21 Jul 2026 15:31:27 -0700 Subject: [PATCH 2/3] Clean up .coveragerc: * Remove legacy `omit` filename patterns. * Ignore `PyrtlInternalError` and exception error messages. --- .coveragerc | 22 +++++++++++----------- 1 file changed, 11 insertions(+), 11 deletions(-) diff --git a/.coveragerc b/.coveragerc index 58bfaab0..88bc15c6 100644 --- a/.coveragerc +++ b/.coveragerc @@ -5,32 +5,32 @@ source = pyrtl [report] # Regexes for lines to exclude from consideration exclude_lines = - # Have to re-enable the standard pragma + # The standard pragma. pragma: no cover - # ignore pass statements that are not run: - # sometimes you need them to make a block (especially for unittests) - # need to be really careful with this one because 'pass' shows up in other places as well + # ignore pass statements that are not run: sometimes you need them to make + # a block (especially for unittests). Need to be really careful with these + # regexes because 'pass' shows up in other places as well \spass\s \spass$ - # Don't complain about missing debug-only code: + # Ignore debug-only code. def __repr__ if self\.debug - # Don't complain if tests don't hit defensive assertion code: + # Ignore assertion code. raise AssertionError raise NotImplementedError + raise PyrtlInternalError - # Don't complain if non-runnable code isn't run: + # Ignore exception error messages. + \smsg = + + # Ignore unreachable code. if 0: if __name__ == .__main__.: ignore_errors = True omit = - .tox/* */tests/* - */ctypes/* - six.py - pyparsing.py From aa02467600d95d9aa07dcd50fb2e0928b144f091 Mon Sep 17 00:00:00 2001 From: Jeremy Lau <30300826+fdxmw@users.noreply.github.com> Date: Tue, 21 Jul 2026 15:43:14 -0700 Subject: [PATCH 3/3] Upgrade GitHub actions. --- .github/workflows/python-release.yml | 10 +++++----- .github/workflows/python-test.yml | 6 +++--- 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/.github/workflows/python-release.yml b/.github/workflows/python-release.yml index 58a97bb2..b744f3bd 100644 --- a/.github/workflows/python-release.yml +++ b/.github/workflows/python-release.yml @@ -17,14 +17,14 @@ jobs: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v4 - - uses: astral-sh/setup-uv@v6 + - uses: actions/checkout@v7 + - uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true - name: Build distribution archives run: uv build - name: Upload distribution archives - uses: actions/upload-artifact@v4 + uses: actions/upload-artifact@v7 with: name: python-distribution-archives path: dist/ @@ -46,7 +46,7 @@ jobs: steps: - name: Download distribution archives - uses: actions/download-artifact@v4 + uses: actions/download-artifact@v7 with: name: python-distribution-archives path: dist/ @@ -72,7 +72,7 @@ jobs: steps: - name: Download distribution archives - uses: actions/download-artifact@v4 + uses: actions/download-artifact@v8 with: name: python-distribution-archives path: dist/ diff --git a/.github/workflows/python-test.yml b/.github/workflows/python-test.yml index e1a0cad7..c076a4a1 100644 --- a/.github/workflows/python-test.yml +++ b/.github/workflows/python-test.yml @@ -9,8 +9,8 @@ jobs: runs-on: ubuntu-latest steps: - - uses: actions/checkout@v4 - - uses: astral-sh/setup-uv@v6 + - uses: actions/checkout@v7 + - uses: astral-sh/setup-uv@c771a70e6277c0a99b617c7a806ffedaca235ff9 # v9.0.0 with: enable-cache: true - name: Install apt dependencies @@ -18,6 +18,6 @@ jobs: - name: Run tests run: uv run just tests - name: Upload coverage to Codecov - uses: codecov/codecov-action@v4 + uses: codecov/codecov-action@v7 env: CODECOV_TOKEN: ${{ secrets.CODECOV_TOKEN }}