from __future__ import annotations from dataclasses import dataclass, field from typing import Callable, Iterable, Literal, Protocol AccessKind = Literal["read", "write", "prefetch"] MissKind = Literal["hit", "cold", "capacity", "conflict"] class PipelineMemoryAccess(Protocol): address: int write: bool size: int @dataclass(frozen=True, slots=True) class CacheConfig: name: str sets: int ways: int block_size: int hit_latency: int = 1 fill_latency: int = 20 write_back: bool = True write_allocate: bool = True def __post_init__(self) -> None: if self.sets < 1 or self.ways < 1 or self.block_size < 1: raise ValueError("sets, ways and block_size must be positive") if self.hit_latency < 0 or self.fill_latency < 0: raise ValueError("hit_latency and fill_latency must be non-negative") @property def lines(self) -> int: return self.sets * self.ways @dataclass(slots=True) class CacheLine: valid: bool = False tag: int = 0 block: int = -1 dirty: bool = False touched_at: int = 0 prefetched: bool = False @dataclass(frozen=True, slots=True) class CacheAccess: step: int kind: AccessKind address: int block: int index: int tag: int hit: bool miss_kind: MissKind evicted_block: int | None dirty_eviction: bool fill_bytes: int writeback_bytes: int lower_write_bytes: int latency: int @dataclass(slots=True) class CacheStats: demand: int = 0 hits: int = 0 misses: int = 0 cold: int = 0 capacity: int = 0 conflict: int = 0 fills: int = 0 writebacks: int = 0 fill_bytes: int = 0 writeback_bytes: int = 0 lower_write_bytes: int = 0 prefetches: int = 0 useful_prefetches: int = 0 useless_prefetches: int = 0 @dataclass(slots=True) class Cache: config: CacheConfig sets: list[list[CacheLine]] = field(init=False) clock: int = 0 stats: CacheStats = field(default_factory=CacheStats) def __post_init__(self) -> None: self.sets = [ [CacheLine() for _ in range(self.config.ways)] for _ in range(self.config.sets) ] def split(self, address: int) -> tuple[int, int, int]: block = address // self.config.block_size index = block % self.config.sets tag = block // self.config.sets return block, index, tag def probe(self, address: int) -> bool: block, index, tag = self.split(address) return any(line.valid and line.tag == tag and line.block == block for line in self.sets[index]) def access(self, address: int, kind: AccessKind = "read", miss_kind: MissKind = "cold") -> CacheAccess: self.clock += 1 block, index, tag = self.split(address) line_set = self.sets[index] hit_line = next((line for line in line_set if line.valid and line.tag == tag), None) demand = kind in ("read", "write") if demand: self.stats.demand += 1 else: self.stats.prefetches += 1 if hit_line is not None: hit_line.touched_at = self.clock if kind == "write": if self.config.write_back: hit_line.dirty = True else: self.stats.lower_write_bytes += 4 if demand and hit_line.prefetched: self.stats.useful_prefetches += 1 hit_line.prefetched = False self.stats.hits += int(demand) return CacheAccess(self.clock, kind, address, block, index, tag, True, "hit", None, False, 0, 0, 4 if kind == "write" and not self.config.write_back else 0, self.config.hit_latency) allocate = kind != "write" or self.config.write_allocate lower_write_bytes = 4 if kind == "write" and (not self.config.write_back or not allocate) else 0 if not allocate: if demand: self.stats.misses += 1 setattr(self.stats, miss_kind, getattr(self.stats, miss_kind) + 1) self.stats.lower_write_bytes += lower_write_bytes return CacheAccess(self.clock, kind, address, block, index, tag, False, miss_kind, None, False, 0, 0, lower_write_bytes, self.config.hit_latency + self.config.fill_latency) victim = next((line for line in line_set if not line.valid), None) if victim is None: victim = min(line_set, key=lambda line: line.touched_at) evicted_block = victim.block if victim.valid else None dirty_eviction = victim.valid and victim.dirty writeback_bytes = self.config.block_size if dirty_eviction else 0 if victim.valid and victim.prefetched: self.stats.useless_prefetches += 1 victim.valid = True victim.tag = tag victim.block = block victim.dirty = kind == "write" and self.config.write_back victim.touched_at = self.clock victim.prefetched = kind == "prefetch" if demand: self.stats.misses += 1 setattr(self.stats, miss_kind, getattr(self.stats, miss_kind) + 1) self.stats.fills += 1 self.stats.fill_bytes += self.config.block_size self.stats.writebacks += int(dirty_eviction) self.stats.writeback_bytes += writeback_bytes self.stats.lower_write_bytes += lower_write_bytes return CacheAccess(self.clock, kind, address, block, index, tag, False, miss_kind, evicted_block, dirty_eviction, self.config.block_size, writeback_bytes, lower_write_bytes, self.config.hit_latency + self.config.fill_latency + (self.config.fill_latency if dirty_eviction else 0)) @dataclass(slots=True) class ClassifiedCache: cache: Cache shadow: Cache seen_blocks: set[int] = field(default_factory=set) @classmethod def create(cls, config: CacheConfig) -> "ClassifiedCache": shadow = Cache(CacheConfig("fully-assoc-shadow", 1, config.lines, config.block_size)) return cls(Cache(config), shadow) def miss_kind(self, address: int) -> MissKind: block = address // self.cache.config.block_size if self.cache.probe(address): return "hit" if block not in self.seen_blocks: return "cold" return "conflict" if self.shadow.probe(address) else "capacity" def access(self, address: int, kind: AccessKind = "read") -> CacheAccess: miss_kind = self.miss_kind(address) result = self.cache.access(address, kind, miss_kind) self.shadow.access(address, "read", "cold") self.seen_blocks.add(result.block) return result def hierarchy_amat(l1_hit: int, l1_miss_rate: float, l2_hit: int, l2_local_miss_rate: float, memory_penalty: int) -> float: if not 0 <= l1_miss_rate <= 1 or not 0 <= l2_local_miss_rate <= 1: raise ValueError("miss rates must be in [0, 1]") return l1_hit + l1_miss_rate * (l2_hit + l2_local_miss_rate * memory_penalty) def latency_function(classified: ClassifiedCache) -> Callable[[PipelineMemoryAccess], int]: def latency(access: PipelineMemoryAccess) -> int: block_size = classified.cache.config.block_size if access.size != 4: raise ValueError("pipeline latency interface supports 4-byte accesses only") if access.address < 0 or access.address % access.size != 0: raise ValueError("pipeline latency interface requires non-negative aligned addresses") if access.address // block_size != (access.address + access.size - 1) // block_size: raise ValueError("pipeline latency interface does not model cross-block accesses") return classified.access(access.address, "write" if access.write else "read").latency return latency @dataclass(frozen=True, slots=True) class DramConfig: banks: int = 2 row_size: int = 64 block_size: int = 16 t_rcd: int = 3 t_cl: int = 2 t_rp: int = 3 t_rfc: int = 8 t_refi: int = 30 def __post_init__(self) -> None: if self.banks < 1 or self.row_size < 1 or self.block_size < 1: raise ValueError("banks, row_size and block_size must be positive") if self.block_size > self.row_size: raise ValueError("block_size must not exceed row_size") if any(value < 0 for value in (self.t_rcd, self.t_cl, self.t_rp, self.t_rfc)): raise ValueError("DRAM timing values must be non-negative") if self.t_refi <= 0: raise ValueError("t_refi must be positive") if self.t_rfc >= self.t_refi: raise ValueError("t_rfc must be smaller than t_refi in this finite teaching model") @dataclass(frozen=True, slots=True) class DramRequest: name: str address: int arrival: int @dataclass(frozen=True, slots=True) class DramEvent: name: str bank: int row: int row_event: Literal["empty", "hit", "conflict"] arrival: int start: int service: int done: int wait: int refresh_wait: int @dataclass(slots=True) class DramController: config: DramConfig = field(default_factory=DramConfig) open_rows: list[int | None] = field(init=False) bank_free_at: list[int] = field(init=False) bus_free_at: int = 0 next_refresh: int = field(init=False) def __post_init__(self) -> None: self.open_rows = [None] * self.config.banks self.bank_free_at = [0] * self.config.banks self.next_refresh = self.config.t_refi def map_address(self, address: int) -> tuple[int, int]: block = address // self.config.block_size bank = block % self.config.banks row = address // self.config.row_size return bank, row def issue(self, request: DramRequest) -> DramEvent: bank, row = self.map_address(request.address) start = max(request.arrival, self.bank_free_at[bank], self.bus_free_at) refresh_wait = 0 while start >= self.next_refresh: refresh_start = max(self.next_refresh, self.bus_free_at) refresh_done = refresh_start + self.config.t_rfc refresh_wait += max(0, refresh_done - start) self.bus_free_at = refresh_done self.bank_free_at = [max(t, refresh_done) for t in self.bank_free_at] self.open_rows = [None] * self.config.banks self.next_refresh += self.config.t_refi start = max(request.arrival, self.bank_free_at[bank], self.bus_free_at) current = self.open_rows[bank] if current is None: row_event = "empty" service = self.config.t_rcd + self.config.t_cl elif current == row: row_event = "hit" service = self.config.t_cl else: row_event = "conflict" service = self.config.t_rp + self.config.t_rcd + self.config.t_cl done = start + service self.open_rows[bank] = row self.bank_free_at[bank] = done self.bus_free_at = done return DramEvent(request.name, bank, row, row_event, request.arrival, start, service, done, start - request.arrival, refresh_wait) @dataclass(frozen=True, slots=True) class MemoryRequest: name: str address: int ready_at: int = 0 @dataclass(frozen=True, slots=True) class MshrEvent: name: str block: int issue: int done: int merged: bool def simulate_mshr(requests: Iterable[MemoryRequest], block_size: int = 16, miss_latency: int = 10, mshr_entries: int = 2, dependent: bool = False) -> list[MshrEvent]: if block_size < 1: raise ValueError("block_size must be positive") if miss_latency < 0: raise ValueError("miss_latency must be non-negative") if mshr_entries < 1: raise ValueError("mshr_entries must be positive") events: list[MshrEvent] = [] outstanding: dict[int, int] = {} chain_ready = 0 scheduler_time = 0 for request in requests: block = request.address // block_size issue = max(request.ready_at, chain_ready, scheduler_time) outstanding = {b: done for b, done in outstanding.items() if done > issue} if block in outstanding: done = outstanding[block] events.append(MshrEvent(request.name, block, issue, done, True)) else: while len(outstanding) >= mshr_entries: earliest = min(outstanding.values()) issue = max(issue, earliest) outstanding = {b: done for b, done in outstanding.items() if done > issue} done = issue + miss_latency outstanding[block] = done events.append(MshrEvent(request.name, block, issue, done, False)) scheduler_time = issue if dependent: chain_ready = done return events def run_next_line_prefetch(addresses: Iterable[int], config: CacheConfig) -> tuple[list[CacheAccess], CacheStats]: cache = ClassifiedCache.create(config) events: list[CacheAccess] = [] for address in addresses: demand = cache.access(address, "read") events.append(demand) next_line = (demand.block + 1) * config.block_size events.append(cache.cache.access(next_line, "prefetch", "cold")) return events, cache.cache.stats