Files
PSPRecomp/profiles/vcs/tools/apply_v812_code_density.py
T
Jessica_Natalia 430e45bdb0 more updates
more updates
2026-08-27 04:39:50 -03:00

496 lines
21 KiB
Python

#!/usr/bin/env python3
from __future__ import annotations
import argparse
import json
import re
from dataclasses import asdict, dataclass
from pathlib import Path
STAGE = "V8.12 generated code density"
SHARED_BLOCK_MARKER = "PSPRECOMP_V812_SHARED_SCHED_BLOCK"
SHARED_JUMP_MARKER = "PSPRECOMP_V812_SHARED_SCHED_JUMP"
REG_MASK_MARKER = "PSPRECOMP_V812_COMPACT_REGISTRATION"
RUNTIME_HELPER = "register_generated_entry_mask"
EXPECTED_UNITS = 234
EXPECTED_REGISTERED_ENTRIES = 181789
EXPECTED_LOCAL_LINKS = 11299
MASK_RE = re.compile(
r"static constexpr std::uint64_t (?P<name>kEntryMasks_recomp_unit_(?P<unit>\d{4}))"
r"\[(?P<count>\d+)\] = \{\n(?P<body>.*?)\n\};",
re.DOTALL,
)
REGISTER_FN_RE = re.compile(
r'^[ \t]*runtime\.register_function\(0x[0-9A-Fa-f]+u,\s*&recomp_unit_\d{4},\s*"recomp_unit_\d{4}"\);\s*$',
re.MULTILINE,
)
REGISTER_BODY_RE = re.compile(
r"void register_generated_unit_(?P<bucket>\d+)\(Runtime &runtime\) \{\n"
r"(?P<body>.*?)\n\}",
re.DOTALL,
)
REGISTER_UNIT_CALL_RE = re.compile(
r'^[ \t]*runtime\.register_generated_unit\('
r'(?P<bucket>\d+)u,\s*(?P<base>0x[0-9A-Fa-f]+)u,\s*(?P<span>\d+)u,\s*'
r'&(?P<fn>recomp_unit_(?P<unit>\d{4})),\s*&(?P<entry>recomp_unit_\d{4}_entry)\);\s*$',
re.MULTILINE,
)
DISPATCH_ORIGIN_RE = re.compile(
r"const std::uint32_t entry_delta = local_pc - (?P<origin>0x[0-9A-Fa-f]+)u;"
)
COMPACT_REG_CALL_RE = re.compile(
r'runtime\.register_generated_entry_mask\((?P<origin>0x[0-9A-Fa-f]+)u,\s*'
r'&(?P<fn>recomp_unit_\d{4}),\s*"(?P=fn)",\s*'
r'(?P<mask>kEntryMasks_recomp_unit_\d{4}),\s*(?P<groups>\d+)u\);',
re.DOTALL,
)
# Exact V8.11 multi-line scheduler blocks, generated by apply_v811_perf.py.
REGCACHE_LOCAL_RE = re.compile(
r"(?P<i>[ \t]*)ctx\.pc = jump_target;\n"
r"(?P=i)AOT_REGCACHE_SYNC_OUT\(\);\n"
r"(?P=i)// PSPRECOMP_V811_LOCAL_LINK: preserve the 256-transfer scheduler boundary, skip outer trampoline\.\n"
r"(?P=i)if \(local_redispatch_rounds < 7u && rt\.continue_generated_local_dispatch\(ctx\)\) \{\n"
r"(?P=i) \+\+local_redispatch_rounds;\n"
r"(?P=i) AOT_REGCACHE_SYNC_IN\(\);\n"
r"(?P=i) local_transfers = 0u;\n"
r"(?P=i) local_pc = ctx\.pc;\n"
r"(?P=i) entry_id = 0u;\n"
r"(?P=i) goto LOCAL_DISPATCH;\n"
r"(?P=i)\}\n"
r"(?P=i)return;"
)
PLAIN_LOCAL_RE = re.compile(
r"(?P<i>[ \t]*)ctx\.pc = jump_target;\n"
r"(?P=i)// PSPRECOMP_V811_LOCAL_LINK: preserve the 256-transfer scheduler boundary, skip outer trampoline\.\n"
r"(?P=i)if \(local_redispatch_rounds < 7u && rt\.continue_generated_local_dispatch\(ctx\)\) \{\n"
r"(?P=i) \+\+local_redispatch_rounds;\n"
r"(?P=i) local_transfers = 0u;\n"
r"(?P=i) local_pc = ctx\.pc;\n"
r"(?P=i) entry_id = 0u;\n"
r"(?P=i) goto LOCAL_DISPATCH;\n"
r"(?P=i)\}\n"
r"(?P=i)return;"
)
@dataclass
class UnitStats:
path: str
bytes_before: int
bytes_after: int
entry_count: int
registration_lines_removed: int
local_blocks_shared: int
had_regcache: bool
def read(path: Path) -> str:
return path.read_text(encoding="utf-8", errors="strict")
def write_if_changed(path: Path, old: str, new: str) -> bool:
if old == new:
return False
tmp = path.with_suffix(path.suffix + ".v812tmp")
tmp.write_text(new, encoding="utf-8", newline="\n")
tmp.replace(path)
return True
def require_once(text: str, needle: str, path: Path, description: str) -> None:
count = text.count(needle)
if count != 1:
raise RuntimeError(f"{path}: expected one {description}, found {count}")
def parse_masks(text: str, path: Path) -> tuple[str, int, list[int], int]:
m = MASK_RE.search(text)
if not m:
raise RuntimeError(f"{path}: V8.11 compact entry masks not found")
declared = int(m.group("count"))
values = [int(x, 16) for x in re.findall(r"0x([0-9A-Fa-f]{16})ull", m.group("body"))]
if len(values) != declared:
raise RuntimeError(f"{path}: compact mask declared {declared}, parsed {len(values)}")
entries = sum(x.bit_count() for x in values)
return m.group("name"), declared, values, entries
def patch_runtime_hpp(path: Path) -> bool:
old = read(path)
if RUNTIME_HELPER in old:
return False
anchor = (
" void register_generated_unit(std::uint32_t unit_index, std::uint32_t unit_address,\n"
" std::uint32_t unit_span, RecompiledFunction function,\n"
" RecompiledEntryFunction entry_function = nullptr);\n"
)
require_once(old, anchor, path, "register_generated_unit declaration")
addition = anchor + (
" // V8.12: register all valid per-PC entries from the compact 64-slot\n"
" // occupancy masks already emitted for generated-unit dispatch. This\n"
" // preserves the exact legacy function/direct-chain tables while avoiding\n"
" // hundreds of source-level registration call sites per translation unit.\n"
" void register_generated_entry_mask(std::uint32_t unit_address,\n"
" RecompiledFunction function,\n"
" std::string_view name,\n"
" const std::uint64_t *entry_masks,\n"
" std::size_t group_count);\n"
)
return write_if_changed(path, old, old.replace(anchor, addition, 1))
def patch_runtime_cpp(path: Path) -> bool:
old = read(path)
text = old
if "#include <bit>\n" not in text:
anchor = "#include <array>\n"
require_once(text, anchor, path, "<array> include")
text = text.replace(anchor, anchor + "#include <bit>\n", 1)
if f"Runtime::{RUNTIME_HELPER}" not in text:
# Insert immediately after register_generated_unit(), before invoke_isolated_aot().
anchor = "\nbool Runtime::invoke_isolated_aot(std::uint32_t address, AllegrexContext &ctx) {\n"
require_once(text, anchor, path, "invoke_isolated_aot implementation")
impl = (
"\n// PSPRECOMP_V812_COMPACT_REGISTRATION\n"
"void Runtime::register_generated_entry_mask(std::uint32_t unit_address,\n"
" RecompiledFunction function,\n"
" std::string_view name,\n"
" const std::uint64_t *entry_masks,\n"
" std::size_t group_count) {\n"
" if (function == nullptr || entry_masks == nullptr || group_count == 0u) return;\n"
" for (std::size_t group = 0; group < group_count; ++group) {\n"
" std::uint64_t bits = entry_masks[group];\n"
" while (bits != 0u) {\n"
" const auto bit = static_cast<std::uint32_t>(std::countr_zero(bits));\n"
" const auto slot = static_cast<std::uint32_t>(group * 64u) + bit;\n"
" register_function(unit_address + slot * 4u, function, std::string(name));\n"
" bits &= bits - 1u;\n"
" }\n"
" }\n"
"}\n"
)
text = text.replace(anchor, impl + anchor, 1)
return write_if_changed(path, old, text)
def insert_shared_scheduler_block(text: str, path: Path, has_regcache: bool) -> str:
if SHARED_BLOCK_MARKER in text:
return text
# First guest label follows the LOCAL_DISPATCH switch and is at outer function scope.
dispatch = text.find("LOCAL_DISPATCH:\n")
if dispatch < 0:
raise RuntimeError(f"{path}: LOCAL_DISPATCH not found")
m = re.search(r"(?m)^L_[0-9A-Fa-f]+:\n", text[dispatch:])
if not m:
raise RuntimeError(f"{path}: first guest label after LOCAL_DISPATCH not found")
pos = dispatch + m.start()
lines = [
f"// {SHARED_BLOCK_MARKER}: one scheduler-exact slow boundary per unit.",
"LOCAL_SCHED_BOUNDARY:",
" ctx.pc = jump_target;",
]
if has_regcache:
lines.append(" AOT_REGCACHE_SYNC_OUT();")
lines.extend([
" if (local_redispatch_rounds < 7u && rt.continue_generated_local_dispatch(ctx)) {",
" ++local_redispatch_rounds;",
])
if has_regcache:
lines.append(" AOT_REGCACHE_SYNC_IN();")
lines.extend([
" local_transfers = 0u;",
" local_pc = ctx.pc;",
" entry_id = 0u;",
" goto LOCAL_DISPATCH;",
" }",
" return;",
"",
])
return text[:pos] + "\n".join(lines) + text[pos:]
def share_local_scheduler_blocks(text: str, path: Path) -> tuple[str, int]:
if SHARED_BLOCK_MARKER in text:
return text, text.count(SHARED_JUMP_MARKER)
has_regcache = "#define AOT_REGCACHE_SYNC_OUT()" in text
count = 0
def repl(m: re.Match[str]) -> str:
nonlocal count
count += 1
i = m.group("i")
return f"{i}// {SHARED_JUMP_MARKER}\n{i}goto LOCAL_SCHED_BOUNDARY;"
if has_regcache:
text = REGCACHE_LOCAL_RE.sub(repl, text)
# A register-cached unit must not contain an unwrapped V8.11 block.
if "PSPRECOMP_V811_LOCAL_LINK" in text:
raise RuntimeError(f"{path}: unrecognized V8.11 local-link form remains in regcache unit")
else:
text = PLAIN_LOCAL_RE.sub(repl, text)
if "PSPRECOMP_V811_LOCAL_LINK" in text:
raise RuntimeError(f"{path}: unrecognized V8.11 local-link form remains in plain unit")
if count:
text = insert_shared_scheduler_block(text, path, has_regcache)
return text, count
def compact_registration(text: str, path: Path, mask_name: str, group_count: int, entry_count: int) -> tuple[str, int]:
# Already compacted.
if REG_MASK_MARKER in text:
return text, entry_count
matches = list(REGISTER_BODY_RE.finditer(text))
if len(matches) != 1:
raise RuntimeError(f"{path}: expected one register_generated_unit_N body, found {len(matches)}")
m = matches[0]
body = m.group("body")
unit_call = REGISTER_UNIT_CALL_RE.search(body)
if unit_call is None:
raise RuntimeError(f"{path}: register_generated_unit call not found")
unit_suffix = unit_call.group("unit")
fn = unit_call.group("fn")
bucket = unit_call.group("bucket")
unit_base = unit_call.group("base")
origin_match = DISPATCH_ORIGIN_RE.search(text)
if origin_match is None:
raise RuntimeError(f"{path}: compact-dispatch entry origin not found")
entry_origin = origin_match.group("origin")
if mask_name != f"kEntryMasks_{fn}":
raise RuntimeError(f"{path}: mask/function mismatch {mask_name} vs {fn}")
registration_lines = REGISTER_FN_RE.findall(body)
reg_count = len(registration_lines)
if reg_count != entry_count:
raise RuntimeError(
f"{path}: registration lines={reg_count} but compact masks contain {entry_count} entries"
)
# Require the old body to consist only of the unit registration plus exact per-PC registrations.
scrubbed = REGISTER_FN_RE.sub("", body)
scrubbed = REGISTER_UNIT_CALL_RE.sub("", scrubbed)
if scrubbed.strip():
raise RuntimeError(f"{path}: unexpected code in generated registration body: {scrubbed.strip()[:120]!r}")
unit_call_line = unit_call.group(0).strip()
new_body = (
f"void register_generated_unit_{m.group('bucket')}(Runtime &runtime) {{\n"
f" {unit_call_line}\n"
f" // {REG_MASK_MARKER}: same entry set, looped once in Runtime instead of source-expanded calls.\n"
# IMPORTANT: kEntryMasks is relative to the compact dispatch origin,
# which is not always the 16 KiB generated-unit base (some units begin
# at base+4). Using unit_base shifts every registered PC and breaks boot.
f" runtime.register_generated_entry_mask({entry_origin}u, &{fn}, \"{fn}\",\n"
f" {mask_name}, {group_count}u);\n"
f"}}"
)
text = text[:m.start()] + new_body + text[m.end():]
return text, reg_count
def transform_unit(path: Path) -> UnitStats:
old = read(path)
text = old
mask_name, group_count, _masks, entry_count = parse_masks(text, path)
had_regcache = "#define AOT_REGCACHE_SYNC_OUT()" in text
text, local_count = share_local_scheduler_blocks(text, path)
text, registration_count = compact_registration(text, path, mask_name, group_count, entry_count)
write_if_changed(path, old, text)
return UnitStats(
path=path.name,
bytes_before=len(old.encode("utf-8")),
bytes_after=len(text.encode("utf-8")),
entry_count=entry_count,
registration_lines_removed=registration_count if REG_MASK_MARKER not in old else 0,
local_blocks_shared=local_count if SHARED_BLOCK_MARKER not in old else 0,
had_regcache=had_regcache,
)
def validate(root: Path) -> dict:
profile = root / "profiles" / "vcs"
generated = profile / "generated"
units = sorted(generated.glob("generated_unit_*.cpp"))
if len(units) != EXPECTED_UNITS:
raise RuntimeError(f"expected {EXPECTED_UNITS} generated units, found {len(units)}")
runtime_hpp = read(root / "include" / "psprecomp" / "runtime.hpp")
runtime_cpp = read(root / "src" / "runtime.cpp")
if RUNTIME_HELPER not in runtime_hpp or f"Runtime::{RUNTIME_HELPER}" not in runtime_cpp:
raise RuntimeError("V8.12 compact registration Runtime helper is missing")
total_entries = 0
total_jumps = 0
shared_units = 0
compact_registration_units = 0
total_bytes = 0
for path in units:
text = read(path)
_mask_name, _groups, _masks, entries = parse_masks(text, path)
total_entries += entries
total_bytes += len(text.encode("utf-8"))
jumps = text.count(SHARED_JUMP_MARKER)
total_jumps += jumps
if jumps:
if text.count(SHARED_BLOCK_MARKER) != 1:
raise RuntimeError(f"{path}: shared jumps present without exactly one shared block")
if text.count("LOCAL_SCHED_BOUNDARY:") != 1:
raise RuntimeError(f"{path}: expected exactly one LOCAL_SCHED_BOUNDARY label")
shared_units += 1
elif SHARED_BLOCK_MARKER in text:
raise RuntimeError(f"{path}: shared block exists with no shared jumps")
if "PSPRECOMP_V811_LOCAL_LINK" in text:
raise RuntimeError(f"{path}: old duplicated V8.11 local-link block remains")
if text.count(REG_MASK_MARKER) != 1:
raise RuntimeError(f"{path}: compact registration marker missing/duplicated")
if text.count("runtime.register_generated_entry_mask(") != 1:
raise RuntimeError(f"{path}: expected one compact registration call")
origin_match = DISPATCH_ORIGIN_RE.search(text)
compact_match = COMPACT_REG_CALL_RE.search(text)
if origin_match is None or compact_match is None:
raise RuntimeError(f"{path}: compact registration/origin parse failed")
dispatch_origin = int(origin_match.group("origin"), 16)
registration_origin = int(compact_match.group("origin"), 16)
if registration_origin != dispatch_origin:
raise RuntimeError(
f"{path}: compact registration origin {registration_origin:#010x} "
f"!= dispatch origin {dispatch_origin:#010x}"
)
if compact_match.group("mask") != _mask_name or int(compact_match.group("groups")) != _groups:
raise RuntimeError(f"{path}: compact registration mask/group mismatch")
# Prove that every mask bit resolves to exactly the same PC represented
# by the unit's switch table. This catches any future origin drift.
switch_pcs = {
int(x, 16)
for x in re.findall(r"case\s+\d+u:\s+goto\s+L_([0-9A-Fa-f]{8});", text)
}
mask_pcs = {
dispatch_origin + (group * 64 + bit) * 4
for group, mask in enumerate(_masks)
for bit in range(64)
if (mask >> bit) & 1
}
if switch_pcs != mask_pcs:
missing = sorted(switch_pcs - mask_pcs)[:8]
extra = sorted(mask_pcs - switch_pcs)[:8]
raise RuntimeError(
f"{path}: compact mask/switch PC mismatch "
f"missing={[hex(x) for x in missing]} extra={[hex(x) for x in extra]}"
)
# Per-PC generated registration calls must be gone from this unit's registration function.
rm = REGISTER_BODY_RE.search(text)
if rm is None:
raise RuntimeError(f"{path}: generated registration body missing after V8.12")
if REGISTER_FN_RE.search(rm.group("body")):
raise RuntimeError(f"{path}: old per-PC registration calls remain")
compact_registration_units += 1
if total_entries != EXPECTED_REGISTERED_ENTRIES:
raise RuntimeError(
f"entry-mask population changed: expected {EXPECTED_REGISTERED_ENTRIES}, found {total_entries}"
)
if total_jumps != EXPECTED_LOCAL_LINKS:
raise RuntimeError(
f"local-link population changed: expected {EXPECTED_LOCAL_LINKS}, found {total_jumps}"
)
return {
"units": len(units),
"entry_count": total_entries,
"shared_local_jumps": total_jumps,
"shared_local_units": shared_units,
"compact_registration_units": compact_registration_units,
"generated_cpp_bytes": total_bytes,
}
def main() -> int:
parser = argparse.ArgumentParser(description=STAGE)
parser.add_argument("root", nargs="?", default=".", help="PSPRecomp repository root")
parser.add_argument("--check", action="store_true", help="validate only")
args = parser.parse_args()
root = Path(args.root).resolve()
profile = root / "profiles" / "vcs"
generated = profile / "generated"
if not (root / "CMakeLists.txt").is_file() or not generated.is_dir():
raise RuntimeError(f"not a PSPRecomp VCS tree: {root}")
manifest_path = generated / "v812_code_density_manifest.json"
if args.check:
summary = validate(root)
print(
"V8.12 CHECK OK: "
f"units={summary['units']} entries={summary['entry_count']} "
f"shared_jumps={summary['shared_local_jumps']} "
f"generated_cpp={summary['generated_cpp_bytes'] / (1024*1024):.2f} MiB"
)
return 0
units = sorted(generated.glob("generated_unit_*.cpp"))
if len(units) != EXPECTED_UNITS:
raise RuntimeError(f"expected {EXPECTED_UNITS} generated units, found {len(units)}")
before_total = sum(p.stat().st_size for p in units)
previous = {}
if manifest_path.is_file():
try:
previous = json.loads(manifest_path.read_text(encoding="utf-8"))
except Exception:
previous = {}
patch_runtime_hpp(root / "include" / "psprecomp" / "runtime.hpp")
patch_runtime_cpp(root / "src" / "runtime.cpp")
stats = [transform_unit(path) for path in units]
after_total = sum(p.stat().st_size for p in units)
baseline = int(previous.get("baseline_generated_cpp_bytes", before_total))
if baseline < after_total and previous.get("stage") != STAGE:
baseline = before_total
summary = validate(root)
manifest = {
"stage": STAGE,
"baseline_generated_cpp_bytes": baseline,
"generated_cpp_bytes_before_this_run": before_total,
"generated_cpp_bytes_after": after_total,
"source_bytes_saved_this_run": before_total - after_total,
"source_bytes_saved_from_baseline": baseline - after_total,
"registered_entries_preserved": summary["entry_count"],
"per_pc_registration_calls_collapsed_this_run": sum(s.registration_lines_removed for s in stats),
"scheduler_blocks_shared_this_run": sum(s.local_blocks_shared for s in stats),
"shared_local_jumps": summary["shared_local_jumps"],
"shared_local_units": summary["shared_local_units"],
"compact_registration_units": summary["compact_registration_units"],
"units": [asdict(s) for s in stats],
}
manifest_path.write_text(json.dumps(manifest, indent=2) + "\n", encoding="utf-8")
print(
"V8.12 code-density pass applied: "
f"units={len(stats)} entries={summary['entry_count']} "
f"registrations_collapsed={sum(s.registration_lines_removed for s in stats)} "
f"shared_boundaries={sum(s.local_blocks_shared for s in stats)} "
f"generated_cpp={after_total / (1024*1024):.2f} MiB "
f"saved={(before_total-after_total) / (1024*1024):.2f} MiB this run "
f"cumulative={(baseline-after_total) / (1024*1024):.2f} MiB"
)
return 0
if __name__ == "__main__":
raise SystemExit(main())