diff --git a/test/external/external_test_am.py b/test/external/external_test_am.py index 011255bf4a..00f0a4ca20 100644 --- a/test/external/external_test_am.py +++ b/test/external/external_test_am.py @@ -176,6 +176,34 @@ class TestAMPageTable(unittest.TestCase): mm0.map_range(helper_va(0x1000000), 2 << 20, paddrs=[(0x10000, 2 << 20)], aspace=AddrSpace.PHYS) mm0.unmap_range(helper_va(0x1000000), 2 << 20) + def test_inspect_mode(self): + mm0 = self.d[0].mm + + # Map a few disjoint ranges inside a larger region. + mappings = [(0x10000, 0x3000), (0x20000, 0x2000), (0x1000000, 2 << 20)] + for va, sz in mappings: + mm0.map_range(helper_va(va), sz, paddrs=[(va, sz)], aspace=AddrSpace.PHYS) + + # Inspect over the whole region: should visit all mapped pages. + ctx = PageTableTraverseContext(self.d[0], mm0.root_page_table, helper_va(0x0), inspect=True) + visited = set() + for _off, pt, pte_idx, n_ptes, pte_covers in ctx.next(0x4000000): + for i in range(n_ptes): + pte = helper_read_entry_components(pt.entries[pte_idx + i]) + if pte['valid']: + for p in range(0, pte_covers, 0x1000): visited.add(pte['paddr'] + p) + + expected_pages = {va + off for va, sz in mappings for off in range(0, sz, 0x1000)} + assert visited == expected_pages + + for va, sz in mappings: + mm0.unmap_range(helper_va(va), sz) + + # Inspect after unmap: should find no valid entries. + ctx = PageTableTraverseContext(self.d[0], mm0.root_page_table, helper_va(0x0), inspect=True) + for _off, pt, pte_idx, n_ptes, pte_covers in ctx.next(0x4000000): + for i in range(n_ptes): assert not pt.valid(pte_idx + i) + def test_frag_size(self): mm0 = self.d[0].mm diff --git a/tinygrad/runtime/support/memory.py b/tinygrad/runtime/support/memory.py index a63f90c8d2..abaac81e94 100644 --- a/tinygrad/runtime/support/memory.py +++ b/tinygrad/runtime/support/memory.py @@ -113,8 +113,8 @@ class AddrSpace(enum.Enum): PHYS = enum.auto(); SYS = enum.auto(); PEER = enum.a class VirtMapping: va_addr:int; size:int; paddrs:list[tuple[int, int]]; aspace:AddrSpace; uncached:bool=False; snooped:bool=False # noqa: E702 class PageTableTraverseContext: - def __init__(self, dev, pt, vaddr, create_pts=False, free_pts=False, boot=False): - self.dev, self.vaddr, self.create_pts, self.free_pts, self.boot = dev, vaddr - dev.mm.va_base, create_pts, free_pts, boot + def __init__(self, dev, pt, vaddr, create_pts=False, free_pts=False, inspect=False, boot=False): + self.dev, self.vaddr, self.create_pts, self.free_pts, self.inspect, self.boot = dev, vaddr - dev.mm.va_base, create_pts, free_pts, inspect, boot self.pt_stack:list[tuple[Any, int, int]] = [(pt, self._pt_pte_idx(pt, self.vaddr), self._pt_pte_size(pt))] def _pt_pte_cnt(self, lv): return self.dev.mm.pte_cnt[lv] @@ -151,13 +151,17 @@ class PageTableTraverseContext: def next(self, size:int, paddr:int|None=None, off:int=0): while size > 0: pt, pte_idx, pte_covers = self.pt_stack[-1] + + # create_pts goes down until the page covers the request. + # free_pts goes down to the table, it assumses all entries are valid on the range (and validates that) + # inspect just visits any valid ranges and yields them. if self.create_pts: assert paddr is not None, "paddr must be provided when allocating new page tables" while pte_covers > size or not pt.supports_huge_page(paddr+off) or self.vaddr&(pte_covers-1) != 0: pt, pte_idx, pte_covers = self.level_down() else: - while not pt.is_page(pte_idx): pt, pte_idx, pte_covers = self.level_down() + while not pt.is_page(pte_idx) and (self.free_pts or pt.valid(pte_idx)): pt, pte_idx, pte_covers = self.level_down() - entries = min(size // pte_covers, self._pt_pte_cnt(pt.lv) - pte_idx) + entries = max(min(size // pte_covers, self._pt_pte_cnt(pt.lv) - pte_idx), 1 if self.inspect else 0) assert entries > 0, f"Invalid entries {size=:#x}, {pte_covers=:#x}" yield off, pt, pte_idx, entries, pte_covers @@ -197,11 +201,14 @@ class MemoryManager: assert size == sum(p[1] for p in paddrs), f"Size mismatch {size=} {sum(p[1] for p in paddrs)=}" + ctx = PageTableTraverseContext(self.dev, self.root_page_table, vaddr, boot=boot, inspect=True) + for _, pt, pte_idx, pte_cnt, _ in ctx.next(size): + for pte_off in range(pte_cnt): assert not pt.valid(pte_idx + pte_off), f"PTE already mapped: {pt.entry(pte_idx + pte_off):#x}" + ctx = PageTableTraverseContext(self.dev, self.root_page_table, vaddr, create_pts=True, boot=boot) for paddr, psize in paddrs: for off, pt, pte_idx, pte_cnt, pte_covers in ctx.next(psize, paddr=paddr): for pte_off in range(pte_cnt): - assert not pt.valid(pte_idx + pte_off), f"PTE already mapped: {pt.entry(pte_idx + pte_off):#x}" pt.set_entry(pte_idx + pte_off, paddr + off + pte_off * pte_covers, uncached=uncached, aspace=aspace, snooped=snooped, frag=self._frag_size(ctx.vaddr+off, pte_cnt * pte_covers), valid=True)