From 989435ff3b19f50f9cfffc84eda02757123f7b96 Mon Sep 17 00:00:00 2001 From: Seven Gao <799889633@qq.com> Date: Mon, 14 Sep 2026 19:57:22 +0800 Subject: [PATCH 1/6] feat(topic28): activate extended instruction selection with fp64 support --- scratchv/backend/asm_emit.py | 4 +- scratchv/backend/inst_select_ext.py | 171 +++++++--- scratchv/backend/instruction_select.py | 3 + scratchv/backend/riscv_encoder.py | 30 ++ scratchv/compiler.py | 40 ++- scratchv/ir/builder.py | 95 ++++++ scratchv/ir/types.py | 27 ++ scratchv/main.py | 13 + tests/test_extended_isel_cli.py | 127 +++++++ tests/test_inst_select_ext.py | 440 ++++++++++++++++++++++++- tests/test_riscv_encoder_fd.py | 88 +++++ 11 files changed, 993 insertions(+), 45 deletions(-) create mode 100644 tests/test_extended_isel_cli.py create mode 100644 tests/test_riscv_encoder_fd.py diff --git a/scratchv/backend/asm_emit.py b/scratchv/backend/asm_emit.py index 06f4b9f..ac5cb5e 100644 --- a/scratchv/backend/asm_emit.py +++ b/scratchv/backend/asm_emit.py @@ -127,7 +127,9 @@ def emit(self) -> str: def _format_instr(self, instr: MachineInstr) -> str: op_name = _OP_NAMES.get(instr.op) if op_name is None: - return f" # {instr.op.value} {instr.comment}".strip() + raise ValueError( + f"no assembly mapping for MachineOp " + f"{instr.op.name} ({instr.op.value})") # Branch/jump/call use comment as target label if instr.op in (MachineOp.CALL, MachineOp.J, MachineOp.JAL, diff --git a/scratchv/backend/inst_select_ext.py b/scratchv/backend/inst_select_ext.py index c7b138d..266f482 100644 --- a/scratchv/backend/inst_select_ext.py +++ b/scratchv/backend/inst_select_ext.py @@ -16,6 +16,7 @@ from __future__ import annotations +import struct from typing import Optional # moved import above @@ -47,6 +48,13 @@ class ExtendedInstructionSelector(InstructionSelector): If False, emit a library call to ``sqrtf``/``sqrt``. """ + # fp64-specific opcodes gated by ``enable_fp64``. + _FP64_OPCODES: frozenset[str] = frozenset({ + "load_f64", "store_f64", "load_const_f64", + "fadd_d", "fsub_d", "fmul_d", "fdiv_d", + "fcmp_l_d", "fcmp_eq_d", "fcvt_s_d", "fcvt_d_s", + }) + def __init__(self, program: Program, *, enable_fp64: bool = True, use_hardware_sqrt: bool = False): @@ -54,13 +62,23 @@ def __init__(self, program: Program, *, self.enable_fp64 = enable_fp64 self.use_hardware_sqrt = use_hardware_sqrt self._current_dtype: Optional[DataType] = None + self._temp_counter = 0 + + def run(self) -> list: + """Select instructions, resetting the temp counter for determinism.""" + self._temp_counter = 0 + return super().run() # ------------------------------------------------------------------ # Base overrides # ------------------------------------------------------------------ def _select_instruction(self, instr: Instruction) -> None: - """Override to add dtype tracking.""" + """Override to add fp64 gating and dtype tracking.""" + if not self.enable_fp64 and instr.opcode.value in self._FP64_OPCODES: + raise ValueError( + f"opcode '{instr.opcode.value}' requires enable_fp64=True " + f"(ExtendedInstructionSelector)") if instr.dest is not None: self._current_dtype = instr.dest.dtype super()._select_instruction(instr) @@ -76,6 +94,10 @@ def _select_sqrt(self, instr: Instruction) -> None: Otherwise emit a library call to ``sqrtf`` (float) or ``sqrt`` (double). """ + if self._involves_fp64(instr): + self._require_fp64(instr) + self._check_dtype( + instr, (DataType.FLOAT32, DataType.FLOAT64), "sqrt") src = self._op(instr, 0) dst = self._dst(instr) dtype = instr.dest.dtype if instr.dest else DataType.FLOAT32 @@ -89,7 +111,10 @@ def _select_sqrt(self, instr: Instruction) -> None: comment="fsqrt.s (hardware)") else: # Library call: argument in a0, result in a0 - if src and src.kind != "imm": + if src.kind == "imm": + self._emit(MachineOp.LI, MachineOperand.reg("a0"), src, + comment="sqrt arg -> a0") + else: self._emit(MachineOp.MV, MachineOperand.reg("a0"), src, comment="sqrt arg -> a0") func = "sqrt" if dtype == DataType.FLOAT64 else "sqrtf" @@ -111,20 +136,25 @@ def _select_min(self, instr: Instruction) -> None: and tmp, tmp, dst # mask = tmp & diff add dst, a, tmp # dst = a + mask """ + if self._involves_fp64(instr): + self._require_fp64(instr) + self._check_dtype( + instr, + (DataType.INT32, DataType.INT64, DataType.FLOAT64), + "min") a = self._op(instr, 0) b = self._op(instr, 1) dst = self._dst(instr) - if (self.enable_fp64 and instr.dest - and instr.dest.dtype == DataType.FLOAT64): + if instr.dest is not None and instr.dest.dtype == DataType.FLOAT64: # Use FMIN.D pseudo (expands to branchless sequence) self._emit(MachineOp.FMIN_D, dst, a, b, comment="fmin.d") else: - tmp = MachineOperand.vreg("_min_tmp1") + tmp = self._fresh_temp("min_slt") self._emit(MachineOp.SLT, tmp, a, b, comment="min: slt") - diff = MachineOperand.vreg("_min_tmp2") + diff = self._fresh_temp("min_sub") self._emit(MachineOp.SUB, diff, b, a, comment="min: sub") - and_tmp = MachineOperand.vreg("_min_tmp3") + and_tmp = self._fresh_temp("min_and") self._emit( MachineOp.AND, and_tmp, tmp, diff, comment="min: and" ) @@ -139,12 +169,17 @@ def _select_max(self, instr: Instruction) -> None: Uses the existing `max` pseudo-instruction from the base selector, or a branchless sequence if not available. """ + if self._involves_fp64(instr): + self._require_fp64(instr) + self._check_dtype( + instr, + (DataType.INT32, DataType.INT64, DataType.FLOAT64), + "max") a = self._op(instr, 0) b = self._op(instr, 1) dst = self._dst(instr) - if (self.enable_fp64 and instr.dest - and instr.dest.dtype == DataType.FLOAT64): + if instr.dest is not None and instr.dest.dtype == DataType.FLOAT64: self._emit(MachineOp.FMAX_D, dst, a, b, comment="fmax.d") else: # Use existing MAX pseudo (base selector has this) @@ -162,20 +197,25 @@ def _select_abs(self, instr: Instruction) -> None: xor dst, x, tmp # invert bits if negative sub dst, dst, tmp # add 1 if negative """ + if self._involves_fp64(instr): + self._require_fp64(instr) + self._check_dtype( + instr, + (DataType.INT32, DataType.INT64, DataType.FLOAT64), + "abs") src = self._op(instr, 0) dst = self._dst(instr) - if (self.enable_fp64 and instr.dest - and instr.dest.dtype == DataType.FLOAT64): + if instr.dest is not None and instr.dest.dtype == DataType.FLOAT64: # fabs.d: clear the sign bit self._emit(MachineOp.FABS_D, dst, src, comment="fabs.d") else: - tmp1 = MachineOperand.vreg("_abs_tmp1") + tmp1 = self._fresh_temp("abs_srai") imm31 = MachineOperand.immediate(31) self._emit( MachineOp.SRAI, tmp1, src, imm31, comment="abs: srai 31" ) - tmp2 = MachineOperand.vreg("_abs_tmp2") + tmp2 = self._fresh_temp("abs_xor") self._emit( MachineOp.XOR, tmp2, src, tmp1, comment="abs: xor" ) @@ -190,6 +230,7 @@ def _select_abs(self, instr: Instruction) -> None: def _select_idiv(self, instr: Instruction) -> None: """Select instruction for integer division.""" + self._check_dtype(instr, (DataType.INT32,), "idiv") a = self._op(instr, 0) b = self._op(instr, 1) dst = self._dst(instr) @@ -197,6 +238,7 @@ def _select_idiv(self, instr: Instruction) -> None: def _select_rem(self, instr: Instruction) -> None: """Select instruction for integer remainder.""" + self._check_dtype(instr, (DataType.INT32,), "rem") a = self._op(instr, 0) b = self._op(instr, 1) dst = self._dst(instr) @@ -204,6 +246,7 @@ def _select_rem(self, instr: Instruction) -> None: def _select_mod(self, instr: Instruction) -> None: """Select instruction for modulo (synonym of rem for non-negative).""" + self._check_dtype(instr, (DataType.INT32,), "mod") self._select_rem(instr) # ------------------------------------------------------------------ @@ -217,10 +260,19 @@ def _select_load_f64(self, instr: Instruction) -> None: self._emit(MachineOp.FLD, dst, src, comment="fld (load f64)") def _select_store_f64(self, instr: Instruction) -> None: - """Store a 64-bit float to memory.""" - val = self._op(instr, 0) - addr = self._op(instr, 1) - self._emit(MachineOp.FSD, val, addr, comment="fsd (store f64)") + """Store a 64-bit float to memory (operands: [addr, value]).""" + if len(instr.operands) < 2: + raise ValueError( + "store_f64 requires operands [addr, value], got " + f"{len(instr.operands)} operand(s)") + val = instr.operands[1] + if val.dtype != DataType.FLOAT64: + raise ValueError( + "store_f64 requires a FLOAT64 value operand, got " + f"{val.dtype.value}") + addr = self._op(instr, 0) + val_op = self._op(instr, 1) + self._emit(MachineOp.FSD, val_op, addr, comment="fsd (store f64)") def _select_fadd_d(self, instr: Instruction) -> None: """Add two float64 values.""" @@ -277,67 +329,77 @@ def _select_fcvt_d_s(self, instr: Instruction) -> None: self._emit(MachineOp.FCVT_D_S, dst, src, comment="fcvt.d.s") def _select_load_const_f64(self, instr: Instruction) -> None: - """Load a float64 constant.""" - raw_val = instr.attrs.get("value", 0.0) - assert isinstance(raw_val, (int, float)) + """Load a float64 constant (exact IEEE-754 bit pattern).""" + raw_val = instr.attrs.get("value") + if not isinstance(raw_val, (int, float)): + raise ValueError( + "load_const_f64 requires numeric attrs['value'], got " + f"{raw_val!r}") + bits = struct.unpack(" None: - if self._is_fp64(instr): + if self._involves_fp64(instr): + self._require_fp64(instr) self._select_fadd_d(instr) else: super()._select_add(instr) def _select_sub(self, instr: Instruction) -> None: - if self._is_fp64(instr): + if self._involves_fp64(instr): + self._require_fp64(instr) self._select_fsub_d(instr) else: super()._select_sub(instr) def _select_mul(self, instr: Instruction) -> None: - if self._is_fp64(instr): + if self._involves_fp64(instr): + self._require_fp64(instr) self._select_fmul_d(instr) else: super()._select_mul(instr) def _select_div(self, instr: Instruction) -> None: - if self._is_fp64(instr): + if self._involves_fp64(instr): + self._require_fp64(instr) self._select_fdiv_d(instr) - elif instr.dest and instr.dest.dtype == DataType.INT32: + elif instr.dest is not None and instr.dest.dtype == DataType.INT32: self._select_idiv(instr) else: super()._select_div(instr) def _select_load(self, instr: Instruction) -> None: - if self._is_fp64(instr): + if self._involves_fp64(instr): + self._require_fp64(instr) self._select_load_f64(instr) else: super()._select_load(instr) def _select_store(self, instr: Instruction) -> None: - if self._is_fp64(instr): + if self._involves_fp64(instr): + self._require_fp64(instr) self._select_store_f64(instr) else: super()._select_store(instr) def _select_load_const(self, instr: Instruction) -> None: - if self._is_fp64(instr): + if self._involves_fp64(instr): + self._require_fp64(instr) self._select_load_const_f64(instr) else: super()._select_load_const(instr) def _select_neg(self, instr: Instruction) -> None: """Negate: for float64 use fneg.d, for int use sub x0 - x.""" - if self._is_fp64(instr): + if self._involves_fp64(instr): + self._require_fp64(instr) src = self._op(instr, 0) dst = self._dst(instr) self._emit(MachineOp.FNEG_D, dst, src, comment="fneg.d") @@ -348,10 +410,13 @@ def _select_neg(self, instr: Instruction) -> None: # Helpers # ------------------------------------------------------------------ - def _is_fp64(self, instr: Instruction) -> bool: - """Check if an instruction operates on float64 data.""" - if not self.enable_fp64: - return False + def _fresh_temp(self, prefix: str) -> MachineOperand: + """Return a fresh deterministic virtual register temp.""" + self._temp_counter += 1 + return MachineOperand.vreg(f"__{prefix}_{self._temp_counter}") + + def _involves_fp64(self, instr: Instruction) -> bool: + """Detect float64 involvement (pure check, ignores enable_fp64).""" if instr.dest is not None and instr.dest.dtype == DataType.FLOAT64: return True for op in instr.operands: @@ -359,6 +424,34 @@ def _is_fp64(self, instr: Instruction) -> bool: return True return False + def _require_fp64(self, instr: Instruction) -> None: + """Raise if float64 support is disabled.""" + if not self.enable_fp64: + raise ValueError( + f"instruction '{instr.opcode.value}' involves FLOAT64 but " + f"enable_fp64=False (ExtendedInstructionSelector)") + + def _check_dtype(self, instr: Instruction, + allowed: tuple[DataType, ...], + opname: str) -> None: + """Guard dest/operand dtypes against the allowed set.""" + if instr.dest is None: + raise ValueError(f"{opname} requires a destination value") + allowed_vals = ", ".join(d.value for d in allowed) + if instr.dest.dtype not in allowed: + raise ValueError( + f"{opname} requires destination dtype in " + f"({allowed_vals}), got {instr.dest.dtype.value}") + for op in instr.operands: + if op.dtype not in allowed: + raise ValueError( + f"{opname} requires operand dtype in " + f"({allowed_vals}), got {op.dtype.value}") + + def _is_fp64(self, instr: Instruction) -> bool: + """Legacy compatibility: float64 involvement with fp64 enabled.""" + return self.enable_fp64 and self._involves_fp64(instr) + @property def supported_ops(self) -> list[str]: """Return list of all supported opcodes in this selector.""" @@ -382,7 +475,3 @@ def supported_ops(self) -> list[str]: base_ops + extended_ops + (fp64_ops if self.enable_fp64 else []) ) - - -# New MachineOp entries for the extended selector are now defined directly -# in the MachineOp enum in scratchv.backend.register_alloc. diff --git a/scratchv/backend/instruction_select.py b/scratchv/backend/instruction_select.py index 26395d2..5d1d5b2 100644 --- a/scratchv/backend/instruction_select.py +++ b/scratchv/backend/instruction_select.py @@ -45,6 +45,9 @@ def _select_function(self, func: Function) -> None: self._select_instruction(instr) def _select_instruction(self, instr: Instruction) -> None: + # Handler naming contract: ``_select_{opcode.value}``. New opcodes + # from Topic 28 are implemented only in the extended selector; + # this base selector raises ``ValueError`` for them. handler = getattr(self, f"_select_{instr.opcode.value}", None) if handler is None: raise ValueError( diff --git a/scratchv/backend/riscv_encoder.py b/scratchv/backend/riscv_encoder.py index d39ee4b..3de4480 100644 --- a/scratchv/backend/riscv_encoder.py +++ b/scratchv/backend/riscv_encoder.py @@ -2,6 +2,10 @@ Converts assembly text to 32-bit machine code words. Supports the subset of instructions emitted by the ScratchV compiler backend. + +F/D (single/double-precision floating point) instructions are not +supported by this encoder; encoding F/D assembly raises +``UnsupportedInstructionError`` (see Topic 28). """ from __future__ import annotations @@ -11,6 +15,27 @@ from enum import IntEnum +class UnsupportedInstructionError(ValueError): + """Raised when assembly cannot be encoded by the RV32IM encoder.""" + + +# Exact F/D mnemonics plus prefixes covering the F/D instruction families. +_FD_EXACT: frozenset[str] = frozenset({"fld", "fsd", "flw", "fsw", "li.d"}) +_FD_PREFIXES: tuple[str, ...] = ( + "fadd.", "fsub.", "fmul.", "fdiv.", "fsqrt.", "fmin.", "fmax.", + "fabs.", "fneg.", "flt.", "fle.", "feq.", "fcvt.", "fmv.", + "fsgnj", "fsgnjn", "fsgnjx", +) + + +def _is_fd_mnemonic(op: str) -> bool: + """Return True if *op* is an F/D-extension instruction mnemonic.""" + op = op.lower() + if op in _FD_EXACT: + return True + return any(op.startswith(prefix) for prefix in _FD_PREFIXES) + + # ── RISC-V opcodes ──────────────────────────────────────────────────── class RVOpcode(IntEnum): @@ -489,6 +514,11 @@ def _encode_line( elif op == "nop": word = _i_type(0, 0, 0, F3_ADD_SUB) else: + if _is_fd_mnemonic(op): + raise UnsupportedInstructionError( + f"F/D instruction '{op}' is not supported by the " + f"RV32IM encoder (Topic 28: final encoding out of " + f"scope; output is assembly text)") raise ValueError(f"Unknown instruction: {op}") return (word, fixup) diff --git a/scratchv/compiler.py b/scratchv/compiler.py index fa5459e..f012eb8 100644 --- a/scratchv/compiler.py +++ b/scratchv/compiler.py @@ -53,6 +53,10 @@ class CompilerConfig: cycle_stats: Run 5-stage pipeline cycle estimation (detailed). enable_forwarding: Enable forwarding in cycle estimator. branch_predictor: Branch predictor mode for cycle estimator. + extended_isel: Use the extended instruction selector (Topic 28). + enable_fp64: Enable float64 (D extension) support (Topic 28). + use_hardware_sqrt: Use ``fsqrt.s``/``fsqrt.d`` instead of libm + calls (Topic 28). """ backend: str = "riscv" @@ -73,6 +77,9 @@ class CompilerConfig: cycle_stats: bool = False enable_forwarding: bool = True branch_predictor: str = "always_not_taken" + extended_isel: bool = False + enable_fp64: bool = True + use_hardware_sqrt: bool = False # ═══════════════════════════════════════════════════════════════════════════════ @@ -315,6 +322,7 @@ def compile(self, input_path: str, output_path: str | None = None, ) # --- 4. Code generation --- + self._collect_selector_warnings(warnings) try: asm_text = self._generate_code(program) except Exception as e: @@ -428,13 +436,41 @@ def _generate_code(self, program) -> str: return self._generate_riscv_dag(program) return self._generate_riscv_linear(program) + def _collect_selector_warnings(self, warnings: list[str]) -> None: + """Append mode-conflict warnings for selector-related options.""" + if self.config.extended_isel: + if self.config.use_dag_isel: + warnings.append( + "--dag-isel takes precedence; --extended-isel ignored") + if self.config.backend == "llvm": + warnings.append( + "--extended-isel is RISC-V only; ignored for LLVM " + "backend") + elif (not self.config.enable_fp64 + or self.config.use_hardware_sqrt): + warnings.append( + "--no-fp64/--hardware-sqrt have no effect without " + "--extended-isel") + def _generate_riscv_linear(self, program) -> str: """Standard RISC-V pipeline.""" - from scratchv.backend.instruction_select import InstructionSelector from scratchv.backend.register_alloc import RegisterAllocator from scratchv.backend.asm_emit import AsmEmitter - selector = InstructionSelector(program) + if self.config.extended_isel: + from scratchv.backend.inst_select_ext import ( + ExtendedInstructionSelector, + ) + selector = ExtendedInstructionSelector( + program, + enable_fp64=self.config.enable_fp64, + use_hardware_sqrt=self.config.use_hardware_sqrt, + ) + else: + from scratchv.backend.instruction_select import ( + InstructionSelector, + ) + selector = InstructionSelector(program) machine_instrs = selector.run() # Linear-scan: skip greedy allocator, use liveness-driven allocator diff --git a/scratchv/ir/builder.py b/scratchv/ir/builder.py index 5468962..cf9ce49 100644 --- a/scratchv/ir/builder.py +++ b/scratchv/ir/builder.py @@ -207,3 +207,98 @@ def reshape(self, val: Value, shape: tuple) -> Value: dest = self.make_value() self._emit(OpCode.RESHAPE, dest, [val], shape=shape) return dest + + # --- [Topic 28] Extended instruction selection --- + + def sqrt(self, val: Value, + dtype: DataType = DataType.FLOAT32) -> Value: + dest = self.make_value(dtype=dtype) + self._emit(OpCode.SQRT, dest, [val]) + return dest + + def min(self, a: Value, b: Value, + dtype: DataType = DataType.INT32) -> Value: + dest = self.make_value(dtype=dtype) + self._emit(OpCode.MIN, dest, [a, b]) + return dest + + def max(self, a: Value, b: Value, + dtype: DataType = DataType.INT32) -> Value: + dest = self.make_value(dtype=dtype) + self._emit(OpCode.MAX, dest, [a, b]) + return dest + + def abs(self, val: Value, + dtype: DataType = DataType.INT32) -> Value: + dest = self.make_value(dtype=dtype) + self._emit(OpCode.ABS, dest, [val]) + return dest + + def idiv(self, a: Value, b: Value) -> Value: + dest = self.make_value(dtype=DataType.INT32) + self._emit(OpCode.IDIV, dest, [a, b]) + return dest + + def rem(self, a: Value, b: Value) -> Value: + dest = self.make_value(dtype=DataType.INT32) + self._emit(OpCode.REM, dest, [a, b]) + return dest + + def mod(self, a: Value, b: Value) -> Value: + dest = self.make_value(dtype=DataType.INT32) + self._emit(OpCode.MOD, dest, [a, b]) + return dest + + def load_f64(self, addr: Value) -> Value: + dest = self.make_value(dtype=DataType.FLOAT64) + self._emit(OpCode.LOAD_F64, dest, [addr]) + return dest + + def store_f64(self, addr: Value, val: Value) -> Instruction: + return self._emit(OpCode.STORE_F64, operands=[addr, val]) + + def load_const_f64(self, value: float) -> Value: + dest = self.make_value(dtype=DataType.FLOAT64, + is_constant=True, const_value=value) + self._emit(OpCode.LOAD_CONST_F64, dest, value=value) + return dest + + def fadd_d(self, a: Value, b: Value) -> Value: + dest = self.make_value(dtype=DataType.FLOAT64) + self._emit(OpCode.FADD_D, dest, [a, b]) + return dest + + def fsub_d(self, a: Value, b: Value) -> Value: + dest = self.make_value(dtype=DataType.FLOAT64) + self._emit(OpCode.FSUB_D, dest, [a, b]) + return dest + + def fmul_d(self, a: Value, b: Value) -> Value: + dest = self.make_value(dtype=DataType.FLOAT64) + self._emit(OpCode.FMUL_D, dest, [a, b]) + return dest + + def fdiv_d(self, a: Value, b: Value) -> Value: + dest = self.make_value(dtype=DataType.FLOAT64) + self._emit(OpCode.FDIV_D, dest, [a, b]) + return dest + + def fcmp_l_d(self, a: Value, b: Value) -> Value: + dest = self.make_value(dtype=DataType.INT32) + self._emit(OpCode.FCMP_L_D, dest, [a, b]) + return dest + + def fcmp_eq_d(self, a: Value, b: Value) -> Value: + dest = self.make_value(dtype=DataType.INT32) + self._emit(OpCode.FCMP_EQ_D, dest, [a, b]) + return dest + + def fcvt_s_d(self, val: Value) -> Value: + dest = self.make_value(dtype=DataType.FLOAT32) + self._emit(OpCode.FCVT_S_D, dest, [val]) + return dest + + def fcvt_d_s(self, val: Value) -> Value: + dest = self.make_value(dtype=DataType.FLOAT64) + self._emit(OpCode.FCVT_D_S, dest, [val]) + return dest diff --git a/scratchv/ir/types.py b/scratchv/ir/types.py index 059f73d..ad490ca 100644 --- a/scratchv/ir/types.py +++ b/scratchv/ir/types.py @@ -49,6 +49,33 @@ class OpCode(enum.Enum): RESHAPE = "reshape" CONCAT = "concat" + # ───────────────────────────────────────────────────────────────── + # [Topic 28] Extended instruction selection — append-only partition. + # 约定:新成员只能追加在本分区末尾;课题 29 及后续课题必须追加在 + # 本分区之后,禁止插入本分区或重排既有成员(保护序列化与并行合入)。 + # Dispatch 契约:_select_{value}(见 backend/instruction_select.py)。 + # ───────────────────────────────────────────────────────────────── + # 整数 / 通用扩展 + SQRT = "sqrt" + MIN = "min" + MAX = "max" + ABS = "abs" + IDIV = "idiv" + REM = "rem" + MOD = "mod" + # float64(D 扩展) + LOAD_F64 = "load_f64" + STORE_F64 = "store_f64" + LOAD_CONST_F64 = "load_const_f64" + FADD_D = "fadd_d" + FSUB_D = "fsub_d" + FMUL_D = "fmul_d" + FDIV_D = "fdiv_d" + FCMP_L_D = "fcmp_l_d" + FCMP_EQ_D = "fcmp_eq_d" + FCVT_S_D = "fcvt_s_d" + FCVT_D_S = "fcvt_d_s" + def is_arith(self) -> bool: return self in (OpCode.ADD, OpCode.SUB, OpCode.MUL, OpCode.DIV) diff --git a/scratchv/main.py b/scratchv/main.py index 52feff5..21f588b 100644 --- a/scratchv/main.py +++ b/scratchv/main.py @@ -104,6 +104,16 @@ def build_arg_parser() -> argparse.ArgumentParser: "--extended-isel", action="store_true", help="Use extended instruction selector with fp64/sqrt/min/max/abs support (Topic 28)", ) + parser.add_argument( + "--no-fp64", dest="enable_fp64", action="store_false", + help="Disable float64 (D extension) support " + "(requires --extended-isel)", + ) + parser.add_argument( + "--hardware-sqrt", dest="use_hardware_sqrt", action="store_true", + help="Use fsqrt.s/fsqrt.d instead of libm calls " + "(requires --extended-isel)", + ) # ── Cycle estimation ────────────────────────────────────────────── parser.add_argument( @@ -153,6 +163,9 @@ def args_to_config(args: argparse.Namespace) -> CompilerConfig: cycle_stats=args.cycle_stats, enable_forwarding=not args.no_forwarding, branch_predictor=args.branch_predictor, + extended_isel=args.extended_isel, + enable_fp64=args.enable_fp64, + use_hardware_sqrt=args.use_hardware_sqrt, ) diff --git a/tests/test_extended_isel_cli.py b/tests/test_extended_isel_cli.py new file mode 100644 index 0000000..91c5ab1 --- /dev/null +++ b/tests/test_extended_isel_cli.py @@ -0,0 +1,127 @@ +"""Tests for --extended-isel CLI/config wiring (Topic 28).""" + +from unittest import mock + +from scratchv.compiler import CompilerConfig, CompilerDriver +from scratchv.main import args_to_config, build_arg_parser + +DSL_SOURCE = "y = add(a, b)\nreturn y\n" + + +def _write_dsl(tmp_path): + path = tmp_path / "model.dsl" + path.write_text(DSL_SOURCE) + return str(path) + + +def test_cli_flags_map_to_config(): + args = build_arg_parser().parse_args([ + "-o", "x", "--dsl", "a", + "--extended-isel", "--no-fp64", "--hardware-sqrt", + ]) + config = args_to_config(args) + assert config.extended_isel is True + assert config.enable_fp64 is False + assert config.use_hardware_sqrt is True + + +def test_cli_defaults_map_to_config(): + args = build_arg_parser().parse_args(["input.dsl"]) + config = args_to_config(args) + assert config.extended_isel is False + assert config.enable_fp64 is True + assert config.use_hardware_sqrt is False + + +def test_driver_uses_extended_selector(tmp_path): + inp = _write_dsl(tmp_path) + out = str(tmp_path / "out.s") + with mock.patch( + "scratchv.backend.inst_select_ext.ExtendedInstructionSelector" + ) as selector_cls: + selector_cls.return_value.run.return_value = [] + result = CompilerDriver( + CompilerConfig(extended_isel=True)).compile(inp, out) + + assert result.success, result.errors + selector_cls.assert_called_once() + assert selector_cls.call_args.kwargs == { + "enable_fp64": True, "use_hardware_sqrt": False, + } + + +def test_driver_passes_fp64_flags_to_selector(tmp_path): + inp = _write_dsl(tmp_path) + out = str(tmp_path / "out.s") + with mock.patch( + "scratchv.backend.inst_select_ext.ExtendedInstructionSelector" + ) as selector_cls: + selector_cls.return_value.run.return_value = [] + result = CompilerDriver(CompilerConfig( + extended_isel=True, enable_fp64=False, + use_hardware_sqrt=True)).compile(inp, out) + + assert result.success, result.errors + assert selector_cls.call_args.kwargs == { + "enable_fp64": False, "use_hardware_sqrt": True, + } + + +def test_default_uses_base_selector(tmp_path): + inp = _write_dsl(tmp_path) + out = str(tmp_path / "out.s") + with mock.patch( + "scratchv.backend.instruction_select.InstructionSelector" + ) as base_cls, mock.patch( + "scratchv.backend.inst_select_ext.ExtendedInstructionSelector" + ) as ext_cls: + base_cls.return_value.run.return_value = [] + result = CompilerDriver(CompilerConfig()).compile(inp, out) + + assert result.success, result.errors + base_cls.assert_called_once() + ext_cls.assert_not_called() + + +def test_dag_isel_conflict_warns(tmp_path): + inp = _write_dsl(tmp_path) + out = str(tmp_path / "out.s") + driver = CompilerDriver( + CompilerConfig(extended_isel=True, use_dag_isel=True)) + with mock.patch.object(driver, "_generate_code", return_value=""): + result = driver.compile(inp, out) + + assert result.success, result.errors + assert any("precedence" in w for w in result.warnings) + + +def test_llvm_backend_warns(tmp_path): + inp = _write_dsl(tmp_path) + out = str(tmp_path / "out.ll") + driver = CompilerDriver( + CompilerConfig(extended_isel=True, backend="llvm")) + with mock.patch.object(driver, "_generate_code", return_value=""): + result = driver.compile(inp, out) + + assert result.success, result.errors + assert any("RISC-V only" in w for w in result.warnings) + + +def test_fp64_flags_without_extended_warn(tmp_path): + inp = _write_dsl(tmp_path) + out = str(tmp_path / "out.s") + result = CompilerDriver( + CompilerConfig(enable_fp64=False)).compile(inp, out) + + assert result.success, result.errors + assert any("no effect" in w for w in result.warnings) + + +def test_hardware_sqrt_without_extended_warns(tmp_path): + inp = _write_dsl(tmp_path) + out = str(tmp_path / "out.s") + result = CompilerDriver( + CompilerConfig(use_hardware_sqrt=True)).compile(inp, out) + + assert result.success, result.errors + assert any("no effect" in w for w in result.warnings) diff --git a/tests/test_inst_select_ext.py b/tests/test_inst_select_ext.py index 911f4fb..a136a39 100644 --- a/tests/test_inst_select_ext.py +++ b/tests/test_inst_select_ext.py @@ -1,9 +1,12 @@ """Tests for Extended Instruction Selector.""" import pytest +from scratchv.backend.asm_emit import AsmEmitter from scratchv.backend.inst_select_ext import ExtendedInstructionSelector from scratchv.ir.builder import IRBuilder -from scratchv.ir.types import Value, DataType # noqa: F401 +from scratchv.ir.types import ( # noqa: F401 + OpCode, Value, DataType, +) from scratchv.backend.register_alloc import MachineOp @@ -194,5 +197,440 @@ def test_load_store(self): assert MachineOp.SW in ops +# ═══════════════════════════════════════════════════════════════════════ +# [Topic 28] Extended instruction selection +# ═══════════════════════════════════════════════════════════════════════ + +NEW_OPCODES = [ + "sqrt", "min", "max", "abs", "idiv", "rem", "mod", + "load_f64", "store_f64", "load_const_f64", + "fadd_d", "fsub_d", "fmul_d", "fdiv_d", + "fcmp_l_d", "fcmp_eq_d", "fcvt_s_d", "fcvt_d_s", +] + +DISPATCH_EXPECTED_OP = { + "sqrt": MachineOp.CALL, + "min": MachineOp.SLT, + "max": MachineOp.MAX, + "abs": MachineOp.SRAI, + "idiv": MachineOp.DIV, + "rem": MachineOp.REM, + "mod": MachineOp.REM, + "load_f64": MachineOp.FLD, + "store_f64": MachineOp.FSD, + "load_const_f64": MachineOp.LI_D, + "fadd_d": MachineOp.FADD_D, + "fsub_d": MachineOp.FSUB_D, + "fmul_d": MachineOp.FMUL_D, + "fdiv_d": MachineOp.FDIV_D, + "fcmp_l_d": MachineOp.FLT_D, + "fcmp_eq_d": MachineOp.FEQ_D, + "fcvt_s_d": MachineOp.FCVT_S_D, + "fcvt_d_s": MachineOp.FCVT_D_S, +} + + +def _make_single_op_program(op): + """Build a minimal valid program containing one new-op instruction.""" + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + i32 = DataType.INT32 + f32 = DataType.FLOAT32 + f64 = DataType.FLOAT64 + a = builder.make_value(name="a", dtype=i32) + b = builder.make_value(name="b", dtype=i32) + x = builder.make_value(name="x", dtype=f32) + dx = builder.make_value(name="dx", dtype=f64) + dy = builder.make_value(name="dy", dtype=f64) + + if op == "sqrt": + builder.sqrt(x) + elif op == "min": + builder.min(a, b) + elif op == "max": + builder.max(a, b) + elif op == "abs": + builder.abs(a) + elif op == "idiv": + builder.idiv(a, b) + elif op == "rem": + builder.rem(a, b) + elif op == "mod": + builder.mod(a, b) + elif op == "load_f64": + builder.load_f64(a) + elif op == "store_f64": + builder.store_f64(a, dx) + elif op == "load_const_f64": + builder.load_const_f64(1.5) + elif op == "fadd_d": + builder.fadd_d(dx, dy) + elif op == "fsub_d": + builder.fsub_d(dx, dy) + elif op == "fmul_d": + builder.fmul_d(dx, dy) + elif op == "fdiv_d": + builder.fdiv_d(dx, dy) + elif op == "fcmp_l_d": + builder.fcmp_l_d(dx, dy) + elif op == "fcmp_eq_d": + builder.fcmp_eq_d(dx, dy) + elif op == "fcvt_s_d": + builder.fcvt_s_d(dx) + elif op == "fcvt_d_s": + builder.fcvt_d_s(x) + else: + raise AssertionError(f"unhandled opcode {op}") + return builder + + +def _clean_asm_lines(asm: str) -> list: + """Strip comments and blank lines from emitted assembly.""" + return [ + line.split("#")[0].strip() + for line in asm.splitlines() + if line.split("#")[0].strip() + ] + + +class TestDispatchCoverage: + """Every new opcode must dispatch to its handler (Topic 28).""" + + def test_new_opcode_handler_exists(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + selector = ExtendedInstructionSelector(builder.program) + for value in NEW_OPCODES: + assert hasattr(selector, f"_select_{value}"), value + + def test_handler_and_opcode_one_to_one(self): + enum_values = {op.value for op in OpCode} + handlers = { + name[len("_select_"):] + for name in dir(ExtendedInstructionSelector) + if name.startswith("_select_") + } + for value in NEW_OPCODES: + assert value in enum_values + assert value in handlers + # No [Topic 28] handler without a matching OpCode member. + new_handlers = {h for h in handlers if h in set(NEW_OPCODES)} + assert new_handlers == set(NEW_OPCODES) + + @pytest.mark.parametrize("op", NEW_OPCODES) + def test_dispatch_table(self, op): + builder = _make_single_op_program(op) + selector = ExtendedInstructionSelector(builder.program) + instrs = selector.run() + ops = {i.op for i in instrs if i.op != MachineOp.LABEL} + assert DISPATCH_EXPECTED_OP[op] in ops + + def test_unique_temps_two_mins(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + a = builder.make_value(name="a", dtype=DataType.INT32) + b = builder.make_value(name="b", dtype=DataType.INT32) + c = builder.make_value(name="c", dtype=DataType.INT32) + d = builder.make_value(name="d", dtype=DataType.INT32) + builder.min(a, b) + builder.min(c, d) + + instrs = ExtendedInstructionSelector(builder.program).run() + names = [ + i.dst.value for i in instrs + if i.dst is not None and i.dst.value.startswith("__min") + ] + assert len(names) == 6 + assert len(names) == len(set(names)) + assert sorted(names, key=lambda n: int(n.rsplit("_", 1)[1])) == names + + +class TestAsmText: + """Assembly text for MIN/ABS/SQRT/f64 sequences (Topic 28).""" + + def test_min_branchless_asm(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + a = builder.make_value(name="a", dtype=DataType.INT32) + b = builder.make_value(name="b", dtype=DataType.INT32) + dest = builder.min(a, b) + + instrs = ExtendedInstructionSelector(builder.program).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + seq = [ln for ln in lines if "__min" in ln] + assert seq == [ + "slt __min_slt_1, a, b", + "sub __min_sub_2, b, a", + "and __min_and_3, __min_slt_1, __min_sub_2", + f"add {dest.name}, a, __min_and_3", + ] + + def test_abs_branchless_asm(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + x = builder.make_value(name="x", dtype=DataType.INT32) + dest = builder.abs(x) + + instrs = ExtendedInstructionSelector(builder.program).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + seq = [ln for ln in lines if "__abs" in ln] + assert seq == [ + "srai __abs_srai_1, x, 31", + "xor __abs_xor_2, x, __abs_srai_1", + f"sub {dest.name}, __abs_xor_2, __abs_srai_1", + ] + + def test_sqrt_software_uses_a0_and_call(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + x = builder.make_value(name="x", dtype=DataType.FLOAT32) + dest = builder.sqrt(x) + + instrs = ExtendedInstructionSelector(builder.program).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + idx = lines.index("mv a0, x") + assert lines[idx + 1] == "call sqrtf" + assert lines[idx + 2] == f"mv {dest.name}, a0" + + def test_sqrt_software_f64_calls_sqrt(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + builder.sqrt(dx, dtype=DataType.FLOAT64) + + instrs = ExtendedInstructionSelector(builder.program).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + assert "call sqrt" in lines + + def test_sqrt_immediate_uses_li(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + imm = builder.make_const(4.0, dtype=DataType.FLOAT32) + builder.sqrt(imm) + + instrs = ExtendedInstructionSelector(builder.program).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + assert "li a0, 4" in lines + + def test_sqrt_hardware_f32_f64(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + x = builder.make_value(name="x", dtype=DataType.FLOAT32) + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + d32 = builder.sqrt(x) + d64 = builder.sqrt(dx, dtype=DataType.FLOAT64) + + instrs = ExtendedInstructionSelector( + builder.program, use_hardware_sqrt=True).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + assert f"fsqrt.s {d32.name}, x" in lines + assert f"fsqrt.d {d64.name}, dx" in lines + + def test_fp64_dtype_driven_add(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + dy = builder.make_value(name="dy", dtype=DataType.FLOAT64) + dest = builder.make_value(name="dd", dtype=DataType.FLOAT64) + builder._emit(OpCode.ADD, dest, [dx, dy]) + + instrs = ExtendedInstructionSelector(builder.program).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + assert f"fadd.d {dest.name}, dx, dy" in lines + + def test_store_f64_operand_order(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + ptr = builder.make_value(name="p", dtype=DataType.INT32) + val = builder.make_value(name="v", dtype=DataType.FLOAT64) + builder.store_f64(ptr, val) + + instrs = ExtendedInstructionSelector(builder.program).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + assert "fsd v, p" in lines + + +class TestFp64Gate: + """fail-loud when enable_fp64=False (Topic 28).""" + + def test_fadd_d_without_fp64_raises(self): + builder = _make_single_op_program("fadd_d") + selector = ExtendedInstructionSelector( + builder.program, enable_fp64=False) + with pytest.raises(ValueError, match="enable_fp64"): + selector.run() + + def test_f64_add_without_fp64_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + dy = builder.make_value(name="dy", dtype=DataType.FLOAT64) + dest = builder.make_value(name="dd", dtype=DataType.FLOAT64) + builder._emit(OpCode.ADD, dest, [dx, dy]) + + selector = ExtendedInstructionSelector( + builder.program, enable_fp64=False) + with pytest.raises(ValueError, match="enable_fp64"): + selector.run() + + def test_f64_min_without_fp64_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + dy = builder.make_value(name="dy", dtype=DataType.FLOAT64) + builder.min(dx, dy, dtype=DataType.FLOAT64) + + selector = ExtendedInstructionSelector( + builder.program, enable_fp64=False) + with pytest.raises(ValueError, match="enable_fp64"): + selector.run() + + def test_f64_sqrt_without_fp64_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + builder.sqrt(dx, dtype=DataType.FLOAT64) + + selector = ExtendedInstructionSelector( + builder.program, enable_fp64=False) + with pytest.raises(ValueError, match="enable_fp64"): + selector.run() + + def test_integer_ops_without_fp64_still_work(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + a = builder.make_value(name="a", dtype=DataType.INT32) + b = builder.make_value(name="b", dtype=DataType.INT32) + builder.add(a, b) + + selector = ExtendedInstructionSelector( + builder.program, enable_fp64=False) + instrs = selector.run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + assert any(ln.startswith("add ") for ln in lines) + + +class TestIllegalDtype: + """dtype guards raise ValueError with the opcode name (Topic 28).""" + + def test_sqrt_int_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + a = builder.make_value(name="a", dtype=DataType.INT32) + builder.sqrt(a, dtype=DataType.INT32) + with pytest.raises(ValueError, match="sqrt"): + ExtendedInstructionSelector(builder.program).run() + + def test_min_f32_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + x = builder.make_value(name="x", dtype=DataType.FLOAT32) + y = builder.make_value(name="y", dtype=DataType.FLOAT32) + builder.min(x, y, dtype=DataType.FLOAT32) + with pytest.raises(ValueError, match="min"): + ExtendedInstructionSelector(builder.program).run() + + def test_idiv_f64_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + dy = builder.make_value(name="dy", dtype=DataType.FLOAT64) + builder.idiv(dx, dy) + with pytest.raises(ValueError, match="idiv"): + ExtendedInstructionSelector(builder.program).run() + + +class TestLoadConstF64: + """Exact IEEE-754 bit patterns for f64 constants (Topic 28).""" + + @pytest.mark.parametrize("value,bits", [ + (1.5, 4609434218613702656), + (2.0, 4611686018427387904), + (-0.0, 9223372036854775808), + ]) + def test_exact_bits(self, value, bits): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + builder.load_const_f64(value) + + instrs = ExtendedInstructionSelector(builder.program).run() + li_d = [i for i in instrs if i.op == MachineOp.LI_D] + assert len(li_d) == 1 + assert li_d[0].src1.value == bits + + def test_missing_value_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dest = builder.make_value(dtype=DataType.FLOAT64) + builder._emit(OpCode.LOAD_CONST_F64, dest, attrs={}) + with pytest.raises(ValueError, match="load_const_f64"): + ExtendedInstructionSelector(builder.program).run() + + def test_non_numeric_value_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dest = builder.make_value(dtype=DataType.FLOAT64) + builder._emit(OpCode.LOAD_CONST_F64, dest, value="nope") + with pytest.raises(ValueError, match="load_const_f64"): + ExtendedInstructionSelector(builder.program).run() + + +class TestAsmEmitterFailLoud: + """AsmEmitter must not silently drop unmapped MachineOps (Topic 28).""" + + def test_unknown_machine_op_raises(self): + from scratchv.backend.machine_types import ( + MachineInstr, MachineOp as MOp, + ) + instrs = [MachineInstr(MOp.GLOBL, comment="foo")] + with pytest.raises(ValueError, match="GLOBL"): + AsmEmitter(instrs).emit() + + +class TestSupportedOpsNew: + """supported_ops reflects the 18 new opcodes (Topic 28).""" + + def test_new_ops_present(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + ops = ExtendedInstructionSelector(builder.program).supported_ops + for value in NEW_OPCODES: + assert value in ops + + def test_fp64_ops_gated(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + ops = ExtendedInstructionSelector( + builder.program, enable_fp64=False).supported_ops + for value in ("sqrt", "min", "max", "abs", + "idiv", "rem", "mod"): + assert value in ops + for value in ("load_f64", "fadd_d", "fcvt_d_s"): + assert value not in ops + + if __name__ == "__main__": pytest.main([__file__, "-v"]) diff --git a/tests/test_riscv_encoder_fd.py b/tests/test_riscv_encoder_fd.py new file mode 100644 index 0000000..1c8aea5 --- /dev/null +++ b/tests/test_riscv_encoder_fd.py @@ -0,0 +1,88 @@ +"""Tests for RV32IM encoder fail-loud on F/D instructions (Topic 28).""" + +import pytest + +from scratchv.backend.riscv_encoder import ( + UnsupportedInstructionError, + _is_fd_mnemonic, + assemble_to_binary, +) + + +FD_SAMPLES = [ + "fadd.d f0, f1, f2", + "fsub.s f0, f1, f2", + "fmul.d f0, f1, f2", + "fdiv.d f0, f1, f2", + "fsqrt.d f0, f1", + "fmin.d f0, f1, f2", + "fmax.d f0, f1, f2", + "fabs.d f0, f1", + "fneg.d f0, f1", + "fld f0, 0(sp)", + "fsd f0, 0(sp)", + "flw f0, 0(sp)", + "fsw f0, 0(sp)", + "flt.d a0, f1, f2", + "feq.d a0, f1, f2", + "fcvt.s.d f0, f1", + "fcvt.d.s f0, f1", + "fmv.x.w a0, f1", + "li.d f0, 123", +] + + +@pytest.mark.parametrize("asm", FD_SAMPLES) +def test_fd_mnemonics_raise(asm): + mnemonic = asm.split()[0] + with pytest.raises(UnsupportedInstructionError) as excinfo: + assemble_to_binary(asm) + assert mnemonic in str(excinfo.value) + + +def test_unsupported_is_value_error(): + assert issubclass(UnsupportedInstructionError, ValueError) + + +def test_unknown_non_fd_still_value_error(): + with pytest.raises(ValueError) as excinfo: + assemble_to_binary("frobnicate x1, x2") + assert not isinstance(excinfo.value, UnsupportedInstructionError) + assert "Unknown instruction" in str(excinfo.value) + + +@pytest.mark.parametrize("mnemonic,expected", [ + ("add", False), + ("fence", False), + ("frobnicate", False), + ("flw", True), + ("fld", True), + ("li.d", True), + ("fadd.d", True), + ("fsqrt.s", True), + ("fmv.x.w", True), + ("fsgnjx.d", True), +]) +def test_is_fd_mnemonic_predicate(mnemonic, expected): + assert _is_fd_mnemonic(mnemonic) is expected + + +def test_rv32im_regression_still_assembles(): + asm = "\n".join([ + " addi a0, x0, 5", + " add a1, a0, a0", + " sw a1, 0(sp)", + " li t0, 4098", + " beq a0, a1, .Ldone", + " nop", + ".Ldone:", + " ret", + ]) + result = assemble_to_binary(asm) + assert isinstance(result, bytearray) + # 8 encoded words: addi, add, sw, lui+addi (li), beq, nop, ret + assert len(result) == 8 * 4 + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) From 5431568919ac4d2f13eceb3b8672dda984eb1483 Mon Sep 17 00:00:00 2001 From: Seven Gao <799889633@qq.com> Date: Mon, 14 Sep 2026 21:22:01 +0800 Subject: [PATCH 2/6] docs(topic28): add design and development documents --- ...00\345\217\221\346\226\207\346\241\243.md" | 520 +++++++++++++++++ ...76\350\256\241\346\226\207\346\241\243.md" | 527 ++++++++++++++++++ 2 files changed, 1047 insertions(+) create mode 100644 "docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\345\274\200\345\217\221\346\226\207\346\241\243.md" create mode 100644 "docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\350\256\276\350\256\241\346\226\207\346\241\243.md" diff --git "a/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\345\274\200\345\217\221\346\226\207\346\241\243.md" "b/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\345\274\200\345\217\221\346\226\207\346\241\243.md" new file mode 100644 index 0000000..2b509c5 --- /dev/null +++ "b/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\345\274\200\345\217\221\346\226\207\346\241\243.md" @@ -0,0 +1,520 @@ +# 课题28 扩展指令选择开发文档 + +> 文档版本:v1.0 +> 编写日期:2026-09-14 +> 涉及模块:`scratchv/ir/types.py`、`scratchv/ir/builder.py`、`scratchv/backend/instruction_select.py`、`scratchv/backend/inst_select_ext.py`、`scratchv/backend/machine_types.py`、`scratchv/backend/asm_emit.py`、`scratchv/backend/riscv_encoder.py`、`scratchv/compiler.py`、`scratchv/main.py`、`tests/*` +> 配套文档:《设计文档.md》(同目录)——本文件是其实现指南,二者若冲突以设计文档为准 +> 前置环境:Python 3.11+(项目 venv;本机可用 `/usr/local/bin/python3.11`)、pytest、flake8/mypy(`make lint`) +> 基线事实:2026-09-14 实测——`OpCode` 无 sqrt/min/max/abs/idiv/rem/mod/fp64 成员;`--extended-isel` 未进 config;`riscv_encoder` 对 F/D 抛 `Unknown instruction`;仅 `DIV` 的 INT32 分支生效 + +--- + +## 一、实施范围与总览 + +### 1.1 做 + +1. `ir/types.py`:追加 18 个 `OpCode`(末尾注释分区,append-only)。 +2. `ir/builder.py`:新增 18 个构造方法。 +3. `backend/inst_select_ext.py`:fp64 门控、dtype 守卫、唯一临时值、`store_f64` 顺序、sqrt 立即数、常量位模式、失效注释清理(P1–P7)。 +4. `backend/asm_emit.py`:未知 MachineOp 改为 fail-loud。 +5. `backend/riscv_encoder.py`:新增 `UnsupportedInstructionError`,F/D 助记符明确报错。 +6. `compiler.py` + `main.py`:`extended_isel/enable_fp64/use_hardware_sqrt` 字段与 CLI 接线、模式警告。 +7. 测试:扩展 `tests/test_inst_select_ext.py`;新增 `tests/test_riscv_encoder_fd.py`、`tests/test_extended_isel_cli.py`。 + +### 1.2 不做 + +- 不实现 F/D 机器码编码(编码器对 F/D 明确报错)。 +- 不做 RVV(课题 29)。 +- 不做 f-register 寄存器分配、f64 可执行数值验证、常量池。 +- 不修改 DSL/ONNX 前端;不修改 `instruction_select.py` 的分发逻辑;不新增第三方依赖。 +- 不改动 `MachineOp` 枚举成员(仅盘点确认)。 + +### 1.3 改动依赖顺序 + +`ir/types.py` → `ir/builder.py` → `inst_select_ext.py` → `asm_emit.py` / `riscv_encoder.py` → `compiler.py` / `main.py` → 测试。每步保持 import 可用,避免中间态破坏既有测试。 + +--- + +## 二、接口契约(精确名称,实现/评审以此为准) + +### 2.1 IR OpCode 最终定名与分区(`scratchv/ir/types.py`) + +**精确插入位置**:当前 `OpCode` 枚举第 50 行 `CONCAT = "concat"` 之后、第 52 行 `def is_arith(self)` 之前。禁止改动第 14–50 行既有成员、禁止改动 `DataType`。 + +```python + # Shape / data movement + TRANSPOSE = "transpose" + RESHAPE = "reshape" + CONCAT = "concat" + + # ───────────────────────────────────────────────────────────────── + # [Topic 28] Extended instruction selection — append-only partition. + # 约定:新成员只能追加在本分区末尾;课题 29 及后续课题必须追加在 + # 本分区之后,禁止插入本分区或重排既有成员(保护序列化与并行合入)。 + # Dispatch 契约:_select_{value}(见 backend/instruction_select.py)。 + # ───────────────────────────────────────────────────────────────── + # 整数 / 通用扩展 + SQRT = "sqrt" + MIN = "min" + MAX = "max" + ABS = "abs" + IDIV = "idiv" + REM = "rem" + MOD = "mod" + # float64(D 扩展) + LOAD_F64 = "load_f64" + STORE_F64 = "store_f64" + LOAD_CONST_F64 = "load_const_f64" + FADD_D = "fadd_d" + FSUB_D = "fsub_d" + FMUL_D = "fmul_d" + FDIV_D = "fdiv_d" + FCMP_L_D = "fcmp_l_d" + FCMP_EQ_D = "fcmp_eq_d" + FCVT_S_D = "fcvt_s_d" + FCVT_D_S = "fcvt_d_s" +``` + +不变式:`OpCode..value` 必须与 `ExtendedInstructionSelector` 的 `_select_` 后缀一一对应(§2.3.1 全表)。`is_arith()/is_nn()/is_control_flow()` 不改。 + +### 2.2 IRBuilder 方法签名(`scratchv/ir/builder.py`,追加在 `reshape` 之后) + +| 方法 | 签名 | 发射 OpCode | dest dtype | +|------|------|-------------|-----------| +| `sqrt` | `sqrt(self, val: Value, dtype: DataType = DataType.FLOAT32) -> Value` | `SQRT` | 参数 `dtype` | +| `min` | `min(self, a: Value, b: Value, dtype: DataType = DataType.INT32) -> Value` | `MIN` | 参数 `dtype` | +| `max` | `max(self, a: Value, b: Value, dtype: DataType = DataType.INT32) -> Value` | `MAX` | 参数 `dtype` | +| `abs` | `abs(self, val: Value, dtype: DataType = DataType.INT32) -> Value` | `ABS` | 参数 `dtype` | +| `idiv` | `idiv(self, a: Value, b: Value) -> Value` | `IDIV` | `INT32` | +| `rem` | `rem(self, a: Value, b: Value) -> Value` | `REM` | `INT32` | +| `mod` | `mod(self, a: Value, b: Value) -> Value` | `MOD` | `INT32` | +| `load_f64` | `load_f64(self, addr: Value) -> Value` | `LOAD_F64` | `FLOAT64` | +| `store_f64` | `store_f64(self, addr: Value, val: Value) -> Instruction` | `STORE_F64`(`operands=[addr, val]`) | 无 | +| `load_const_f64` | `load_const_f64(self, value: float) -> Value` | `LOAD_CONST_F64`(`value=` attr) | `FLOAT64` | +| `fadd_d` | `fadd_d(self, a: Value, b: Value) -> Value` | `FADD_D` | `FLOAT64` | +| `fsub_d` | `fsub_d(self, a: Value, b: Value) -> Value` | `FSUB_D` | `FLOAT64` | +| `fmul_d` | `fmul_d(self, a: Value, b: Value) -> Value` | `FMUL_D` | `FLOAT64` | +| `fdiv_d` | `fdiv_d(self, a: Value, b: Value) -> Value` | `FDIV_D` | `FLOAT64` | +| `fcmp_l_d` | `fcmp_l_d(self, a: Value, b: Value) -> Value` | `FCMP_L_D` | `INT32` | +| `fcmp_eq_d` | `fcmp_eq_d(self, a: Value, b: Value) -> Value` | `FCMP_EQ_D` | `INT32` | +| `fcvt_s_d` | `fcvt_s_d(self, val: Value) -> Value` | `FCVT_S_D` | `FLOAT32` | +| `fcvt_d_s` | `fcvt_d_s(self, val: Value) -> Value` | `FCVT_D_S` | `FLOAT64` | + +实现骨架(其余同构): + +```python + def idiv(self, a: Value, b: Value) -> Value: + dest = self.make_value(dtype=DataType.INT32) + self._emit(OpCode.IDIV, dest, [a, b]) + return dest + + def store_f64(self, addr: Value, val: Value) -> Instruction: + return self._emit(OpCode.STORE_F64, operands=[addr, val]) + + def load_const_f64(self, value: float) -> Value: + dest = self.make_value(dtype=DataType.FLOAT64, + is_constant=True, const_value=value) + self._emit(OpCode.LOAD_CONST_F64, dest, value=value) + return dest +``` + +`min/max/abs/sqrt` 带 `dtype` 形参是刻意设计:`IRBuilder.make_value` 默认 `FLOAT32`,无显式 dtype 会把 i32 算子误标为 f32 而触发 §2.4 的 dtype 报错。 + +### 2.3 选择器契约(`scratchv/backend/inst_select_ext.py`) + +#### 2.3.1 OpCode ↔ handler 对应表(必须全部可达) + +| `.value` | handler | 关键产物 | +|----------|---------|----------| +| `sqrt` | `_select_sqrt` | MV/CALL 或 `SQRT_S`/`SQRT_D` | +| `min` | `_select_min` | `FMIN_D` 或 SLT/SUB/AND/ADD | +| `max` | `_select_max` | `FMAX_D` 或 `MachineOp.MAX` | +| `abs` | `_select_abs` | `FABS_D` 或 SRAI/XOR/SUB | +| `idiv` | `_select_idiv` | DIV | +| `rem` | `_select_rem` | REM | +| `mod` | `_select_mod` | REM(转发 `_select_rem`) | +| `load_f64` | `_select_load_f64` | FLD | +| `store_f64` | `_select_store_f64` | FSD | +| `load_const_f64` | `_select_load_const_f64` | LI_D | +| `fadd_d` | `_select_fadd_d` | FADD_D | +| `fsub_d` | `_select_fsub_d` | FSUB_D | +| `fmul_d` | `_select_fmul_d` | FMUL_D | +| `fdiv_d` | `_select_fdiv_d` | FDIV_D | +| `fcmp_l_d` | `_select_fcmp_l_d` | FLT_D | +| `fcmp_eq_d` | `_select_fcmp_eq_d` | FEQ_D | +| `fcvt_s_d` | `_select_fcvt_s_d` | FCVT_S_D | +| `fcvt_d_s` | `_select_fcvt_d_s` | FCVT_D_S | + +#### 2.3.2 新增/变更的类成员(精确名称) + +| 名称 | 签名/类型 | 说明 | +|------|-----------|------| +| `_FP64_OPCODES` | `frozenset[str]` 类常量 | 值集合:`{"load_f64","store_f64","load_const_f64","fadd_d","fsub_d","fmul_d","fdiv_d","fcmp_l_d","fcmp_eq_d","fcvt_s_d","fcvt_d_s"}` | +| `_fresh_temp` | `(self, prefix: str) -> MachineOperand` | 返回 `MachineOperand.vreg(f"__{prefix}_{n}")`,`n` 由 `_temp_counter` 单调递增 | +| `_involves_fp64` | `(self, instr: Instruction) -> bool` | **纯检测**(忽略 enable 开关):dest 或任一 operand 的 dtype 为 `FLOAT64` | +| `_require_fp64` | `(self, instr: Instruction) -> None` | 未启用时抛 `ValueError` | +| `_check_dtype` | `(self, instr, allowed: tuple[DataType, ...], opname: str) -> None` | dest dtype 越界时抛 `ValueError`(含 opcode 与 dtype) | +| `run` | `(self) -> list[MachineInstr]` | 覆写:`self._temp_counter = 0` 后调用 `super().run()`(保证文本确定性) | +| `supported_ops` | `property` | 清单与 18 个值一致;fp64 子集按现有语义受 `enable_fp64` 控制 | + +保留但改变语义的成员:`_is_fp64(instr)` ≡ `self.enable_fp64 and self._involves_fp64(instr)`(旧语义),仅作兼容保留;新代码一律用 `_involves_fp64` + `_require_fp64`。 + +#### 2.3.3 异常与消息模板(测试按子串匹配) + +| 场景 | 异常 | 消息模板 | +|------|------|----------| +| fp64 专用 opcode 关闭时 | `ValueError` | `opcode '{value}' requires enable_fp64=True (ExtendedInstructionSelector)` | +| dtype 驱动路径关闭时 | `ValueError` | `instruction '{value}' involves FLOAT64 but enable_fp64=False (ExtendedInstructionSelector)` | +| dtype 越界 | `ValueError` | `{opname} requires destination dtype in {allowed}, got {dtype}` | +| 常量缺失/非法 | `ValueError` | `load_const_f64 requires numeric attrs['value'], got {raw!r}` | + +所有消息必须含子串 `enable_fp64`(前两类)或 opcode 名(后两类)。 + +### 2.4 配置与 CLI 契约 + +#### 2.4.1 `CompilerConfig` 新增字段(`scratchv/compiler.py`) + +| 字段 | 类型 | 默认 | 语义 | +|------|------|------|------| +| `extended_isel` | `bool` | `False` | 线性 RISC-V 管线使用扩展选择器 | +| `enable_fp64` | `bool` | `True` | 传递给 `ExtendedInstructionSelector.enable_fp64` | +| `use_hardware_sqrt` | `bool` | `False` | 传递给 `ExtendedInstructionSelector.use_hardware_sqrt` | + +#### 2.4.2 CLI 参数(`scratchv/main.py`) + +| 参数 | argparse 定义 | 映射 | +|------|----------------|------| +| `--extended-isel` | `action="store_true"`,默认 False | `extended_isel=args.extended_isel` | +| `--no-fp64` | `dest="enable_fp64", action="store_false"`,默认 True | `enable_fp64=args.enable_fp64` | +| `--hardware-sqrt` | `dest="use_hardware_sqrt", action="store_true"`,默认 False | `use_hardware_sqrt=args.use_hardware_sqrt` | + +`--extended-isel` 精确行为矩阵(与设计文档 §2.3.2 一致): + +| 条件 | 行为 | warning 文本(追加到 `CompileResult.warnings`) | +|------|------|--------------------------------------------------| +| `extended_isel` 且 backend=riscv 且非 `use_dag_isel` | 使用 `ExtendedInstructionSelector(program, enable_fp64=..., use_hardware_sqrt=...)` | 无 | +| 否则 | 基础 `InstructionSelector` | — | +| `extended_isel` + `use_dag_isel` | DAG 路径优先,选择器不替换 | `--dag-isel takes precedence; --extended-isel ignored` | +| `extended_isel` + backend=llvm | LLVM 路径优先 | `--extended-isel is RISC-V only; ignored for LLVM backend` | +| 非 `extended_isel` 且(`enable_fp64=False` 或 `use_hardware_sqrt=True`) | 不生效 | `--no-fp64/--hardware-sqrt have no effect without --extended-isel` | + +### 2.5 编码器契约(`scratchv/backend/riscv_encoder.py`) + +| 名称 | 签名/类型 | 说明 | +|------|-----------|------| +| `UnsupportedInstructionError` | 类,继承 `ValueError` | 本课题新增;可由 `from scratchv.backend.riscv_encoder import UnsupportedInstructionError` 导入 | +| `_is_fd_mnemonic` | `(op: str) -> bool` | 精确集 `{"fld","fsd","flw","fsw","li.d"}` + 前缀集 `("fadd.","fsub.","fmul.","fdiv.","fsqrt.","fmin.","fmax.","fabs.","fneg.","flt.","fle.","feq.","fcvt.","fmv.","fsgnj","fsgnjn","fsgnjx")` | +| 兜底报错 | `_encode_line` 末分支 | F/D → `UnsupportedInstructionError`;其他 → 保持 `ValueError(f"Unknown instruction: {op}")` | + +消息模板:`F/D instruction '{op}' is not supported by the RV32IM encoder (Topic 28: final encoding out of scope; output is assembly text)`。 + +### 2.6 发射器契约(`scratchv/backend/asm_emit.py`) + +`AsmEmitter._format_instr` 未知 MachineOp:`raise ValueError(f"no assembly mapping for MachineOp {instr.op.name} ({instr.op.value})")`。`_OP_NAMES` 不需要新增条目(全部在产 F/D MachineOp 已映射)。 + +### 2.7 MachineOp 契约(`scratchv/backend/machine_types.py`) + +**本课题新增 0 个成员**。实现时必须逐一核对以下已存在成员(缺一即阻塞): + +`SQRT_S`(fsqrt.s)、`SQRT_D`(fsqrt.d)、`FMIN_D`、`FMAX_D`、`FABS_D`、`FNEG_D`、`FADD_D`、`FSUB_D`、`FMUL_D`、`FDIV_D`、`FLT_D`、`FEQ_D`、`FCVT_S_D`、`FCVT_D_S`、`LI_D`、`FLD`、`FSD`、`SRAI`、`XOR`、`AND`、`SLT`、`REM`、`MAX`、`DIV`。 + +约定:`MachineOp` 同样 append-only;`SQRT_S`/`FSQRT_S` 为同值别名,保留不动。 + +--- + +## 三、分模块改动清单 + +### 3.1 `scratchv/ir/types.py` + +按 §2.1 精确块插入。插入后自检: + +```bash +python -c "from scratchv.ir.types import OpCode; print([o.value for o in OpCode][-18:])" +# 期望输出以 'sqrt','min','max','abs','idiv','rem','mod', +# 'load_f64','store_f64','load_const_f64','fadd_d','fsub_d','fmul_d', +# 'fdiv_d','fcmp_l_d','fcmp_eq_d','fcvt_s_d','fcvt_d_s' 结尾 +``` + +### 3.2 `scratchv/ir/builder.py` + +按 §2.2 表追加 18 个方法(置于文件末尾,`reshape` 之后)。无既有方法签名变更。 + +### 3.3 `scratchv/backend/instruction_select.py` + +- **功能零改动**。dispatch 仍为 `getattr(self, f"_select_{instr.opcode.value}", None)` + 未命中 `ValueError`。 +- 建议在 `_select_instruction` 上方补一行契约注释(`# handler 命名契约:_select_{opcode.value}`),并在 §4 的不变式测试中断言:对 `OpCode` 中属于 `[Topic 28]` 分区的成员,扩展选择器均有同名 handler。 +- 基础选择器对新 OpCode 的 `ValueError` 行为是预期结果,不修。 + +### 3.4 `scratchv/backend/inst_select_ext.py`(P1–P7) + +`__init__` 增加 `self._temp_counter = 0`;覆写 `run()` 重置计数器。 + +**P1 fp64 门控**(替换现有 `_select_instruction`,`inst_select_ext.py:62-66`): + +```python + def _select_instruction(self, instr: Instruction) -> None: + if not self.enable_fp64 and instr.opcode.value in self._FP64_OPCODES: + raise ValueError( + f"opcode '{instr.opcode.value}' requires enable_fp64=True " + f"(ExtendedInstructionSelector)") + if instr.dest is not None: + self._current_dtype = instr.dest.dtype + super()._select_instruction(instr) +``` + +类型驱动覆写(`_select_add/_select_sub/_select_mul/_select_div/_select_neg/_select_load/_select_store/_select_load_const`)统一改为: + +```python + def _select_add(self, instr: Instruction) -> None: + if self._involves_fp64(instr): + self._require_fp64(instr) + self._select_fadd_d(instr) + else: + super()._select_add(instr) +``` + +`_select_min/_select_max/_select_abs/_select_sqrt` 的条件同样由 `self.enable_fp64 and dest.dtype == FLOAT64` 改为 `_involves_fp64` + `_require_fp64`。 + +**P2 唯一临时值**: + +```python + def _fresh_temp(self, prefix: str) -> MachineOperand: + self._temp_counter += 1 + return MachineOperand.vreg(f"__{prefix}_{self._temp_counter}") +``` + +`_select_min` 的三个临时值 `__min_slt_{n}/__min_sub_{n}/__min_and_{n}`,`_select_abs` 的两个临时值 `__abs_srai_{n}/__abs_xor_{n}`。 + +**P3 `store_f64` 顺序**: + +```python + def _select_store_f64(self, instr: Instruction) -> None: + addr = self._op(instr, 0) + val = self._op(instr, 1) + self._emit(MachineOp.FSD, val, addr, comment="fsd (store f64)") +``` + +**P4 `sqrt` 立即数**:软件路径中 `src.kind == "imm"` 时发射 `LI a0, src`,否则 `MV a0, src`;随后 `CALL`,再 `MV dst, a0`(f64 时函数名 `sqrt`,f32 为 `sqrtf`)。 + +**P5 常量位模式**: + +```python + def _select_load_const_f64(self, instr: Instruction) -> None: + raw_val = instr.attrs.get("value") + if not isinstance(raw_val, (int, float)): + raise ValueError( + f"load_const_f64 requires numeric attrs['value'], got {raw_val!r}") + bits = struct.unpack("` | +| `TestAsmText::test_sqrt_hardware_f32_f64` | `fsqrt.s d, x`;`fsqrt.d d, x` | +| `TestAsmText::test_fp64_dtype_driven_add` | FLOAT64 的 `ADD` 输出 `fadd.d` | +| `TestAsmText::test_store_f64_operand_order` | `store_f64(p, v)` 输出 `fsd v, p` | +| `TestFp64Gate::test_fadd_d_without_fp64_raises` | `pytest.raises(ValueError, match="enable_fp64")` | +| `TestFp64Gate::test_f64_add_without_fp64_raises` | 同上 | +| `TestFp64Gate::test_f64_min_without_fp64_raises` | 同上 | +| `TestFp64Gate::test_integer_ops_without_fp64_still_work` | i32 `ADD` 输出 `add` | +| `TestIllegalDtype::test_sqrt_int_raises` / `test_min_f32_raises` / `test_idiv_f64_raises` | `ValueError` 且消息含 opcode 名 | +| `TestLoadConstF64::test_exact_bits` | 1.5→`4609434218613702656`;2.0→`4611686018427387904`;−0.0→`9223372036854775808` | +| `TestLoadConstF64::test_missing_value_raises` | 手工构造 `Instruction(OpCode.LOAD_CONST_F64, dest, attrs={})` → `ValueError` | +| `TestSupportedOps::test_new_ops_present` | `supported_ops` 含 18 个值中 f32 子集;`enable_fp64=True` 时含 fp64 子集 | + +### 4.2 `tests/test_riscv_encoder_fd.py`(新增) + +| 用例 | 断言要点 | +|------|----------| +| `test_fd_mnemonics_raise`(parametrize:`fadd.d`/`fsub.s`/`fmul.d`/`fdiv.d`/`fsqrt.d`/`fmin.d`/`fmax.d`/`fabs.d`/`fneg.d`/`fld`/`fsd`/`flw`/`fsw`/`flt.d`/`feq.d`/`fcvt.s.d`/`fcvt.d.s`/`fmv.x.w`/`li.d`) | `pytest.raises(UnsupportedInstructionError)`,消息含助记符 | +| `test_unsupported_is_value_error` | `issubclass(UnsupportedInstructionError, ValueError)` | +| `test_unknown_non_fd_still_value_error` | `"frobnicate x1, x2"` → `ValueError` 且消息 `Unknown instruction` | +| `test_is_fd_mnemonic_predicate` | `_is_fd_mnemonic("add") is False`;`"fence"`/`"flw"` 等边界 | +| `test_rv32im_regression_still_assembles` | 既有 `li`/`add`/`sw`/`beq` 小序列编码成功(与 `test_simulator.py` 用例不重复即可) | + +### 4.3 `tests/test_extended_isel_cli.py`(新增) + +| 用例 | 断言要点 | +|------|----------| +| `test_cli_flags_map_to_config` | `build_arg_parser().parse_args(["-o","x","--dsl","a","--extended-isel","--no-fp64","--hardware-sqrt"])` → `args_to_config` 三字段正确 | +| `test_driver_uses_extended_selector` | `unittest.mock.patch("scratchv.backend.inst_select_ext.ExtendedInstructionSelector")` 后 `CompilerDriver(CompilerConfig(extended_isel=True)).compile(...)`,断言以 `enable_fp64=True, use_hardware_sqrt=False` 调用 | +| `test_default_uses_base_selector` | 未开旗标时 `InstructionSelector` 被调用(patch 基础类) | +| `test_dag_isel_conflict_warns` | `extended_isel=True, use_dag_isel=True` → `result.warnings` 含 `precedence` | +| `test_llvm_backend_warns` | `extended_isel=True, backend="llvm"` → warning 含 `RISC-V only` | +| `test_fp64_flags_without_extended_warn` | `enable_fp64=False, extended_isel=False` → warning 含 `no effect` | + +(编译输入用 `tmp_path` 写最小 DSL 文件,避免依赖 ONNX。) + +### 4.4 运行命令 + +```bash +python3 -m pytest tests/test_inst_select_ext.py tests/test_riscv_encoder_fd.py tests/test_extended_isel_cli.py -v +make test # 全量回归 +make lint # flake8 + mypy +``` + +若本地 `python3` 为 3.8(venv 现状),请改用支持 3.9+ 语法的解释器:`/usr/local/bin/python3.11 -m pytest ...`。 + +--- + +## 五、实施顺序与验收标准 + +### 5.1 推荐提交序列(每步可独立测试) + +1. `ir/types.py` + `ir/builder.py`(纯新增,不影响既有路径)。 +2. `inst_select_ext.py` P1–P7 + `tests/test_inst_select_ext.py`。 +3. `riscv_encoder.py` + `asm_emit.py` + `tests/test_riscv_encoder_fd.py`。 +4. `compiler.py` + `main.py` + `tests/test_extended_isel_cli.py`。 +5. 文档同步:修正 `docs/topics/28-扩展指令选择.md` 的状态与算子表(可选,但建议随实现提交)。 + +### 5.2 验收检查表 + +| # | 验收项 | 通过标准 | +|---|--------|----------| +| A1 | dispatch 全覆盖 | 18 个 `.value` 均有 handler;parametrize 用例全绿 | +| A2 | 汇编文本 | MIN/ABS/SQRT/f64 用例逐行匹配设计文档 §3 与附录 | +| A3 | fp64 关闭 | 四类 f64 输入均抛 `ValueError`(含 `enable_fp64`),无整数兜底 | +| A4 | 常量精度 | `1.5` 不再被截断为 `1`;三个位模式精确匹配 | +| A5 | 编码器 fail-loud | 采样 F/D 助记符全部抛 `UnsupportedInstructionError`;未知非 F/D 仍 `Unknown instruction`;子类关系成立 | +| A6 | CLI 接线 | 旗标→config→选择器链路有 mock 断言;冲突组合产生 warning;默认路径不变 | +| A7 | 全量回归 | `make test` 0 失败;默认管线输出与改动前逐字节一致(抽样对比) | +| A8 | 静态检查 | `make lint` 对改动文件无新增告警(既有白名单除外) | +| A9 | 变更边界 | 未新增依赖;未改 `MachineOp` 成员;未改前端;未改 `instruction_select.py` 逻辑 | + +--- + +## 六、风险与回退 + +| 编号 | 风险 | 触发条件 | 缓解/回退 | +|------|------|----------|-----------| +| R1 | 与课题 29 并行合入冲突 | 两者都追加 `OpCode` | append-only 分区 + 注释约定;冲突时双分区共存,禁止重排;合入后跑全量 `make test` | +| R2 | f64 路径不可执行被误用 | 用户在默认 `--extended-isel` 下期望运行 f64 | 文档与 help 明确「汇编文本级」;编码器 fail-loud 使错误尽早暴露;后续课题补 f-reg 分配/常量池 | +| R3 | fail-loud 改变既有行为 | 外部调用者依赖「未知 MachineOp 输出注释」或「f64 静默走整数」 | 两类行为均属静默错误,按本课题修正;如需兼容,可在调用侧显式捕获 `ValueError` 并处理 | +| R4 | `AsmEmitter` 抛错影响 `GLOBL/SIZE/TYPE` | 外部代码手工构造这些 MachineOp | 当前库内无生产者;如出现,优先为其补 `_OP_NAMES` 映射而非恢复静默 | +| R5 | dtype 守卫误伤 | `make_value()` 默认 FLOAT32,调用方未传 dtype | builder 方法默认 `INT32`(min/max/abs/idiv/rem/mod);报错信息含 dtype,便于定位;必要时在 docs 中给用法示例 | +| R6 | 回归难以定位 | 大批量文本断言 | 断言只针对去注释后的关键行;用 `parametrize` 拆分,失败可精确到 opcode | + +**回退策略**:`extended_isel` 默认 False,默认编译路径零影响;如需整体回退,revert 课题提交即可(新增 OpCode 在不被引用时为惰性成员,不影响既有枚举语义)。`UnsupportedInstructionError` 继承 `ValueError`,即使外部代码只捕获 `ValueError` 也保持兼容。 + +--- + +## 七、接口契约速查(一页) + +``` +IR 枚举(18): SQRT=sqrt, MIN=min, MAX=max, ABS=abs, IDIV=idiv, REM=rem, MOD=mod, + LOAD_F64=load_f64, STORE_F64=store_f64, LOAD_CONST_F64=load_const_f64, + FADD_D=fadd_d, FSUB_D=fsub_d, FMUL_D=fmul_d, FDIV_D=fdiv_d, + FCMP_L_D=fcmp_l_d, FCMP_EQ_D=fcmp_eq_d, FCVT_S_D=fcvt_s_d, FCVT_D_S=fcvt_d_s + +Builder: sqrt(val,dtype=F32), min/max(a,b,dtype=I32), abs(val,dtype=I32), + idiv/rem/mod(a,b), load_f64(addr), store_f64(addr,val), load_const_f64(value), + fadd_d/fsub_d/fmul_d/fdiv_d(a,b), fcmp_l_d/fcmp_eq_d(a,b), + fcvt_s_d(val), fcvt_d_s(val) + +Selector: ExtendedInstructionSelector(program, enable_fp64=True, use_hardware_sqrt=False) + 内部: _FP64_OPCODES, _fresh_temp(prefix), _involves_fp64(instr), + _require_fp64(instr), _check_dtype(instr, allowed, opname) + +Config: extended_isel=False, enable_fp64=True, use_hardware_sqrt=False +CLI: --extended-isel --no-fp64 --hardware-sqrt +异常: UnsupportedInstructionError(ValueError) # riscv_encoder + _is_fd_mnemonic(op) -> bool +``` + +--- + +## 实现结果(2026-09-14 集成) + +> **集成 commit**:`d5af621`(`feat(topic28): activate extended instruction selection with fp64 support`) +> **集成位置**:`Seven_big_summary` 上第 10 个 topic commit(顺序 … → 15 → **28** → 29) +> **集成后全量**:`PYTHONPATH=. python3.11 -m pytest tests/ -q` → **1011 passed / 13 xfailed / 20 xpassed / 0 failed** + +### 实现文件与要点 + +| 文件 | 要点 | +|------|------| +| `scratchv/ir/types.py` | `[Topic 28]` 分区 18 个 OpCode(SQRT / MIN / MAX / ABS / IDIV / REM / MOD / LOAD_F64 / STORE_F64 / LOAD_CONST_F64 / FADD_D / FSUB_D / FMUL_D / FDIV_D / FCMP_L_D / FCMP_EQ_D / FCVT_S_D / FCVT_D_S) | +| `scratchv/ir/builder.py` | 对应 builder API | +| `scratchv/backend/inst_select_ext.py` | `ExtendedInstructionSelector` 激活:fp64 门控、dtype 守卫、常量位模式物化 | +| `scratchv/backend/instruction_select.py` | 接入扩展选择器 | +| `scratchv/backend/riscv_encoder.py` | F/D 指令 `UnsupportedInstructionError`(fail-loud) | +| `scratchv/backend/asm_emit.py` | fail-loud(未知 opcode 不再静默) | +| `scratchv/compiler.py`、`scratchv/main.py` | `--extended-isel` / `--no-fp64` / `--hardware-sqrt` 双侧接线 | +| `tests/test_inst_select_ext.py`、`tests/test_riscv_encoder_fd.py`、`tests/test_extended_isel_cli.py` | 共 99 个定向用例 | + +### 测试数字 + +| 口径 | 结果 | +|------|------| +| 定向(3 个测试文件) | 99 passed | +| 分支全量(cherry-pick 前) | 651 passed | +| 集成后全量 | 1011 passed / 13 xfailed / 20 xpassed / 0 failed | + +### 与本文档的偏差 / 未完成项 + +- `_check_dtype` 同时校验操作数 dtype(文档只写了 dest)。 +- `_select_store_f64` 额外校验操作数个数与 value 的 dtype。 + +### 已知限制 + +- f64 路径为**汇编文本级**,不可执行(无 f-reg 分配 / 常量池,§六 R2)。 +- `--extended-isel + --minimal-call-codegen` 组合未覆盖(低风险)。 +- 编码器/发射器对未知指令 fail-loud,属有意行为变更(R3)。 diff --git "a/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\350\256\276\350\256\241\346\226\207\346\241\243.md" "b/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\350\256\276\350\256\241\346\226\207\346\241\243.md" new file mode 100644 index 0000000..b51bba7 --- /dev/null +++ "b/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\350\256\276\350\256\241\346\226\207\346\241\243.md" @@ -0,0 +1,527 @@ +# 课题28 扩展指令选择技术设计文档 + +> 文档版本:v1.0 +> 编写日期:2026-09-14 +> 涉及模块:`scratchv/ir/types.py`(OpCode)、`scratchv/ir/builder.py`(IRBuilder)、`scratchv/backend/instruction_select.py`(dispatch 契约)、`scratchv/backend/inst_select_ext.py`(主改动)、`scratchv/backend/machine_types.py`、`scratchv/backend/asm_emit.py`、`scratchv/backend/riscv_encoder.py`、`scratchv/compiler.py`、`scratchv/main.py` +> 功能范围:18 个新增 IR OpCode(7 个整数/通用 + 11 个 float64)的指令选择;RV32IM 软件序列与 F/D 硬件指令的选择策略;`--extended-isel` 的 CLI→config→selector 全链路接线;F/D 编码的 fail-loud 行为 +> 状态说明:本文档描述**目标实现**。当前代码存在 §1.1 所列已实测断链(2026-09-14),设计阶段未修改仓库任何文件;实现指南见同目录《开发文档.md》,二者冲突时以本文档为准。 + +--- + +## 一、功能介绍 + +### 1.1 功能概述 + +扩展指令选择器 `ExtendedInstructionSelector`(`scratchv/backend/inst_select_ext.py`)已写出 sqrt/min/max/abs/idiv/rem/mod 与 float64(D 扩展)的指令选择代码,但实测存在三处断链,导致其中大部分分支不可达。`docs/topics/28-扩展指令选择.md` 标注的状态「✅ 已完成」与实测不符。 + +#### 断链 1:IR 层缺 OpCode,dispatch 永远命不中 + +`InstructionSelector._select_instruction`(`instruction_select.py:47`)按命名约定分发: + +```python +handler = getattr(self, f"_select_{instr.opcode.value}", None) +if handler is None: + raise ValueError(f"No instruction selection for opcode: {instr.opcode}") +``` + +而 `ir/types.py` 的 `OpCode` 枚举(当前 28 个成员)没有 `sqrt`/`min`/`max`/`abs`/`idiv`/`rem`/`mod`/`load_f64`/… 等成员(实测 `'sqrt' in [o.value for o in OpCode] == False`)。因此 `_select_sqrt`、`_select_min`、`_select_max`、`_select_abs`、`_select_idiv`、`_select_rem`、`_select_mod` 及全部 fp64 handler 都不可达。当前唯一实际生效的是 `_select_div` 内部对 `DataType.INT32` 的分支(`inst_select_ext.py:312-318`): + +```python +def _select_div(self, instr): + if self._is_fp64(instr): # 永不成立(OpCode 层可写,dtype 层可达) + ... + elif instr.dest and instr.dest.dtype == DataType.INT32: + self._select_idiv(instr) # 实际生效 + else: + super()._select_div(instr) +``` + +#### 断链 2:`--extended-isel` 被解析但未接线 + +`main.py:104` 定义了 `--extended-isel`,但 `args_to_config`(`main.py:135-156`)没有把它写入 `CompilerConfig`(`compiler.py:34-75` 无对应字段),`CompilerDriver._generate_riscv_linear`(`compiler.py:397`)固定使用基础 `InstructionSelector`。该旗标当前是 no-op。 + +#### 断链 3:编码器不支持 F/D + +`riscv_encoder.py` 是 RV32IM 编码器,对 F/D 指令(含 `li.d` 伪指令)一律抛 `ValueError: Unknown instruction: ...`(实测:`fadd.d`/`fld`/`fsqrt.d`/`li.d` 均如此)。错误信息未区分「本课题范围外的不支持」与「拼写错误」。 + +#### 本课题交付的能力 + +1. **整数/通用扩展**:`sqrt`、`min`、`max`、`abs`、`idiv`、`rem`、`mod` 七个 OpCode 可达;策略为 RV32IM 软件序列(平方根走 libm 调用)或 M 扩展原生指令。 +2. **float64(D 扩展)**:11 个显式 f64 OpCode,外加既有 `ADD/SUB/MUL/DIV/NEG/LOAD/STORE/LOAD_CONST` 的 f64 类型驱动路由;`enable_fp64=True` 时输出 D 扩展汇编文本。 +3. **CLI 接线**:`--extended-isel`(配套 `--no-fp64`、`--hardware-sqrt`)→ `CompilerConfig` → `ExtendedInstructionSelector`。 +4. **fail-loud**:fp64 关闭时凡涉及 f64 的选择必须报错;编码器对 F/D 明确报错,绝不静默产出错误机器码;发射器对未知 MachineOp 报错。 +5. **并行安全**:`OpCode` 扩展采用末尾追加 + 注释分区,与课题 29(RVV)的枚举扩展可并行合入。 + +#### 非目标(范围边界) + +- **不实现完整 F/D 编码器**:`riscv_encoder.py` 对 F/D 明确不支持并报错;本课题只承诺汇编文本级语义。 +- **不实现 RVV**(课题 29)。 +- **不实现 f-register 寄存器分配与可执行 f64 数值验证**;正确性以汇编文本与软件序列语义为准(不依赖模拟器)。 +- **不改动 DSL/ONNX 前端**对上述 OpCode 的生成(`dsl_parser.py`/`dsl_extended.py`/`onnx_parser.py` 均无对应算子映射,属后续课题)。 +- **不做 `MOD` 的 floor 语义**(见 §2.1.3),不改变 `REM/IDIV` 的除零行为。 + +### 1.2 设计目标 + +- **可达性**:每个新 OpCode 的 `.value` 与 `_select_{value}` handler 严格一一对应,并有不变式测试防回归。 +- **策略显式**:同一 OpCode 的整数/浮点分支、软件/硬件分支由 `dtype` 与显式开关决定,不依赖隐式默认值。 +- **禁止静默**:不可选、不可编码、未定义的类型组合必须抛错——禁止整数序列兜底 f64、禁止 `int()` 截断浮点常量、禁止发射器用注释吞掉未知 MachineOp。 +- **零回归**:`extended_isel=False` 为默认,基础管线输出与改动前逐字节一致。 +- **并行安全**:`OpCode` 枚举 append-only + 注释分区,与课题 29 的扩展不冲突。 +- **可验证**:全部新行为可用纯汇编文本断言 + pytest 覆盖,不依赖外部模拟器或编码器。 + +--- + +## 二、设计规范 + +### 2.1 新增 IR OpCode 清单与语义 + +#### 2.1.1 IR 记法(BNF 等价定义) + +```bnf +so_instr ::= opcode_name operands [ "->" dest ] [ "[" attrs "]" ] +opcode_name ::= name ; 见 2.1.3 定名表,Dispatch 键为 opcode.value +operands ::= value { "," value } +value ::= "%" name ":" dtype +dtype ::= "i32" | "i64" | "f32" | "f64" +attrs ::= "value" "=" number ; 仅 LOAD_CONST_F64 使用 +``` + +ScratchV 中的等价 Python 表示为 `Instruction(opcode=..., dest=Value, operands=[Value, ...], attrs={...})`;`Instruction.opcode.value` 同时是 dispatch 键。 + +#### 2.1.2 元素说明 + +| 元素 | 说明 | +|------|------| +| `opcode_name` | `OpCode` 枚举成员的 `.value`,决定 handler 名 `_select_{value}` | +| `dest` | 目标 `Value`(`STORE_F64` 除外,必备);其 `dtype` 决定选择策略 | +| `operands` | 输入 `Value` 列表,顺序固定(见 §2.5) | +| `dtype` | `DataType`(i32/i64/f32/f64);f64 相关一律受 `enable_fp64` 约束 | +| `attrs["value"]` | 仅 `LOAD_CONST_F64`:Python `float` 常量(不得被 `int()` 截断) | + +#### 2.1.3 最终定名与语义表(18 个) + +| # | 枚举名 | `.value` | 元数 | 合法 dtype(dest / operands) | 语义 | +|---|--------|----------|------|-------------------------------|------| +| 1 | `OpCode.SQRT` | `sqrt` | 1 | f32 或 f64 | `dst = √src`(IEEE-754 舍入) | +| 2 | `OpCode.MIN` | `min` | 2 | i32/i64(整数)或 f64 | `dst = min(a, b)`;整数按有符号比较 | +| 3 | `OpCode.MAX` | `max` | 2 | 同上 | `dst = max(a, b)` | +| 4 | `OpCode.ABS` | `abs` | 1 | i32/i64 或 f64 | `dst = |src|`(整数用位运算序列) | +| 5 | `OpCode.IDIV` | `idiv` | 2 | i32 | `dst = a / b`,向零截断(RISC-V `div`) | +| 6 | `OpCode.REM` | `rem` | 2 | i32 | `dst = a − trunc(a/b)·b`,符号随被除数(RISC-V `rem`) | +| 7 | `OpCode.MOD` | `mod` | 2 | i32 | 本课题定义为 REM 的同义操作(lowering 相同);仅当 a,b ≥ 0 时等于数学取模;floor 语义不在范围 | +| 8 | `OpCode.LOAD_F64` | `load_f64` | 1 | src: i32/i64 地址;dst: f64 | `dst = *(f64*)addr` | +| 9 | `OpCode.STORE_F64` | `store_f64` | 2 | operands=[addr, value];无 dest | `*(f64*)addr = value` | +| 10 | `OpCode.LOAD_CONST_F64` | `load_const_f64` | 0 | dst: f64;`attrs["value"]`: float | 将 64 位 IEEE-754 位模式物化到 f64 值 | +| 11 | `OpCode.FADD_D` | `fadd_d` | 2 | f64 | `dst = a + b`(双精度) | +| 12 | `OpCode.FSUB_D` | `fsub_d` | 2 | f64 | `dst = a − b` | +| 13 | `OpCode.FMUL_D` | `fmul_d` | 2 | f64 | `dst = a × b` | +| 14 | `OpCode.FDIV_D` | `fdiv_d` | 2 | f64 | `dst = a ÷ b` | +| 15 | `OpCode.FCMP_L_D` | `fcmp_l_d` | 2 | f64 输入;i32 结果 | `dst = (a < b) ? 1 : 0`(无序时 0,`flt.d`) | +| 16 | `OpCode.FCMP_EQ_D` | `fcmp_eq_d` | 2 | 同上 | `dst = (a == b) ? 1 : 0`(`feq.d`) | +| 17 | `OpCode.FCVT_S_D` | `fcvt_s_d` | 1 | src: f64;dst: f32 | f64→f32 转换(`fcvt.s.d`) | +| 18 | `OpCode.FCVT_D_S` | `fcvt_d_s` | 1 | src: f32;dst: f64 | f32→f64 转换(`fcvt.d.s`) | + +定名约定:枚举名 = `.value` 全大写,且 `.value` 与现有 handler 后缀严格一致(例:`OpCode.FCMP_L_D.value == "fcmp_l_d"` ↔ `_select_fcmp_l_d`)。`IDIV` 为补充定名(用户清单未列,但 `_select_idiv` 已存在;补上使该 handler 可直接 dispatch,同时保留 `DIV`+INT32 的既有路径)。 + +#### 2.1.4 兼容性与分区约定 + +- 新成员**只追加**在 `OpCode` 枚举末尾(当前 `CONCAT = "concat"` 之后),并加 `[Topic 28]` 注释分区;禁止插入既有成员之间或重排(保护序列化与并行合入)。 +- 课题 29 及后续课题只能追加在 `[Topic 28]` 分区**之后**。 +- 不修改 `is_arith()`/`is_nn()`/`is_control_flow()`(当前无调用方),新 OpCode 不进入这些谓词,不影响 optimizer 分类。 + +### 2.2 指令选择策略表 + +#### 2.2.1 策略总表(RV32IM 软件序列 vs F/D 硬件指令) + +| OpCode | dtype | RV32IM 软件序列(默认) | F/D 硬件指令(开关条件) | 发射的 MachineOp | +|--------|-------|--------------------------|---------------------------|------------------| +| `SQRT` | f32 | `mv a0, src` → `call sqrtf` → `mv dst, a0` | `use_hardware_sqrt=True`:`fsqrt.s` | MV / CALL / SQRT_S | +| `SQRT` | f64 | `call sqrt`(libm double) | `use_hardware_sqrt=True` 且 `enable_fp64=True`:`fsqrt.d` | CALL / SQRT_D | +| `MIN` | i32/i64 | 分支无:`slt; sub; and; add` | — | SLT / SUB / AND / ADD | +| `MIN` | f64 | — | `fmin.d`(需 `enable_fp64=True`) | FMIN_D | +| `MAX` | i32/i64 | `max` 伪指令(编码器展开为 bge 分支序列) | — | MAX | +| `MAX` | f64 | — | `fmax.d`(需 `enable_fp64=True`) | FMAX_D | +| `ABS` | i32/i64 | 分支无:`srai 31; xor; sub` | — | SRAI / XOR / SUB | +| `ABS` | f64 | — | `fabs.d`(等价 `fsgnjx.d`;需 `enable_fp64=True`) | FABS_D | +| `IDIV` | i32 | — | M 扩展原生:`div` | DIV | +| `REM` | i32 | — | M 扩展原生:`rem` | REM | +| `MOD` | i32 | — | 同 REM(同义) | REM | +| `LOAD_F64` | f64 | — | `fld` | FLD | +| `STORE_F64` | f64 | — | `fsd`(RISC-V 序:`fsd value, 0(addr)`) | FSD | +| `LOAD_CONST_F64` | f64 | ScratchV 伪指令 `li.d dst, `(非 GAS 指令) | — | LI_D | +| `FADD_D` | f64 | — | `fadd.d` | FADD_D | +| `FSUB_D` | f64 | — | `fsub.d` | FSUB_D | +| `FMUL_D` | f64 | — | `fmul.d` | FMUL_D | +| `FDIV_D` | f64 | — | `fdiv.d` | FDIV_D | +| `FCMP_L_D` | f64→i32 | — | `flt.d` | FLT_D | +| `FCMP_EQ_D` | f64→i32 | — | `feq.d` | FEQ_D | +| `FCVT_S_D` | f64→f32 | — | `fcvt.s.d` | FCVT_S_D | +| `FCVT_D_S` | f32→f64 | — | `fcvt.d.s` | FCVT_D_S | + +#### 2.2.2 类型驱动路由(与显式 OpCode 等价的第二条路径) + +对 `ADD/SUB/MUL/DIV/NEG/LOAD/STORE/LOAD_CONST`:dest 或任一 operand 的 dtype 为 `FLOAT64` 时,`enable_fp64=True` 下分别发射 `FADD_D/FSUB_D/FMUL_D/FDIV_D/FNEG_D/FLD/FSD/LI_D`;`enable_fp64=False` 时报错(§2.4),**不执行整数兜底**。该路径与显式 f64 OpCode 产物一致,便于 ONNX f64 模型直接经类型推断进入。 + +#### 2.2.3 关键序列正确性 + +- **整数 MIN(分支无)**:`slt t,a,b; sub d,b,a; and t,t,d; add dst,a,t`。a 上表为当前仓库实测路径;若后续重构移动文件,以「同职责文件」对应,实际路径可能不同。 + +### 4.2 `ir/types.py`:追加 OpCode 分区 + +- 位置:`CONCAT = "concat"`(当前第 50 行)之后、`def is_arith(self)`(当前第 52 行)之前。 +- 内容:`[Topic 28]` 注释分区 + 18 个成员(顺序见 §2.1.3),整数/通用在前、f64 在后。 +- 禁止:重排/插入既有 28 个成员;改动 `DataType`。 + +### 4.3 `ir/builder.py`:新增 18 个构造方法 + +在文件末尾(当前 `reshape` 之后)追加,命名与签名见《开发文档.md》§2.2。约定: + +- `sqrt(val, dtype=FLOAT32)`、`min/max(a, b, dtype=INT32)`、`abs(val, dtype=INT32)` 显式传 dtype 以消除默认 FLOAT32 的歧义; +- `idiv/rem/mod` 固定 INT32;`fcmp_l_d/fcmp_eq_d` 结果 INT32; +- `load_f64(addr)`、`store_f64(addr, val)`(顺序与 `store(ptr, val)` 一致)、`load_const_f64(value)`; +- 显式 f64 运算 `fadd_d/fsub_d/fmul_d/fdiv_d`、转换 `fcvt_s_d/fcvt_d_s` 各自固定对应 dtype。 + +### 4.4 `instruction_select.py`:dispatch 接线方式 + +dispatch 无注册表,靠命名约定:`_select_{opcode.value}`(`instruction_select.py:48`)。因此「接线」= 补齐 OpCode 成员并保证 handler 名匹配,**无需修改基础选择器逻辑**。实现时仅: + +1. 在 `_select_instruction` 上加注释,声明该命名契约与「新 opcode 只在扩展选择器实现」; +2. 在测试中加不变式(§3.1)防回归; +3. 基础选择器对新 opcode 的 `ValueError` 保持现状。 + +### 4.5 `inst_select_ext.py`:修正点(P1–P7) + +| 编号 | 修正点 | 说明 | +|------|--------|------| +| P1 | fp64 门控 | `_select_instruction` 中对 f64 专用 opcode 统一检查 `enable_fp64`;新增 `_involves_fp64()`(纯检测)与 `_require_fp64()`(报错);`min/max/abs/sqrt` 由「`enable_fp64 and dtype==FLOAT64`」改为「`_involves_fp64` → 必须先通过门控」,消除 fp64 关闭时落入整数序列的静默路径 | +| P2 | 唯一临时值 | 新增 `_fresh_temp(prefix)`(计数器 `__{prefix}_{n}`),替换 `_min_tmp1/_abs_tmp1` 等固定名;`run()` 时重置计数器,保证文本确定性 | +| P3 | `store_f64` 顺序 | 与 `IRBuilder.store(ptr, val)` 对齐:`addr = _op(0)`、`val = _op(1)`,发射 `FSD(val, addr)`(当前实现读反) | +| P4 | `sqrt` 立即数 | 软件调用路径中,源为 immediate 时先 `LI a0, imm`,再 `CALL`(当前直接 `CALL`,a0 未定义) | +| P5 | 常量位模式 | `_select_load_const_f64` 用 `struct.unpack(" 扩展选择器对 `ADD/SUB/MUL/DIV/NEG/LOAD/STORE/LOAD_CONST` 的类型驱动分支做同样的门控与顺序修正(`STORE` → `_select_store_f64` 同样使用 P3 修正后的顺序)。 + +### 4.6 `machine_types.py`:盘点确认(本课题无新增) + +逐项核对 `MachineOp`(`machine_types.py:27-95`)——所需成员全部已存在:`SQRT_S/SQRT_D/FMIN_D/FMAX_D/FABS_D/FNEG_D/FADD_D/FSUB_D/FMUL_D/FDIV_D/FLT_D/FEQ_D/FCVT_S_D/FCVT_D_S/LI_D/FLD/FSD/SRAI/XOR/AND/SLT/REM/MAX/DIV`。结论: + +- **不新增任何 MachineOp 成员**;`SQRT_S` 与 `FSQRT_S` 为同值别名(`enum` alias),保持不动; +- `asm_emit._OP_NAMES` 已覆盖全部所需成员(含 `LI_D`),无需新增映射; +- MachineOp 的 append-only 约定同样适用于课题 29。 + +### 4.7 `asm_emit.py`:未知 MachineOp fail-loud + +`AsmEmitter._format_instr` 当前对 `_OP_NAMES` 未命中项返回 `# {op.value} {comment}` 注释行(静默丢弃)。改为抛 `ValueError(f"no assembly mapping for MachineOp {instr.op.name} ({instr.op.value})")`。经核对,全部在产 MachineOp 均已映射(`LABEL`/`SECTION` 在 `emit()` 中先行处理,`GLOBL/SIZE/TYPE` 无生产者),改动安全。 + +### 4.8 `riscv_encoder.py`:F/D fail-loud + +新增: + +```python +class UnsupportedInstructionError(ValueError): + """Raised when assembly cannot be encoded by the RV32IM encoder.""" +``` + +并在 `_encode_line` 的兜底分支中先判 F/D(`fld/fsd/flw/fsw`、`fadd.*/fsub.*/fmul.*/fdiv.*/fsqrt.*/fmin.*/fmax.*/fabs.*/fneg.*/flt.*/fle.*/feq.*/fcvt.*/fmv.*/fsgnj*`、`li.d`),抛 `UnsupportedInstructionError`,信息含助记符与「not supported by the RV32IM encoder」;其他未知指令保持 `ValueError: Unknown instruction: ...`。`_expand_pseudo` 对 F/D 行不做展开、原样透传(保持现状),保证最终仍走该报错路径而非被悄悄改写。 + +### 4.9 `compiler.py` / `main.py`:配置与 CLI 接线 + +`CompilerConfig` 新增字段(默认值即零回归): + +```python +extended_isel: bool = False +enable_fp64: bool = True +use_hardware_sqrt: bool = False +``` + +`CompilerDriver._generate_riscv_linear`:`extended_isel=True` 时构造 `ExtendedInstructionSelector(program, enable_fp64=..., use_hardware_sqrt=...)`,否则保持 `InstructionSelector`。`compile()` 在 codegen 前收集模式警告(§2.3.2 矩阵),写入 `warnings`。 + +`main.py`:新增/接线 `--extended-isel`、`--no-fp64`、`--hardware-sqrt`,并在 `args_to_config` 中映射 3 个字段。 + +### 4.10 测试与集成回归 + +- 用例与文件见 §3;新增文件 `tests/test_riscv_encoder_fd.py`、`tests/test_extended_isel_cli.py`,扩展 `tests/test_inst_select_ext.py`。 +- 全量:`make test`;静态检查:`make lint`(flake8 + mypy)。 +- 零回归基准:对同一 DSL/ONNX 输入,比较改动前后默认管线输出(应逐字节一致)。 +- 与课题 29 并行:OpCode 分区合入时若冲突,只允许「双分区共存」,不得重排。 + +--- + +## 五、附录 + +### 5.1 生成汇编示例(目标行为) + +**示例 A:整数 MIN + ABS**(IR:`d = min(a,b)`,`r = abs(x)`;去注释) + +``` + slt __min_slt_1, a, b + sub __min_sub_2, b, a + and __min_and_3, __min_slt_1, __min_sub_2 + add d, a, __min_and_3 + srai __abs_srai_4, x, 31 + xor __abs_xor_5, x, __abs_srai_4 + sub r, __abs_xor_5, __abs_srai_4 +``` + +**示例 B:SQRT 软件与硬件** + +``` +# use_hardware_sqrt=False(f32) + mv a0, x + call sqrtf + mv d, a0 + +# use_hardware_sqrt=True(f32) + fsqrt.s d, x + +# use_hardware_sqrt=True(f64,enable_fp64=True) + fsqrt.d d, x +``` + +**示例 C:float64(D 扩展)** + +IR: +`d1 = load_f64(p)`;`store_f64(p, d1)`;`d2 = load_const_f64(1.5)`;`d3 = fadd_d(d1, d2)`;`t = fcmp_l_d(d1, d3)`;`y = fcvt_s_d(d3)` + +汇编(去注释): + +``` + fld d1, p + fsd d1, p + li.d d2, 4609434218613702656 # 1.5 的 IEEE-754 位模式 + fadd.d d3, d1, d2 + flt.d t, d1, d3 + fcvt.s.d y, d3 +``` + +**示例 D:IDIV / REM / MOD** + +``` + div d1, a, b + rem d2, a, b + rem d3, a, b +``` + +### 5.2 报错示例(目标行为) + +``` +>>> selector = ExtendedInstructionSelector(prog, enable_fp64=False) +>>> selector.run() # IR 含 OpCode.FADD_D +ValueError: opcode 'fadd_d' requires enable_fp64=True (ExtendedInstructionSelector) + +>>> selector.run() # IR 为 f64 类型的 ADD +ValueError: instruction 'add' involves FLOAT64 but enable_fp64=False (ExtendedInstructionSelector) + +>>> assemble_to_binary("fadd.d f0, f1, f2") +UnsupportedInstructionError: F/D instruction 'fadd.d' is not supported by the +RV32IM encoder (Topic 28: final encoding out of scope; output is assembly text) + +>>> assemble_to_binary("frobnicate x1, x2") +ValueError: Unknown instruction: frobnicate +``` + +### 5.3 参考资料 + +- `docs/topics/28-扩展指令选择.md`(课题描述;其中「✅ 已完成」状态与实测不符,实现后需同步修正) +- `docs/ARCHITECTURE.md`(ONNX→RISC-V 双路径总架构) +- 同目录《开发文档.md》(接口契约与实现指南) +- RISC-V ISA Manual Vol.1:RV32I/M 基础指令、F/D 扩展(`fld`/`fsd`/`fadd.d`/`flt.d`/`fcvt.*`) +- RISC-V Calling Convention(浮点 ABI:f0–f7 传参,f0 返回) +- Bit Twiddling Hacks(分支无 min/abs 位技巧) From fbecc12e052ec41fb74735677af28869cbcb5290 Mon Sep 17 00:00:00 2001 From: Seven Gao <799889633@qq.com> Date: Mon, 14 Sep 2026 22:52:33 +0800 Subject: [PATCH 3/6] fix(topic28): materialize float literals and enforce dtype signatures - sqrt f32 literals materialize the exact IEEE-754 bit pattern instead of int()-truncating them; f64 literals fail loud (no 64-bit materialization) - min/max/abs/idiv/rem materialize literal operands before SUB/AND/ADD - fix broken 0/1 mask in integer min (returned max) and stop using the encoder max pseudo whose fallback branch is a pre-existing defect - _check_dtype enforces operand/dest same-type; fp64 handlers validate dest presence and exact (dest, operands) dtype signatures - sync design/development docs with the corrected sequences and guards --- ...00\345\217\221\346\226\207\346\241\243.md" | 10 +- ...76\350\256\241\346\226\207\346\241\243.md" | 27 +- scratchv/backend/inst_select_ext.py | 253 +++++++++++++----- 3 files changed, 206 insertions(+), 84 deletions(-) diff --git "a/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\345\274\200\345\217\221\346\226\207\346\241\243.md" "b/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\345\274\200\345\217\221\346\226\207\346\241\243.md" index 2b509c5..b0530d0 100644 --- "a/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\345\274\200\345\217\221\346\226\207\346\241\243.md" +++ "b/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\345\274\200\345\217\221\346\226\207\346\241\243.md" @@ -128,7 +128,7 @@ |----------|---------|----------| | `sqrt` | `_select_sqrt` | MV/CALL 或 `SQRT_S`/`SQRT_D` | | `min` | `_select_min` | `FMIN_D` 或 SLT/SUB/AND/ADD | -| `max` | `_select_max` | `FMAX_D` 或 `MachineOp.MAX` | +| `max` | `_select_max` | `FMAX_D` 或 SLT/SUB/AND/ADD(分支无,不复用 `max` 伪指令) | | `abs` | `_select_abs` | `FABS_D` 或 SRAI/XOR/SUB | | `idiv` | `_select_idiv` | DIV | | `rem` | `_select_rem` | REM | @@ -283,7 +283,7 @@ python -c "from scratchv.ir.types import OpCode; print([o.value for o in OpCode] return MachineOperand.vreg(f"__{prefix}_{self._temp_counter}") ``` -`_select_min` 的三个临时值 `__min_slt_{n}/__min_sub_{n}/__min_and_{n}`,`_select_abs` 的两个临时值 `__abs_srai_{n}/__abs_xor_{n}`。 +`_select_min/_select_max` 的三个临时值 `__{min,max}_slt_{n}/__{min,max}_mask_{n}/__{min,max}_sub_{n}`(`and` 就地写 `diff`,掩码为 0/-1),`_select_abs` 的两个临时值 `__abs_srai_{n}/__abs_xor_{n}`。 **P3 `store_f64` 顺序**: @@ -294,7 +294,7 @@ python -c "from scratchv.ir.types import OpCode; print([o.value for o in OpCode] self._emit(MachineOp.FSD, val, addr, comment="fsd (store f64)") ``` -**P4 `sqrt` 立即数**:软件路径中 `src.kind == "imm"` 时发射 `LI a0, src`,否则 `MV a0, src`;随后 `CALL`,再 `MV dst, a0`(f64 时函数名 `sqrt`,f32 为 `sqrtf`)。 +**P4 `sqrt` 立即数**:软件路径中,源为 f32 字面量时按 IEEE-754 位模式发射 `LI a0, `(禁止 `int()` 截断),否则 `MV a0, src`;f64 字面量抛 `ValueError`(无 64 位物化手段);随后 `CALL`,再 `MV dst, a0`(f64 时函数名 `sqrt`,f32 为 `sqrtf`)。 **P5 常量位模式**: @@ -312,7 +312,7 @@ python -c "from scratchv.ir.types import OpCode; print([o.value for o in OpCode] (新增 `import struct`;删除对 `int(raw_val)` 的截断与默认 0.0。) -**P6 dtype 守卫**:新增 `_check_dtype`;在 `_select_sqrt`(允许 f32/f64)、`_select_min/_select_max/_select_abs`(允许 i32/i64/f64)、`_select_idiv/_select_rem/_select_mod`(仅 i32)入口调用;`STORE_F64` 只校验 operand dtype 为 f64。 +**P6 dtype 守卫**:新增 `_check_dtype`(dest/operand 均须在允许集合内且同型);在 `_select_sqrt`(允许 f32/f64)、`_select_min/_select_max/_select_abs`(允许 i32/i64/f64)、`_select_idiv/_select_rem/_select_mod`(仅 i32)入口调用;`_check_signature` 为 f64 专用 handler 校验 dest 必备与精确 (dest, operands) 签名;`STORE_F64` 校验地址为 i32/i64、值为 f64;整数 handler 的常量 operand 先 `LI` 物化(f32 按位模式,f64 字面量抛 `ValueError`)。 **P7 清理**:删除文件尾部「New MachineOp entries ... register_alloc」失效注释;核对 `supported_ops` 与 §2.3.1 一一对应。 @@ -367,7 +367,7 @@ python -c "from scratchv.ir.types import OpCode; print([o.value for o in OpCode] | `TestAsmText::test_min_branchless_asm` | 逐行等于设计文档 §3 用例 2 的 4 行序列(去注释) | | `TestAsmText::test_abs_branchless_asm` | 逐行等于 3 行序列 | | `TestAsmText::test_sqrt_software_uses_a0_and_call` | 含 `mv a0, x`、`call sqrtf`、`mv d, a0`;f64 为 `call sqrt` | -| `TestAsmText::test_sqrt_immediate_uses_li` | 立即数源含 `li a0, ` | +| `TestAsmText::test_sqrt_immediate_uses_bit_pattern` | 立即数源含 `li a0, `(2.5→1075838976、4.0→1082130432) | | `TestAsmText::test_sqrt_hardware_f32_f64` | `fsqrt.s d, x`;`fsqrt.d d, x` | | `TestAsmText::test_fp64_dtype_driven_add` | FLOAT64 的 `ADD` 输出 `fadd.d` | | `TestAsmText::test_store_f64_operand_order` | `store_f64(p, v)` 输出 `fsd v, p` | diff --git "a/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\350\256\276\350\256\241\346\226\207\346\241\243.md" "b/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\350\256\276\350\256\241\346\226\207\346\241\243.md" index b51bba7..6c67f85 100644 --- "a/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\350\256\276\350\256\241\346\226\207\346\241\243.md" +++ "b/docs/topics/28-\346\211\251\345\261\225\346\214\207\344\273\244\351\200\211\346\213\251-\350\256\276\350\256\241\346\226\207\346\241\243.md" @@ -137,9 +137,9 @@ ScratchV 中的等价 Python 表示为 `Instruction(opcode=..., dest=Value, oper |--------|-------|--------------------------|---------------------------|------------------| | `SQRT` | f32 | `mv a0, src` → `call sqrtf` → `mv dst, a0` | `use_hardware_sqrt=True`:`fsqrt.s` | MV / CALL / SQRT_S | | `SQRT` | f64 | `call sqrt`(libm double) | `use_hardware_sqrt=True` 且 `enable_fp64=True`:`fsqrt.d` | CALL / SQRT_D | -| `MIN` | i32/i64 | 分支无:`slt; sub; and; add` | — | SLT / SUB / AND / ADD | +| `MIN` | i32/i64 | 分支无:`slt; sub mask; sub diff; and; add` | — | SLT / SUB / AND / ADD | | `MIN` | f64 | — | `fmin.d`(需 `enable_fp64=True`) | FMIN_D | -| `MAX` | i32/i64 | `max` 伪指令(编码器展开为 bge 分支序列) | — | MAX | +| `MAX` | i32/i64 | 分支无:`slt; sub mask; sub diff; and; add`(不使用 `max` 伪指令) | — | SLT / SUB / AND / ADD | | `MAX` | f64 | — | `fmax.d`(需 `enable_fp64=True`) | FMAX_D | | `ABS` | i32/i64 | 分支无:`srai 31; xor; sub` | — | SRAI / XOR / SUB | | `ABS` | f64 | — | `fabs.d`(等价 `fsgnjx.d`;需 `enable_fp64=True`) | FABS_D | @@ -164,9 +164,10 @@ ScratchV 中的等价 Python 表示为 `Instruction(opcode=..., dest=Value, oper #### 2.2.3 关键序列正确性 -- **整数 MIN(分支无)**:`slt t,a,b; sub d,b,a; and t,t,d; add dst,a,t`。a`(禁止 `int()` 截断),否则 `MV a0, src`;f64 字面量抛 `ValueError`(无 64 位物化手段);随后 `CALL`,再 `MV dst, a0` | | P5 | 常量位模式 | `_select_load_const_f64` 用 `struct.unpack(" 扩展选择器对 `ADD/SUB/MUL/DIV/NEG/LOAD/STORE/LOAD_CONST` 的类型驱动分支做同样的门控与顺序修正(`STORE` → `_select_store_f64` 同样使用 P3 修正后的顺序)。 diff --git a/scratchv/backend/inst_select_ext.py b/scratchv/backend/inst_select_ext.py index 266f482..a861dba 100644 --- a/scratchv/backend/inst_select_ext.py +++ b/scratchv/backend/inst_select_ext.py @@ -24,7 +24,7 @@ from scratchv.backend.machine_types import ( MachineOp, MachineOperand, ) -from scratchv.ir.types import DataType, Instruction, Program +from scratchv.ir.types import DataType, Instruction, Program, Value class ExtendedInstructionSelector(InstructionSelector): @@ -63,10 +63,21 @@ def __init__(self, program: Program, *, self.use_hardware_sqrt = use_hardware_sqrt self._current_dtype: Optional[DataType] = None self._temp_counter = 0 + # Names of values defined by an emitted instruction (e.g. the dest + # of ``load_const_f64``). Such "constants" are addressable vregs + # and must not be re-materialized as literals. + self._defined_names: set[str] = set() def run(self) -> list: """Select instructions, resetting the temp counter for determinism.""" self._temp_counter = 0 + self._defined_names = { + instr.dest.name + for func in self.program.functions + for block in func.blocks + for instr in block.instructions + if instr.dest is not None + } return super().run() # ------------------------------------------------------------------ @@ -93,35 +104,43 @@ def _select_sqrt(self, instr: Instruction) -> None: If hardware F extension is available, use ``fsqrt.s``. Otherwise emit a library call to ``sqrtf`` (float) or ``sqrt`` (double). + + Float32 literal arguments are materialized by their exact IEEE-754 + bit pattern (never ``int()``-truncated); float64 literals raise + ``ValueError`` (no 64-bit materialization in RV32IM). """ if self._involves_fp64(instr): self._require_fp64(instr) self._check_dtype( instr, (DataType.FLOAT32, DataType.FLOAT64), "sqrt") - src = self._op(instr, 0) dst = self._dst(instr) - dtype = instr.dest.dtype if instr.dest else DataType.FLOAT32 + dtype = instr.dest.dtype if self.use_hardware_sqrt: + src = self._materialized_op(instr, 0, prefix="sqrt_imm") if dtype == DataType.FLOAT64: self._emit(MachineOp.SQRT_D, dst, src, comment="fsqrt.d (hardware)") else: self._emit(MachineOp.SQRT_S, dst, src, comment="fsqrt.s (hardware)") + return + + # Library call: argument in a0, result in a0 + if self._literal_operand(instr.operands[0]): + bits = self._constant_bits(instr, instr.operands[0]) + self._emit(MachineOp.LI, MachineOperand.reg("a0"), + MachineOperand.immediate(bits), + comment="sqrt arg -> a0") else: - # Library call: argument in a0, result in a0 - if src.kind == "imm": - self._emit(MachineOp.LI, MachineOperand.reg("a0"), src, - comment="sqrt arg -> a0") - else: - self._emit(MachineOp.MV, MachineOperand.reg("a0"), src, - comment="sqrt arg -> a0") - func = "sqrt" if dtype == DataType.FLOAT64 else "sqrtf" - self._emit(MachineOp.CALL, comment=func) - if dst: - self._emit(MachineOp.MV, dst, MachineOperand.reg("a0"), - comment="sqrt result") + src = self._op(instr, 0) + self._emit(MachineOp.MV, MachineOperand.reg("a0"), src, + comment="sqrt arg -> a0") + func = "sqrt" if dtype == DataType.FLOAT64 else "sqrtf" + self._emit(MachineOp.CALL, comment=func) + if dst: + self._emit(MachineOp.MV, dst, MachineOperand.reg("a0"), + comment="sqrt result") # ------------------------------------------------------------------ # min / max @@ -130,11 +149,16 @@ def _select_sqrt(self, instr: Instruction) -> None: def _select_min(self, instr: Instruction) -> None: """Select instruction for min(a, b). - Integer min (branchless): - slt tmp, a, b # tmp = (a < b) - sub dst, b, a # diff = b - a - and tmp, tmp, dst # mask = tmp & diff - add dst, a, tmp # dst = a + mask + Integer min (branchless, 0/-1 mask): + slt tmp, a, b # tmp = (a < b) + sub mask, x0, tmp # mask = 0 or -1 + sub diff, a, b # diff = a - b + and diff, diff, mask # diff = (a < b) ? a - b : 0 + add dst, b, diff # dst = (a < b) ? a : b + + Literal operands are materialized first: ``sub``/``and``/``add`` + have no immediate form, and f64 literals have no materialization + at all (fail-loud). """ if self._involves_fp64(instr): self._require_fp64(instr) @@ -142,32 +166,39 @@ def _select_min(self, instr: Instruction) -> None: instr, (DataType.INT32, DataType.INT64, DataType.FLOAT64), "min") - a = self._op(instr, 0) - b = self._op(instr, 1) + a = self._materialized_op(instr, 0, prefix="min_const") + b = self._materialized_op(instr, 1, prefix="min_const") dst = self._dst(instr) - if instr.dest is not None and instr.dest.dtype == DataType.FLOAT64: + if instr.dest.dtype == DataType.FLOAT64: # Use FMIN.D pseudo (expands to branchless sequence) self._emit(MachineOp.FMIN_D, dst, a, b, comment="fmin.d") else: tmp = self._fresh_temp("min_slt") self._emit(MachineOp.SLT, tmp, a, b, comment="min: slt") + mask = self._fresh_temp("min_mask") + self._emit(MachineOp.SUB, mask, MachineOperand.reg("x0"), tmp, + comment="min: mask") diff = self._fresh_temp("min_sub") - self._emit(MachineOp.SUB, diff, b, a, comment="min: sub") - and_tmp = self._fresh_temp("min_and") + self._emit(MachineOp.SUB, diff, a, b, comment="min: sub") + self._emit(MachineOp.AND, diff, diff, mask, comment="min: and") self._emit( - MachineOp.AND, and_tmp, tmp, diff, comment="min: and" + MachineOp.ADD, dst, b, diff, comment="min result" ) - if dst: - self._emit( - MachineOp.ADD, dst, a, and_tmp, comment="min result" - ) def _select_max(self, instr: Instruction) -> None: """Select instruction for max(a, b). - Uses the existing `max` pseudo-instruction from the base selector, - or a branchless sequence if not available. + Integer max (branchless, 0/-1 mask, symmetric to min): + slt tmp, a, b # tmp = (a < b) + sub mask, x0, tmp # mask = 0 or -1 + sub diff, b, a # diff = b - a + and diff, diff, mask # diff = (a < b) ? b - a : 0 + add dst, a, diff # dst = (a < b) ? b : a + + The machine-level ``max`` pseudo is not used: its encoder fallback + branch is broken (falls back to x0 instead of rs2), so the pseudo + computes max incorrectly for a < b. """ if self._involves_fp64(instr): self._require_fp64(instr) @@ -175,15 +206,24 @@ def _select_max(self, instr: Instruction) -> None: instr, (DataType.INT32, DataType.INT64, DataType.FLOAT64), "max") - a = self._op(instr, 0) - b = self._op(instr, 1) + a = self._materialized_op(instr, 0, prefix="max_const") + b = self._materialized_op(instr, 1, prefix="max_const") dst = self._dst(instr) - if instr.dest is not None and instr.dest.dtype == DataType.FLOAT64: + if instr.dest.dtype == DataType.FLOAT64: self._emit(MachineOp.FMAX_D, dst, a, b, comment="fmax.d") else: - # Use existing MAX pseudo (base selector has this) - self._emit(MachineOp.MAX, dst, a, b, comment="max") + tmp = self._fresh_temp("max_slt") + self._emit(MachineOp.SLT, tmp, a, b, comment="max: slt") + mask = self._fresh_temp("max_mask") + self._emit(MachineOp.SUB, mask, MachineOperand.reg("x0"), tmp, + comment="max: mask") + diff = self._fresh_temp("max_sub") + self._emit(MachineOp.SUB, diff, b, a, comment="max: sub") + self._emit(MachineOp.AND, diff, diff, mask, comment="max: and") + self._emit( + MachineOp.ADD, dst, a, diff, comment="max result" + ) # ------------------------------------------------------------------ # abs @@ -203,10 +243,10 @@ def _select_abs(self, instr: Instruction) -> None: instr, (DataType.INT32, DataType.INT64, DataType.FLOAT64), "abs") - src = self._op(instr, 0) + src = self._materialized_op(instr, 0, prefix="abs_const") dst = self._dst(instr) - if instr.dest is not None and instr.dest.dtype == DataType.FLOAT64: + if instr.dest.dtype == DataType.FLOAT64: # fabs.d: clear the sign bit self._emit(MachineOp.FABS_D, dst, src, comment="fabs.d") else: @@ -231,16 +271,16 @@ def _select_abs(self, instr: Instruction) -> None: def _select_idiv(self, instr: Instruction) -> None: """Select instruction for integer division.""" self._check_dtype(instr, (DataType.INT32,), "idiv") - a = self._op(instr, 0) - b = self._op(instr, 1) + a = self._materialized_op(instr, 0, prefix="idiv_const") + b = self._materialized_op(instr, 1, prefix="idiv_const") dst = self._dst(instr) self._emit(MachineOp.DIV, dst, a, b, comment="div") def _select_rem(self, instr: Instruction) -> None: """Select instruction for integer remainder.""" self._check_dtype(instr, (DataType.INT32,), "rem") - a = self._op(instr, 0) - b = self._op(instr, 1) + a = self._materialized_op(instr, 0, prefix="rem_const") + b = self._materialized_op(instr, 1, prefix="rem_const") dst = self._dst(instr) self._emit(MachineOp.REM, dst, a, b, comment="rem") @@ -255,6 +295,9 @@ def _select_mod(self, instr: Instruction) -> None: def _select_load_f64(self, instr: Instruction) -> None: """Load a 64-bit float from memory.""" + self._check_signature( + instr, (DataType.FLOAT64,), + (DataType.INT32, DataType.INT64), "load_f64") src = self._op(instr, 0) dst = self._dst(instr) self._emit(MachineOp.FLD, dst, src, comment="fld (load f64)") @@ -265,71 +308,90 @@ def _select_store_f64(self, instr: Instruction) -> None: raise ValueError( "store_f64 requires operands [addr, value], got " f"{len(instr.operands)} operand(s)") + addr = instr.operands[0] + if addr.dtype not in (DataType.INT32, DataType.INT64): + raise ValueError( + "store_f64 requires an INT32/INT64 address operand, got " + f"{addr.dtype.value}") val = instr.operands[1] if val.dtype != DataType.FLOAT64: raise ValueError( "store_f64 requires a FLOAT64 value operand, got " f"{val.dtype.value}") - addr = self._op(instr, 0) - val_op = self._op(instr, 1) - self._emit(MachineOp.FSD, val_op, addr, comment="fsd (store f64)") + addr_op = self._op(instr, 0) + val_op = self._materialized_op(instr, 1, prefix="store_f64_val") + self._emit(MachineOp.FSD, val_op, addr_op, comment="fsd (store f64)") def _select_fadd_d(self, instr: Instruction) -> None: """Add two float64 values.""" - a = self._op(instr, 0) - b = self._op(instr, 1) + self._check_dtype(instr, (DataType.FLOAT64,), "fadd_d") + a = self._materialized_op(instr, 0, prefix="fadd_d_const") + b = self._materialized_op(instr, 1, prefix="fadd_d_const") dst = self._dst(instr) self._emit(MachineOp.FADD_D, dst, a, b, comment="fadd.d") def _select_fsub_d(self, instr: Instruction) -> None: """Subtract two float64 values.""" - a = self._op(instr, 0) - b = self._op(instr, 1) + self._check_dtype(instr, (DataType.FLOAT64,), "fsub_d") + a = self._materialized_op(instr, 0, prefix="fsub_d_const") + b = self._materialized_op(instr, 1, prefix="fsub_d_const") dst = self._dst(instr) self._emit(MachineOp.FSUB_D, dst, a, b, comment="fsub.d") def _select_fmul_d(self, instr: Instruction) -> None: """Multiply two float64 values.""" - a = self._op(instr, 0) - b = self._op(instr, 1) + self._check_dtype(instr, (DataType.FLOAT64,), "fmul_d") + a = self._materialized_op(instr, 0, prefix="fmul_d_const") + b = self._materialized_op(instr, 1, prefix="fmul_d_const") dst = self._dst(instr) self._emit(MachineOp.FMUL_D, dst, a, b, comment="fmul.d") def _select_fdiv_d(self, instr: Instruction) -> None: """Divide two float64 values.""" - a = self._op(instr, 0) - b = self._op(instr, 1) + self._check_dtype(instr, (DataType.FLOAT64,), "fdiv_d") + a = self._materialized_op(instr, 0, prefix="fdiv_d_const") + b = self._materialized_op(instr, 1, prefix="fdiv_d_const") dst = self._dst(instr) self._emit(MachineOp.FDIV_D, dst, a, b, comment="fdiv.d") def _select_fcmp_l_d(self, instr: Instruction) -> None: """Float64 less-than comparison.""" - a = self._op(instr, 0) - b = self._op(instr, 1) + self._check_signature( + instr, (DataType.INT32,), (DataType.FLOAT64,), "fcmp_l_d") + a = self._materialized_op(instr, 0, prefix="fcmp_l_d_const") + b = self._materialized_op(instr, 1, prefix="fcmp_l_d_const") dst = self._dst(instr) self._emit(MachineOp.FLT_D, dst, a, b, comment="flt.d") def _select_fcmp_eq_d(self, instr: Instruction) -> None: """Float64 equality comparison.""" - a = self._op(instr, 0) - b = self._op(instr, 1) + self._check_signature( + instr, (DataType.INT32,), (DataType.FLOAT64,), "fcmp_eq_d") + a = self._materialized_op(instr, 0, prefix="fcmp_eq_d_const") + b = self._materialized_op(instr, 1, prefix="fcmp_eq_d_const") dst = self._dst(instr) self._emit(MachineOp.FEQ_D, dst, a, b, comment="feq.d") def _select_fcvt_s_d(self, instr: Instruction) -> None: """Convert float64 to float32.""" - src = self._op(instr, 0) + self._check_signature( + instr, (DataType.FLOAT32,), (DataType.FLOAT64,), "fcvt_s_d") + src = self._materialized_op(instr, 0, prefix="fcvt_s_d_const") dst = self._dst(instr) self._emit(MachineOp.FCVT_S_D, dst, src, comment="fcvt.s.d") def _select_fcvt_d_s(self, instr: Instruction) -> None: """Convert float32 to float64.""" - src = self._op(instr, 0) + self._check_signature( + instr, (DataType.FLOAT64,), (DataType.FLOAT32,), "fcvt_d_s") + src = self._materialized_op(instr, 0, prefix="fcvt_d_s_const") dst = self._dst(instr) self._emit(MachineOp.FCVT_D_S, dst, src, comment="fcvt.d.s") def _select_load_const_f64(self, instr: Instruction) -> None: """Load a float64 constant (exact IEEE-754 bit pattern).""" + self._check_signature( + instr, (DataType.FLOAT64,), (), "load_const_f64") raw_val = instr.attrs.get("value") if not isinstance(raw_val, (int, float)): raise ValueError( @@ -400,7 +462,8 @@ def _select_neg(self, instr: Instruction) -> None: """Negate: for float64 use fneg.d, for int use sub x0 - x.""" if self._involves_fp64(instr): self._require_fp64(instr) - src = self._op(instr, 0) + self._check_dtype(instr, (DataType.FLOAT64,), "neg") + src = self._materialized_op(instr, 0, prefix="neg_const") dst = self._dst(instr) self._emit(MachineOp.FNEG_D, dst, src, comment="fneg.d") else: @@ -431,23 +494,77 @@ def _require_fp64(self, instr: Instruction) -> None: f"instruction '{instr.opcode.value}' involves FLOAT64 but " f"enable_fp64=False (ExtendedInstructionSelector)") - def _check_dtype(self, instr: Instruction, - allowed: tuple[DataType, ...], - opname: str) -> None: - """Guard dest/operand dtypes against the allowed set.""" + def _materialized_op(self, instr: Instruction, idx: int, + prefix: str = "const") -> MachineOperand: + """Return operand *idx* as a register, materializing literals. + + The base ``_op`` routes every constant through ``int()``; for float + literals that silently changes the value. Literal operands are + instead materialized with an exact ``LI`` (f32 bit pattern) into a + fresh temp; f64 literals fail loud (RV32IM has no 64-bit constant + materialization). + """ + op = instr.operands[idx] + if not self._literal_operand(op): + return MachineOperand.vreg(op.name) + bits = self._constant_bits(instr, op) + tmp = self._fresh_temp(prefix) + self._emit(MachineOp.LI, tmp, MachineOperand.immediate(bits), + comment=f"const {op.const_value!r}") + return tmp + + def _literal_operand(self, op: Value) -> bool: + """True if *op* is a literal constant (no defining instruction).""" + if not (op.is_constant and op.const_value is not None): + return False + return op.name not in self._defined_names + + def _constant_bits(self, instr: Instruction, op: Value) -> int: + """Exact 32-bit materialization of a literal constant operand.""" + if op.dtype == DataType.FLOAT64: + raise ValueError( + f"{instr.opcode.value} cannot materialize a FLOAT64 literal " + f"({op.const_value!r}): RV32IM has no 64-bit constant " + f"materialization; use load_const_f64") + if op.dtype == DataType.FLOAT32: + return self._fp32_bits(op.const_value) + return int(op.const_value) + + @staticmethod + def _fp32_bits(value) -> int: + """IEEE-754 f32 bit pattern as a signed RV32 immediate.""" + return struct.unpack(" None: + """Validate dest existence and exact dtype sets for an opcode.""" if instr.dest is None: raise ValueError(f"{opname} requires a destination value") - allowed_vals = ", ".join(d.value for d in allowed) - if instr.dest.dtype not in allowed: + if instr.dest.dtype not in dest_allowed: + allowed_vals = ", ".join(d.value for d in dest_allowed) raise ValueError( f"{opname} requires destination dtype in " f"({allowed_vals}), got {instr.dest.dtype.value}") for op in instr.operands: - if op.dtype not in allowed: + if op.dtype not in operand_allowed: + allowed_vals = ", ".join(d.value for d in operand_allowed) raise ValueError( f"{opname} requires operand dtype in " f"({allowed_vals}), got {op.dtype.value}") + def _check_dtype(self, instr: Instruction, + allowed: tuple[DataType, ...], + opname: str) -> None: + """Guard dest/operand dtypes; all operands must match the dest type.""" + self._check_signature(instr, allowed, allowed, opname) + for op in instr.operands: + if op.dtype != instr.dest.dtype: + raise ValueError( + f"{opname} requires operands to match destination dtype " + f"{instr.dest.dtype.value}, got {op.dtype.value}") + def _is_fp64(self, instr: Instruction) -> bool: """Legacy compatibility: float64 involvement with fp64 enabled.""" return self.enable_fp64 and self._involves_fp64(instr) From 1ca947c090cf060f3c9bb3bc6c502c9cede607f8 Mon Sep 17 00:00:00 2001 From: Seven Gao <799889633@qq.com> Date: Mon, 14 Sep 2026 22:52:39 +0800 Subject: [PATCH 4/6] fix(topic28): fail loud for fused and classified F/D mnemonics Extend _is_fd_mnemonic with a single-suffix pattern plus fmadd/fnmadd/ fmsub/fnmsub/fclass/fli/fround prefixes so every F/D mnemonic raises UnsupportedInstructionError instead of a plain Unknown instruction error. --- scratchv/backend/riscv_encoder.py | 13 +++++++++---- 1 file changed, 9 insertions(+), 4 deletions(-) diff --git a/scratchv/backend/riscv_encoder.py b/scratchv/backend/riscv_encoder.py index 3de4480..c057964 100644 --- a/scratchv/backend/riscv_encoder.py +++ b/scratchv/backend/riscv_encoder.py @@ -19,13 +19,16 @@ class UnsupportedInstructionError(ValueError): """Raised when assembly cannot be encoded by the RV32IM encoder.""" -# Exact F/D mnemonics plus prefixes covering the F/D instruction families. +# Exact F/D mnemonics plus prefixes/patterns covering the F/D families. _FD_EXACT: frozenset[str] = frozenset({"fld", "fsd", "flw", "fsw", "li.d"}) _FD_PREFIXES: tuple[str, ...] = ( - "fadd.", "fsub.", "fmul.", "fdiv.", "fsqrt.", "fmin.", "fmax.", - "fabs.", "fneg.", "flt.", "fle.", "feq.", "fcvt.", "fmv.", - "fsgnj", "fsgnjn", "fsgnjx", + "fcvt.", # e.g. fcvt.s.d / fcvt.d.s (multi-suffix) + "fmv.", # e.g. fmv.x.w / fmv.w.x + "fmadd.", "fnmadd.", "fmsub.", "fnmsub.", + "fclass.", "fli.", "fround.", ) +# Single-suffix families: fadd.d, fsqrt.s, fsgnjx.d, fmin.s, ... +_FD_PATTERN = re.compile(r"^f[a-z0-9]+\.[sd]$") def _is_fd_mnemonic(op: str) -> bool: @@ -33,6 +36,8 @@ def _is_fd_mnemonic(op: str) -> bool: op = op.lower() if op in _FD_EXACT: return True + if _FD_PATTERN.match(op): + return True return any(op.startswith(prefix) for prefix in _FD_PREFIXES) From 8ad1ceb626774d0c85280e0291b10500f005532c Mon Sep 17 00:00:00 2001 From: Seven Gao <799889633@qq.com> Date: Mon, 14 Sep 2026 22:52:44 +0800 Subject: [PATCH 5/6] test(topic28): cover literal, dtype, min/max and encoder regressions - sqrt immediate bit patterns, f64 literal rejection, load_const_f64 operand - mixed-dtype rejection and fp64 missing-dest / signature guards - min/max/abs semantics via a local evaluator and tinyfive execution of the greedily allocated binary (immediate operands, both comparison orders) - fused/classified F/D mnemonics in the encoder fail-loud suite --- tests/test_inst_select_ext.py | 524 ++++++++++++++++++++++++++++++++- tests/test_riscv_encoder_fd.py | 15 + 2 files changed, 528 insertions(+), 11 deletions(-) diff --git a/tests/test_inst_select_ext.py b/tests/test_inst_select_ext.py index a136a39..6440f32 100644 --- a/tests/test_inst_select_ext.py +++ b/tests/test_inst_select_ext.py @@ -3,11 +3,15 @@ import pytest from scratchv.backend.asm_emit import AsmEmitter from scratchv.backend.inst_select_ext import ExtendedInstructionSelector +from scratchv.backend.register_alloc import ( + MachineOp, RegisterAllocator, +) +from scratchv.backend.riscv_encoder import REG_MAP, assemble_to_binary from scratchv.ir.builder import IRBuilder from scratchv.ir.types import ( # noqa: F401 OpCode, Value, DataType, ) -from scratchv.backend.register_alloc import MachineOp +from scratchv.simulator.tinyfive import ProfiledMachine class TestExtendedSelectorBasic: @@ -211,7 +215,7 @@ def test_load_store(self): DISPATCH_EXPECTED_OP = { "sqrt": MachineOp.CALL, "min": MachineOp.SLT, - "max": MachineOp.MAX, + "max": MachineOp.SLT, # branchless sequence (see Topic 28 review F4) "abs": MachineOp.SRAI, "idiv": MachineOp.DIV, "rem": MachineOp.REM, @@ -294,6 +298,61 @@ def _clean_asm_lines(asm: str) -> list: ] +def _is_immediate_token(token: str) -> bool: + try: + int(token, 0) + return True + except ValueError: + return False + + +def _eval_int_sequence(asm: str, regs: dict) -> dict: + """Evaluate the integer sequences emitted by the extended selector. + + Test-local mini evaluator over the emitted assembly (vreg names act as + registers). It exists so literal materialization and min/max/abs + semantics can be checked without an external simulator. + """ + values = {"x0": 0} + values.update(regs) + + def val(token: str) -> int: + if _is_immediate_token(token): + return int(token, 0) + return values.get(token, 0) + + for line in _clean_asm_lines(asm): + tokens = line.replace(",", " ").split() + if not tokens or line.startswith(".") or line.endswith(":"): + continue + op, args = tokens[0], tokens[1:] + if op in ("li", "mv"): + values[args[0]] = val(args[1]) + elif op == "slt": + values[args[0]] = 1 if val(args[1]) < val(args[2]) else 0 + elif op == "sub": + values[args[0]] = val(args[1]) - val(args[2]) + elif op == "and": + values[args[0]] = val(args[1]) & val(args[2]) + elif op == "add": + values[args[0]] = val(args[1]) + val(args[2]) + elif op == "srai": + values[args[0]] = val(args[1]) >> val(args[2]) + elif op == "xor": + values[args[0]] = val(args[1]) ^ val(args[2]) + else: + raise AssertionError(f"unexpected integer sequence op: {line}") + return values + + +def _allocated_pipeline(builder, **selector_kwargs): + """Select, greedily allocate and emit one builder program.""" + instrs = ExtendedInstructionSelector( + builder.program, **selector_kwargs).run() + allocated = RegisterAllocator(instrs, mode="greedy").run() + return allocated, AsmEmitter(allocated).emit() + + class TestDispatchCoverage: """Every new opcode must dispatch to its handler (Topic 28).""" @@ -343,9 +402,11 @@ def test_unique_temps_two_mins(self): i.dst.value for i in instrs if i.dst is not None and i.dst.value.startswith("__min") ] - assert len(names) == 6 - assert len(names) == len(set(names)) - assert sorted(names, key=lambda n: int(n.rsplit("_", 1)[1])) == names + # ``and`` writes its result in place, so the sub temp repeats. + unique = list(dict.fromkeys(names)) + assert len(unique) == 6 + assert sorted( + unique, key=lambda n: int(n.rsplit("_", 1)[1])) == unique class TestAsmText: @@ -364,11 +425,33 @@ def test_min_branchless_asm(self): seq = [ln for ln in lines if "__min" in ln] assert seq == [ "slt __min_slt_1, a, b", - "sub __min_sub_2, b, a", - "and __min_and_3, __min_slt_1, __min_sub_2", - f"add {dest.name}, a, __min_and_3", + "sub __min_mask_2, x0, __min_slt_1", + "sub __min_sub_3, a, b", + "and __min_sub_3, __min_sub_3, __min_mask_2", + f"add {dest.name}, b, __min_sub_3", ] + def test_max_branchless_asm(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + a = builder.make_value(name="a", dtype=DataType.INT32) + b = builder.make_value(name="b", dtype=DataType.INT32) + dest = builder.max(a, b) + + instrs = ExtendedInstructionSelector(builder.program).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + seq = [ln for ln in lines if "__max" in ln] + assert seq == [ + "slt __max_slt_1, a, b", + "sub __max_mask_2, x0, __max_slt_1", + "sub __max_sub_3, b, a", + "and __max_sub_3, __max_sub_3, __max_mask_2", + f"add {dest.name}, a, __max_sub_3", + ] + ops = {i.op for i in instrs if i.op != MachineOp.LABEL} + assert MachineOp.MAX not in ops + def test_abs_branchless_asm(self): builder = IRBuilder() builder.new_function("test") @@ -409,16 +492,51 @@ def test_sqrt_software_f64_calls_sqrt(self): lines = _clean_asm_lines(AsmEmitter(instrs).emit()) assert "call sqrt" in lines - def test_sqrt_immediate_uses_li(self): + @pytest.mark.parametrize("value,bits", [ + (2.5, 1075838976), # 0x40200000 + (4.0, 1082130432), # 0x40800000 + ]) + def test_sqrt_immediate_uses_bit_pattern(self, value, bits): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + imm = builder.make_const(value, dtype=DataType.FLOAT32) + dest = builder.sqrt(imm) + + instrs = ExtendedInstructionSelector(builder.program).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + assert f"li a0, {bits}" in lines + assert "call sqrtf" in lines + assert f"mv {dest.name}, a0" in lines + + def test_sqrt_immediate_is_not_truncated(self): builder = IRBuilder() builder.new_function("test") builder.new_block("entry") - imm = builder.make_const(4.0, dtype=DataType.FLOAT32) + imm = builder.make_const(2.5, dtype=DataType.FLOAT32) builder.sqrt(imm) instrs = ExtendedInstructionSelector(builder.program).run() lines = _clean_asm_lines(AsmEmitter(instrs).emit()) - assert "li a0, 4" in lines + # Regression: the old path emitted ``li a0, 2`` for 2.5. + assert "li a0, 2" not in lines + + def test_sqrt_hardware_immediate_is_materialized(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + imm = builder.make_const(2.0, dtype=DataType.FLOAT32) + dest = builder.sqrt(imm) + + instrs = ExtendedInstructionSelector( + builder.program, use_hardware_sqrt=True).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + assert "li __sqrt_imm_1, 1073741824" in lines + assert f"fsqrt.s {dest.name}, __sqrt_imm_1" in lines + # No hardware instruction may carry an immediate source operand. + assert not any( + ln.startswith("fsqrt.s") and "__sqrt_imm" not in ln + for ln in lines) def test_sqrt_hardware_f32_f64(self): builder = IRBuilder() @@ -632,5 +750,389 @@ def test_fp64_ops_gated(self): assert value not in ops +class TestFloatLiteralMaterialization: + """F1: float literals keep their exact value (Topic 28 review).""" + + @pytest.mark.parametrize("value", [2.5, -2.5, 1.5]) + def test_sqrt_f64_immediate_raises(self, value): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + imm = builder.make_const(value, dtype=DataType.FLOAT64) + builder.sqrt(imm, dtype=DataType.FLOAT64) + with pytest.raises(ValueError, match="FLOAT64 literal"): + ExtendedInstructionSelector(builder.program).run() + + def test_sqrt_hardware_f64_immediate_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + imm = builder.make_const(2.5, dtype=DataType.FLOAT64) + builder.sqrt(imm, dtype=DataType.FLOAT64) + with pytest.raises(ValueError, match="FLOAT64 literal"): + ExtendedInstructionSelector( + builder.program, use_hardware_sqrt=True).run() + + def test_fadd_d_f64_immediate_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + imm = builder.make_const(1.5, dtype=DataType.FLOAT64) + builder.fadd_d(imm, dx) + with pytest.raises(ValueError, match="FLOAT64 literal"): + ExtendedInstructionSelector(builder.program).run() + + def test_store_f64_f64_immediate_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + ptr = builder.make_value(name="p", dtype=DataType.INT32) + imm = builder.make_const(1.5, dtype=DataType.FLOAT64) + builder.store_f64(ptr, imm) + with pytest.raises(ValueError, match="FLOAT64 literal"): + ExtendedInstructionSelector(builder.program).run() + + def test_load_const_f64_result_is_an_operand(self): + """A value materialized by load_const_f64 is a vreg, not a literal.""" + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + const = builder.load_const_f64(1.5) + dest = builder.fadd_d(const, dx) + + instrs = ExtendedInstructionSelector(builder.program).run() + lines = _clean_asm_lines(AsmEmitter(instrs).emit()) + assert f"fadd.d {dest.name}, {const.name}, dx" in lines + + +class TestDtypeSignatures: + """F2/F5: dest and operands must match the opcode dtype signature.""" + + def test_min_mixed_i32_f64_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + a = builder.make_value(name="a", dtype=DataType.INT32) + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + dest = builder.make_value(name="r", dtype=DataType.INT32) + builder._emit(OpCode.MIN, dest, [a, dx]) + with pytest.raises(ValueError, match="min"): + ExtendedInstructionSelector(builder.program).run() + + def test_max_mixed_i64_f64_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + a = builder.make_value(name="a", dtype=DataType.INT64) + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + dest = builder.make_value(name="r", dtype=DataType.INT64) + builder._emit(OpCode.MAX, dest, [a, dx]) + with pytest.raises(ValueError, match="max"): + ExtendedInstructionSelector(builder.program).run() + + def test_sqrt_f64_operand_f32_dest_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + builder.sqrt(dx) # dest defaults to FLOAT32 + with pytest.raises(ValueError, match="sqrt"): + ExtendedInstructionSelector(builder.program).run() + + def test_abs_f64_operand_i64_dest_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + builder.abs(dx, dtype=DataType.INT64) + with pytest.raises(ValueError, match="abs"): + ExtendedInstructionSelector(builder.program).run() + + def test_fadd_d_i32_operands_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + a = builder.make_value(name="a", dtype=DataType.INT32) + c = builder.make_value(name="c", dtype=DataType.INT32) + dest = builder.make_value(name="r", dtype=DataType.FLOAT64) + builder._emit(OpCode.FADD_D, dest, [a, c]) + with pytest.raises(ValueError, match="fadd_d"): + ExtendedInstructionSelector(builder.program).run() + + def test_fcmp_l_d_f32_operands_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + x = builder.make_value(name="x", dtype=DataType.FLOAT32) + y = builder.make_value(name="y", dtype=DataType.FLOAT32) + dest = builder.make_value(name="r", dtype=DataType.INT32) + builder._emit(OpCode.FCMP_L_D, dest, [x, y]) + with pytest.raises(ValueError, match="fcmp_l_d"): + ExtendedInstructionSelector(builder.program).run() + + def test_fcmp_l_d_f64_dest_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + dy = builder.make_value(name="dy", dtype=DataType.FLOAT64) + dest = builder.make_value(name="r", dtype=DataType.FLOAT64) + builder._emit(OpCode.FCMP_L_D, dest, [dx, dy]) + with pytest.raises(ValueError, match="fcmp_l_d"): + ExtendedInstructionSelector(builder.program).run() + + def test_fcvt_s_d_f32_operand_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + x = builder.make_value(name="x", dtype=DataType.FLOAT32) + dest = builder.make_value(name="r", dtype=DataType.FLOAT32) + builder._emit(OpCode.FCVT_S_D, dest, [x]) + with pytest.raises(ValueError, match="fcvt_s_d"): + ExtendedInstructionSelector(builder.program).run() + + def test_fcvt_d_s_f64_operand_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + dx = builder.make_value(name="dx", dtype=DataType.FLOAT64) + dest = builder.make_value(name="r", dtype=DataType.FLOAT64) + builder._emit(OpCode.FCVT_D_S, dest, [dx]) + with pytest.raises(ValueError, match="fcvt_d_s"): + ExtendedInstructionSelector(builder.program).run() + + def test_load_f64_i32_dest_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + ptr = builder.make_value(name="p", dtype=DataType.INT32) + dest = builder.make_value(name="r", dtype=DataType.INT32) + builder._emit(OpCode.LOAD_F64, dest, [ptr]) + with pytest.raises(ValueError, match="load_f64"): + ExtendedInstructionSelector(builder.program).run() + + def test_store_f64_f64_addr_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + addr = builder.make_value(name="p", dtype=DataType.FLOAT64) + val = builder.make_value(name="v", dtype=DataType.FLOAT64) + builder.store_f64(addr, val) + with pytest.raises(ValueError, match="address"): + ExtendedInstructionSelector(builder.program).run() + + def test_store_f64_i32_value_raises(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + ptr = builder.make_value(name="p", dtype=DataType.INT32) + val = builder.make_value(name="v", dtype=DataType.INT32) + builder.store_f64(ptr, val) + with pytest.raises(ValueError, match="FLOAT64 value"): + ExtendedInstructionSelector(builder.program).run() + + +FP64_MISSING_DEST_CASES = [ + ("fadd_d", OpCode.FADD_D, + [DataType.FLOAT64, DataType.FLOAT64], {}), + ("fsub_d", OpCode.FSUB_D, + [DataType.FLOAT64, DataType.FLOAT64], {}), + ("fmul_d", OpCode.FMUL_D, + [DataType.FLOAT64, DataType.FLOAT64], {}), + ("fdiv_d", OpCode.FDIV_D, + [DataType.FLOAT64, DataType.FLOAT64], {}), + ("fcmp_l_d", OpCode.FCMP_L_D, + [DataType.FLOAT64, DataType.FLOAT64], {}), + ("fcmp_eq_d", OpCode.FCMP_EQ_D, + [DataType.FLOAT64, DataType.FLOAT64], {}), + ("fcvt_s_d", OpCode.FCVT_S_D, [DataType.FLOAT64], {}), + ("fcvt_d_s", OpCode.FCVT_D_S, [DataType.FLOAT32], {}), + ("load_f64", OpCode.LOAD_F64, [DataType.INT32], {}), + ("load_const_f64", OpCode.LOAD_CONST_F64, [], {"value": 1.5}), +] + + +class TestFp64MissingDest: + """F5: fp64 handlers must reject instructions without a dest.""" + + @pytest.mark.parametrize( + "opname,opcode,dtypes,attrs", FP64_MISSING_DEST_CASES, + ids=[case[0] for case in FP64_MISSING_DEST_CASES]) + def test_missing_dest_raises(self, opname, opcode, dtypes, attrs): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + operands = [builder.make_value(dtype=d) for d in dtypes] + builder._emit(opcode, None, operands, **attrs) + with pytest.raises(ValueError, match=opname): + ExtendedInstructionSelector(builder.program).run() + + +class TestIntegerMinMaxSemantics: + """F3/F4: integer MIN/MAX/ABS sequences are semantically correct.""" + + @pytest.mark.parametrize("a,b", [ + (5, 2), (2, 5), (0, 0), (-3, 2), (2, -3), (-7, -2), + ]) + def test_min_registers(self, a, b): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + av = builder.make_value(name="a", dtype=DataType.INT32) + bv = builder.make_value(name="b", dtype=DataType.INT32) + dest = builder.min(av, bv) + + asm = AsmEmitter(ExtendedInstructionSelector( + builder.program).run()).emit() + values = _eval_int_sequence(asm, {av.name: a, bv.name: b}) + assert values[dest.name] == min(a, b) + + @pytest.mark.parametrize("a,b", [ + (5, 2), (2, 5), (0, 0), (-3, 2), (2, -3), (-7, -2), + ]) + def test_max_registers(self, a, b): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + av = builder.make_value(name="a", dtype=DataType.INT32) + bv = builder.make_value(name="b", dtype=DataType.INT32) + dest = builder.max(av, bv) + + asm = AsmEmitter(ExtendedInstructionSelector( + builder.program).run()).emit() + values = _eval_int_sequence(asm, {av.name: a, bv.name: b}) + assert values[dest.name] == max(a, b) + + @pytest.mark.parametrize("a,const", [ + (5, 2), (1, 2), (-3, 2), (7, 2), + ]) + def test_min_immediate_operand(self, a, const): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + av = builder.make_value(name="a", dtype=DataType.INT32) + dest = builder.min(av, builder.make_const(const, dtype=DataType.INT32)) + + asm = AsmEmitter(ExtendedInstructionSelector( + builder.program).run()).emit() + lines = _clean_asm_lines(asm) + # Regression: the constant must not appear as a register operand. + for line in lines: + tokens = line.replace(",", " ").split() + if tokens[0] in ("sub", "and", "add", "slt"): + assert all( + not _is_immediate_token(tok) for tok in tokens[2:] + ), line + assert any(ln.startswith("li ") for ln in lines) + values = _eval_int_sequence(asm, {av.name: a}) + assert values[dest.name] == min(a, const) + + @pytest.mark.parametrize("a,const", [ + (5, 2), (1, 2), (-3, 2), (7, 2), + ]) + def test_max_immediate_operand(self, a, const): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + av = builder.make_value(name="a", dtype=DataType.INT32) + dest = builder.max(av, builder.make_const(const, dtype=DataType.INT32)) + + asm = AsmEmitter(ExtendedInstructionSelector( + builder.program).run()).emit() + lines = _clean_asm_lines(asm) + for line in lines: + tokens = line.replace(",", " ").split() + if tokens[0] in ("sub", "and", "add", "slt"): + assert all( + not _is_immediate_token(tok) for tok in tokens[2:] + ), line + assert any(ln.startswith("li ") for ln in lines) + values = _eval_int_sequence(asm, {av.name: a}) + assert values[dest.name] == max(a, const) + + @pytest.mark.parametrize("a", [-7, -1, 0, 1, 7]) + def test_abs(self, a): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + av = builder.make_value(name="a", dtype=DataType.INT32) + dest = builder.abs(av) + + asm = AsmEmitter(ExtendedInstructionSelector( + builder.program).run()).emit() + values = _eval_int_sequence(asm, {av.name: a}) + assert values[dest.name] == abs(a) + + +class TestMinMaxExecution: + """F3/F4 end-to-end: allocated MIN/MAX/ABS executes correctly.""" + + def setup_method(self): + pytest.importorskip("tinyfive") + + @staticmethod + def _execute(allocated, asm, inputs, result_comment): + result = next(i for i in allocated if i.comment == result_comment) + binary = assemble_to_binary(asm) + words = [ + int.from_bytes(binary[i:i + 4], "little") + for i in range(0, len(binary), 4) + ] + machine = ProfiledMachine(mem_size=4096) + machine.load_binary(words, origin=0) + for name, value in inputs.items(): + machine.set_reg(REG_MAP[name], value) + machine.run(instructions=len(words), start=0, strict=True) + assert machine.last_error is None + return machine.get_reg(REG_MAP[result.dst.value]) + + def test_min_immediate_executes(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + av = builder.make_value(name="a", dtype=DataType.INT32) + builder.min(av, builder.make_const(2, dtype=DataType.INT32)) + allocated, asm = _allocated_pipeline(builder) + + slt = next(i for i in allocated if i.comment == "min: slt") + for value, expected in [(5, 2), (1, 1), (2, 2), (-4, -4)]: + got = self._execute( + allocated, asm, {slt.src1.value: value}, "min result") + assert got == expected + + def test_max_registers_executes(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + av = builder.make_value(name="a", dtype=DataType.INT32) + bv = builder.make_value(name="b", dtype=DataType.INT32) + builder.max(av, bv) + allocated, asm = _allocated_pipeline(builder) + + slt = next(i for i in allocated if i.comment == "max: slt") + for a, b in [(3, 7), (7, 3), (-2, 5), (5, -2)]: + got = self._execute( + allocated, asm, + {slt.src1.value: a, slt.src2.value: b}, "max result") + assert got == max(a, b) + + def test_abs_executes(self): + builder = IRBuilder() + builder.new_function("test") + builder.new_block("entry") + av = builder.make_value(name="a", dtype=DataType.INT32) + builder.abs(av) + allocated, asm = _allocated_pipeline(builder) + + srai = next( + i for i in allocated if i.comment == "abs: srai 31") + for value in (-7, -1, 0, 5): + got = self._execute( + allocated, asm, {srai.src1.value: value}, "abs: sub") + assert got == abs(value) + + if __name__ == "__main__": pytest.main([__file__, "-v"]) diff --git a/tests/test_riscv_encoder_fd.py b/tests/test_riscv_encoder_fd.py index 1c8aea5..0e70f24 100644 --- a/tests/test_riscv_encoder_fd.py +++ b/tests/test_riscv_encoder_fd.py @@ -29,6 +29,14 @@ "fcvt.d.s f0, f1", "fmv.x.w a0, f1", "li.d f0, 123", + # Fused / classified / Zfa families (Topic 28 review F6). + "fmadd.d f0, f1, f2, f3", + "fnmadd.d f0, f1, f2, f3", + "fmsub.s f0, f1, f2, f3", + "fnmsub.s f0, f1, f2, f3", + "fclass.s a0, f1", + "fli.s f0, 1", + "fround.s f0, f1", ] @@ -62,6 +70,13 @@ def test_unknown_non_fd_still_value_error(): ("fsqrt.s", True), ("fmv.x.w", True), ("fsgnjx.d", True), + ("fmadd.d", True), + ("fnmadd.d", True), + ("fmsub.s", True), + ("fnmsub.s", True), + ("fclass.s", True), + ("fli.s", True), + ("fround.s", True), ]) def test_is_fd_mnemonic_predicate(mnemonic, expected): assert _is_fd_mnemonic(mnemonic) is expected From 69b3af9fbbf094584aadc69e893467ee6b070d06 Mon Sep 17 00:00:00 2001 From: opencode Date: Tue, 15 Sep 2026 00:59:17 +0800 Subject: [PATCH 6/6] feat(topic28): add extended-isel feature case report and CI regressions --- .github/workflows/ci.yml | 20 + .../cases/topic28_extended_isel_feature.dsl | 20 + benchmarks/run_topic28_extended_isel_case.py | 729 ++++++++++++++++++ .../test_topic28_extended_isel_case_report.py | 165 ++++ 4 files changed, 934 insertions(+) create mode 100644 benchmarks/cases/topic28_extended_isel_feature.dsl create mode 100644 benchmarks/run_topic28_extended_isel_case.py create mode 100644 tests/test_topic28_extended_isel_case_report.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 15aae4b..27d9260 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -107,6 +107,15 @@ jobs: run: | python3.12 -m pytest tests/test_pr37_regression.py -v --tb=short + - name: Run topic28 extended-isel regressions + run: | + python3.12 -m pytest \ + tests/test_extended_isel_cli.py \ + tests/test_inst_select_ext.py \ + tests/test_riscv_encoder_fd.py \ + tests/test_topic28_extended_isel_case_report.py \ + -v --tb=short + - name: Generate test visualization page if: github.ref == 'refs/heads/main' run: | @@ -218,6 +227,14 @@ jobs: --json benchmark_reports/const_merge_report.json \ --markdown benchmark_reports/const_merge_report.md + # ── 3.1.3 课题28:扩展指令选择 case 报告(A/B + 编码器门禁) ─────── + - name: Topic 28 extended instruction-selection case report + run: | + mkdir -p benchmark_reports + python3.12 benchmarks/run_topic28_extended_isel_case.py \ + --json benchmark_reports/extended_isel_report.json \ + --markdown benchmark_reports/extended_isel_report.md + # ── 3.2 DSL 用例编译 + 模拟基准 ──────────────────────────────────── - name: DSL case compilation benchmarks run: | @@ -363,6 +380,9 @@ jobs: if [ -f benchmark_reports/const_merge_report.md ]; then cat benchmark_reports/const_merge_report.md >> $GITHUB_STEP_SUMMARY fi + if [ -f benchmark_reports/extended_isel_report.md ]; then + cat benchmark_reports/extended_isel_report.md >> $GITHUB_STEP_SUMMARY + fi echo "" >> $GITHUB_STEP_SUMMARY if [ -f benchmark_reports/github_summary.md ]; then cat benchmark_reports/github_summary.md >> $GITHUB_STEP_SUMMARY diff --git a/benchmarks/cases/topic28_extended_isel_feature.dsl b/benchmarks/cases/topic28_extended_isel_feature.dsl new file mode 100644 index 0000000..6c60ee1 --- /dev/null +++ b/benchmarks/cases/topic28_extended_isel_feature.dsl @@ -0,0 +1,20 @@ +# Topic 28 extended instruction-selection feature case (deterministic). +# +# The stock DSL frontend (scratchv/frontend/dsl_parser.py) exposes a closed op +# table: add/sub/mul/div/neg/exp/relu/gelu/dot/matmul/softmax/maxpool. The +# Topic 28 opcodes (sqrt/min/max/abs/idiv/rem/mod and the float64 family) have +# no DSL syntax, so this file drives the CompilerDriver A/B wiring check while +# the extended-only instruction shapes are probed at IR level by +# benchmarks/run_topic28_extended_isel_case.py (see its honesty note). +# +# Loop trip count 0..3 leaves acc = 6, sq = 36, bias = relu(36) = 36. +# neg_acc = -6; total = 36 - (-6) = 42; res = 42 / 6 = 7. +for i = 0, 4 + acc = add(i, i) + sq = mul(acc, acc) + bias = relu(sq) +endfor +neg_acc = neg(acc) +total = sub(bias, neg_acc) +res = div(total, acc) +return res diff --git a/benchmarks/run_topic28_extended_isel_case.py b/benchmarks/run_topic28_extended_isel_case.py new file mode 100644 index 0000000..17d3ba8 --- /dev/null +++ b/benchmarks/run_topic28_extended_isel_case.py @@ -0,0 +1,729 @@ +#!/usr/bin/env python3 +"""Run the Topic 28 extended instruction-selection feature case. + +The report proves four separate facts, all deterministic and auditable: + +1. the ``extended_isel`` opt-in flag is wired through ``CompilerDriver``: the + feature DSL case compiles in both configurations, the assembly is + reproducible, and the RV32 emulator executes both products to the same + architectural state and the expected result; +2. the extended-only opcodes (sqrt / min / max / abs / idiv / rem and the + float64 family) select FP mnemonics under ``extended_isel=True``; raising + the encoder gate, ``assemble_to_binary`` rejects F/D assembly with + ``UnsupportedInstructionError`` instead of silently mis-encoding it, while + the FP-mnemonic-free extended assembly is accepted by the RV32IM encoder; +3. the counterexample matrix is recorded: ``--no-fp64`` on a float64 program + and the base selector on extended opcodes fail loud, and the LLVM backend / + DAG-selection combinations warn instead of silently ignoring the flag; +4. no execution equivalence is claimed for the FP probe: the RV32IM encoder + rejects every F/D mnemonic and ``RV32Emulator`` retires RV32IM only, so the + FP execution entry is explicitly ``skipped``. + +This is a deterministic feature/integration case, not a real-workload speedup +claim. Real workloads remain covered by ``run_benchmark.py``. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import tempfile +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Callable + +from scratchv.backend._asm_parser import parse_asm +from scratchv.backend.riscv_encoder import ( + UnsupportedInstructionError, + _is_fd_mnemonic, + assemble_to_binary, +) +from scratchv.compiler import CompileResult, CompilerConfig, CompilerDriver +from scratchv.ir.builder import IRBuilder +from scratchv.ir.types import DataType, OpCode, Program +from scratchv.simulator.rv32_emulator import REG_ID, RV32Emulator + +SCHEMA_VERSION = "topic28-extended-isel-case/1" +DEFAULT_CASE = ( + Path(__file__).parent / "cases" / "topic28_extended_isel_feature.dsl" +) +DEFAULT_JSON = Path("benchmark_reports/extended_isel_report.json") +DEFAULT_MARKDOWN = Path("benchmark_reports/extended_isel_report.md") +#: a0 after the deterministic DSL case: acc = 6, bias = 36, neg_acc = -6, +#: total = 36 - (-6) = 42, res = 42 / 6 = 7. +EXPECTED_DSL_RESULT = 7 +#: FP mnemonics the extended-only IR probe must emit. The branch lowers +#: float64 arithmetic to ``fadd.d``/``fmul.d`` (``fadd.s`` has no producer). +REQUIRED_FP_MNEMONICS = ("fadd.d", "fmul.d") +#: Text-level evidence uses the greedy allocator path (see the honesty note). +REG_ALLOC = "greedy" +OPTIMIZE_LEVEL = "all" + +HONESTY = ( + "Deterministic feature case, not a workload speedup claim. The stock DSL " + "frontend exposes a closed op table (add/sub/mul/div/neg/exp/relu/gelu/" + "dot/matmul/softmax/maxpool), so the Topic 28-only opcodes are exercised " + "through IR-level probes built with IRBuilder; the A/B DSL comparison " + "therefore proves opt-in wiring and execution equivalence on the shared " + "base path, not a lowering difference (the DSL produces byte-identical " + "assembly in both configurations, which is expected). Only shapes " + "covered by tests/test_inst_select_ext.py are used: f32 sqrt literals go " + "through the exact bit-pattern materialization (the F1 regression fix), " + "float64 literal operands never reach arithmetic (load_const_f64 " + "materializes exact IEEE-754 bits first, and f64 literals fail loud by " + "design), and integer min/max/abs literals are materialized into " + "registers before the register-only branchless sequences. The RV32IM " + "encoder rejects every F/D mnemonic by design (fail-loud, final encoding " + "out of scope), so no FP assembly can be accepted by assemble_to_binary; " + "the gate checked here is explicit rejection plus successful encoding of " + "the FP-mnemonic-free extended assembly. FP execution is skipped: " + "RV32Emulator retires RV32IM only and the F/D product cannot be encoded, " + "so no execution-equivalence claim is made. The extended integer " + "min/max/abs sequences use SLT/SRAI/REM, which the minimal emulator does " + "not implement, so they are validated by encoding only. Execution " + "evidence uses the greedy allocator path because the LinearScanAllocator " + "emits branch targets as trailing comments that assemble_to_binary drops " + "(pre-existing integration gap, orthogonal to Topic 28)." +) + + +# ═══════════════════════════════════════════════════════════════════════ +# IR case builders (Topic 28 extended shapes) +# ═══════════════════════════════════════════════════════════════════════ + +def _runtime_const(builder: IRBuilder, name: str, value: int): + """LOAD_CONST value kept as a runtime register, not a folded literal.""" + value_obj = builder.make_value( + name=name, dtype=DataType.INT32, is_constant=False) + builder._emit(OpCode.LOAD_CONST, value_obj, value=value) + return value_obj + + +def build_fp_feature_program() -> Program: + """FP shapes verified by the branch tests (see the honesty note).""" + builder = IRBuilder() + builder.new_function("fp_feature") + builder.new_block("entry") + x = builder.make_value(name="x", dtype=DataType.FLOAT32) + i = builder.make_value(name="i", dtype=DataType.INT32) + j = builder.make_value(name="j", dtype=DataType.INT32) + # Software sqrt: register operand and exact-bit-pattern f32 literal. + sqrt_reg = builder.sqrt(x) + builder.sqrt(builder.make_const(2.5, dtype=DataType.FLOAT32)) + # f64 arithmetic only on values materialized by load_const_f64. + acc = builder.fadd_d( + builder.load_const_f64(1.5), builder.load_const_f64(2.0)) + builder.fmul_d(acc, builder.load_const_f64(0.5)) + # Integer extended sequences with literal materialization. + lo = builder.min(i, builder.make_const(3, dtype=DataType.INT32)) + builder.max(j, builder.make_const(2, dtype=DataType.INT32)) + builder.abs(lo) + builder.ret(sqrt_reg) + return builder.program + + +def build_integer_extended_program() -> Program: + """Extended integer program without F/D mnemonics (encoder-acceptable).""" + builder = IRBuilder() + builder.new_function("integer_extended") + builder.new_block("entry") + five = _runtime_const(builder, "five", 5) + seventy = _runtime_const(builder, "seventy", 70) + lo = builder.min(five, builder.make_const(3, dtype=DataType.INT32)) + hi = builder.max(lo, seventy) + mag = builder.abs(builder.make_const(-9, dtype=DataType.INT32)) + quotient = builder.idiv(hi, lo) + remainder = builder.rem(hi, lo) + builder.ret(builder.add(builder.add(quotient, remainder), mag)) + return builder.program + + +def build_hardware_sqrt_program() -> Program: + """f32 sqrt literal for the ``--hardware-sqrt`` branch (emits fsqrt.s).""" + builder = IRBuilder() + builder.new_function("hardware_sqrt") + builder.new_block("entry") + builder.ret(builder.sqrt(builder.make_const(2.0, dtype=DataType.FLOAT32))) + return builder.program + + +def build_fp64_only_program() -> Program: + """Pure float64 program for the ``--no-fp64`` counterexample.""" + builder = IRBuilder() + builder.new_function("fp64_only") + builder.new_block("entry") + total = builder.fadd_d( + builder.load_const_f64(1.5), builder.load_const_f64(2.0)) + builder.ret(total) + return builder.program + + +# ═══════════════════════════════════════════════════════════════════════ +# Measurement helpers +# ═══════════════════════════════════════════════════════════════════════ + +def count_asm(asm: str) -> int: + """Count real (non-directive) assembly instructions.""" + return sum( + line.opcode is not None and not line.is_directive + for line in parse_asm(asm) + ) + + +def fp_mnemonics(asm: str) -> list[str]: + """Sorted F/D mnemonics present in *asm* (encoder predicate reused).""" + return sorted({ + line.opcode + for line in parse_asm(asm) + if line.opcode is not None + and not line.is_directive + and _is_fd_mnemonic(line.opcode) + }) + + +def _sha256(text: str) -> str: + return hashlib.sha256(text.encode("utf-8")).hexdigest() + + +def encoder_gate(asm: str) -> dict[str, Any]: + """Run ``assemble_to_binary`` and categorize the outcome.""" + try: + binary = assemble_to_binary(asm) + except UnsupportedInstructionError as exc: + return { + "encoded": False, + "bytes": 0, + "words": 0, + "error_type": "UnsupportedInstructionError", + "error": str(exc), + } + except ValueError as exc: + return { + "encoded": False, + "bytes": 0, + "words": 0, + "error_type": type(exc).__name__, + "error": str(exc), + } + return { + "encoded": True, + "bytes": len(binary), + "words": len(binary) // 4, + "error_type": None, + "error": None, + } + + +def compile_ir_program( + program: Program, *, + extended_isel: bool = True, + enable_fp64: bool = True, + use_hardware_sqrt: bool = False, +) -> dict[str, Any]: + """Compile one IR program through the driver's codegen entry. + + ``CompilerDriver.compile`` only accepts DSL/ONNX sources and the DSL + frontend cannot express Topic 28-only opcodes, so the probe calls the + driver's own RISC-V codegen entry (the same one ``compile`` uses). + """ + driver = CompilerDriver(CompilerConfig( + extended_isel=extended_isel, + enable_fp64=enable_fp64, + use_hardware_sqrt=use_hardware_sqrt, + reg_alloc=REG_ALLOC, + optimize_level=OPTIMIZE_LEVEL, + )) + asm = driver._generate_code(program) + return { + "success": True, + "asm": asm, + "asm_instructions": count_asm(asm), + "fp_mnemonics": fp_mnemonics(asm), + "asm_sha256": _sha256(asm), + "encoder": encoder_gate(asm), + } + + +def _capture_failure(action: Callable[[], Any]) -> dict[str, Any]: + """Run *action*; record an explicit rejection instead of propagating.""" + try: + result = action() + except Exception as exc: # noqa: BLE001 - the rejection is the evidence + return { + "rejected": True, + "error_type": type(exc).__name__, + "error": str(exc), + } + return { + "rejected": False, + "error_type": None, + "error": None, + "unexpected_result": result, + } + + +def run_asm(asm: str) -> dict[str, Any]: + """Assemble and execute *asm*; return register state and counters.""" + binary = assemble_to_binary(asm) + emulator = RV32Emulator() + emulator.load_code(bytes(binary)) + dynamic = emulator.run() + return { + "backend": "rv32-emulator", + "a0": emulator.regs[REG_ID["a0"]], + "registers": {f"x{i}": emulator.regs[i] for i in range(32)}, + "dynamic_instructions": dynamic, + "encoded_words": len(binary) // 4, + } + + +# ═══════════════════════════════════════════════════════════════════════ +# Measurements +# ═══════════════════════════════════════════════════════════════════════ + +def _compile_dsl(case_path: Path, **config_overrides) -> CompileResult: + """Compile the feature DSL case with one driver configuration.""" + source = case_path.read_text(encoding="utf-8") + overrides = { + "reg_alloc": REG_ALLOC, + "optimize_level": OPTIMIZE_LEVEL, + **config_overrides, + } + driver = CompilerDriver(CompilerConfig(**overrides)) + with tempfile.TemporaryDirectory() as tmp: + output = str(Path(tmp) / "case.s") + return driver.compile("", output, dsl_source=source) + + +def _dsl_side(case_path: Path, *, extended_isel: bool) -> dict[str, Any]: + """Compile the case twice and execute the selected product.""" + first = _compile_dsl(case_path, extended_isel=extended_isel) + second = _compile_dsl(case_path, extended_isel=extended_isel) + deterministic = ( + first.success and second.success + and first.output_text == second.output_text + ) + side: dict[str, Any] = { + "extended_isel": extended_isel, + "success": first.success, + "errors": list(first.errors), + "warnings": list(first.warnings), + "asm_instructions": ( + count_asm(first.output_text) if first.success else None), + "fp_mnemonics": ( + fp_mnemonics(first.output_text) if first.success else None), + "asm_sha256": ( + _sha256(first.output_text) if first.success else None), + "asm": first.output_text if first.success else "", + "deterministic": deterministic, + "execution": None, + } + if first.success and second.success: + try: + side["execution"] = run_asm(first.output_text) + except Exception as exc: # noqa: BLE001 - recorded for the hard check + side["execution"] = { + "error": f"{type(exc).__name__}: {exc}", + } + return side + + +def measure_dsl_ab(case_path: Path) -> dict[str, Any]: + """CompilerDriver A/B on the feature DSL case.""" + off = _dsl_side(case_path, extended_isel=False) + on = _dsl_side(case_path, extended_isel=True) + return { + "off": off, + "on": on, + "identical_asm": ( + off["success"] and on["success"] + and off["asm"] == on["asm"] + ), + } + + +def measure_extended_probe() -> dict[str, Any]: + """IR-level evidence for the extended-only opcodes.""" + fp_first = compile_ir_program(build_fp_feature_program()) + fp_second = compile_ir_program(build_fp_feature_program()) + integer = compile_ir_program(build_integer_extended_program()) + hardware = compile_ir_program( + build_hardware_sqrt_program(), use_hardware_sqrt=True) + base_selector = _capture_failure(lambda: compile_ir_program( + build_fp_feature_program(), extended_isel=False)) + return { + "fp": { + **fp_first, + "deterministic": fp_first["asm"] == fp_second["asm"], + }, + "integer": integer, + "hardware_sqrt": hardware, + "base_selector": base_selector, + } + + +def measure_error_matrix(case_path: Path) -> dict[str, Any]: + """Record fail-loud and warning behaviour for the flag combinations.""" + source = case_path.read_text(encoding="utf-8") + + def compile_case(**overrides) -> CompileResult: + config = { + "extended_isel": True, + "reg_alloc": REG_ALLOC, + "optimize_level": OPTIMIZE_LEVEL, + **overrides, + } + driver = CompilerDriver(CompilerConfig(**config)) + with tempfile.TemporaryDirectory() as tmp: + suffix = "ll" if config.get("backend") == "llvm" else "s" + return driver.compile( + "", str(Path(tmp) / f"case.{suffix}"), dsl_source=source) + + off = _capture_failure(lambda: compile_ir_program( + build_fp_feature_program(), extended_isel=False)) + disabled = _capture_failure(lambda: compile_ir_program( + build_fp64_only_program(), extended_isel=True, enable_fp64=False)) + llvm = compile_case(backend="llvm") + dag = compile_case(use_dag_isel=True) + no_ext_fp64 = compile_case(extended_isel=False, enable_fp64=False) + no_ext_hw = compile_case(extended_isel=False, use_hardware_sqrt=True) + + def warns(result: CompileResult, needle: str) -> bool: + return result.success and any(needle in w for w in result.warnings) + + rows = [ + { + "id": "extended_isel_off", + "config": "extended_isel=False on extended-only opcodes", + "expected": "explicit rejection (opt-in feature)", + "observed": off["error"] or "no error", + "status": ( + "error" + if off["rejected"] and off["error_type"] == "ValueError" + else "unexpected" + ), + }, + { + "id": "extended_isel_no_fp64", + "config": "extended_isel=True, enable_fp64=False on float64 IR", + "expected": "explicit rejection mentioning enable_fp64", + "observed": disabled["error"] or "no error", + "status": ( + "error" + if disabled["rejected"] + and "enable_fp64" in (disabled["error"] or "") + else "unexpected" + ), + }, + { + "id": "llvm_backend", + "config": "extended_isel=True, backend=llvm", + "expected": "warning (RISC-V only)", + "observed": ( + "; ".join(llvm.warnings) if llvm.warnings + else "no warning" + ), + "status": ( + "warning" if warns(llvm, "RISC-V only") else "unexpected" + ), + }, + { + "id": "dag_isel_precedence", + "config": "extended_isel=True, use_dag_isel=True", + "expected": "warning (dag-isel takes precedence)", + "observed": ( + "; ".join(dag.warnings) if dag.warnings else "no warning" + ), + "status": "warning" if warns(dag, "precedence") else "unexpected", + }, + { + "id": "fp64_flag_without_extended", + "config": "extended_isel=False, enable_fp64=False", + "expected": "warning (no effect without --extended-isel)", + "observed": ( + "; ".join(no_ext_fp64.warnings) if no_ext_fp64.warnings + else "no warning" + ), + "status": ( + "warning" if warns(no_ext_fp64, "no effect") + else "unexpected" + ), + }, + { + "id": "hardware_sqrt_without_extended", + "config": "extended_isel=False, use_hardware_sqrt=True", + "expected": "warning (no effect without --extended-isel)", + "observed": ( + "; ".join(no_ext_hw.warnings) if no_ext_hw.warnings + else "no warning" + ), + "status": ( + "warning" if warns(no_ext_hw, "no effect") + else "unexpected" + ), + }, + ] + return { + "rows": rows, + "off": off, + "no_fp64": disabled, + "llvm_success": llvm.success, + "dag_success": dag.success, + } + + +# ═══════════════════════════════════════════════════════════════════════ +# Report assembly +# ═══════════════════════════════════════════════════════════════════════ + +def _execution_a0(side: dict[str, Any]) -> Any: + execution = side.get("execution") or {} + return execution.get("a0") + + +def evaluate(case_path: Path) -> dict[str, Any]: + """Build the full report payload and run the hard invariants.""" + dsl_ab = measure_dsl_ab(case_path) + probe = measure_extended_probe() + matrix = measure_error_matrix(case_path) + + off, on = dsl_ab["off"], dsl_ab["on"] + fp = probe["fp"] + integer, hardware = probe["integer"], probe["hardware_sqrt"] + base_selector = probe["base_selector"] + no_fp64 = matrix["no_fp64"] + + off_exec = off.get("execution") or {} + on_exec = on.get("execution") or {} + + hard_checks = { + "dsl_ab_both_compile": off["success"] and on["success"], + "dsl_ab_asm_deterministic": ( + off["deterministic"] and on["deterministic"]), + "dsl_execution_matches_expected": ( + _execution_a0(off) == EXPECTED_DSL_RESULT + and _execution_a0(on) == EXPECTED_DSL_RESULT + and off_exec.get("dynamic_instructions", 0) > 0 + and on_exec.get("dynamic_instructions", 0) > 0 + ), + "dsl_execution_registers_identical": ( + bool(off_exec.get("registers")) + and off_exec.get("registers") == on_exec.get("registers") + ), + "extended_probe_has_fp_mnemonics": ( + set(REQUIRED_FP_MNEMONICS).issubset(fp["fp_mnemonics"])), + "extended_probe_asm_deterministic": fp["deterministic"], + "fp_asm_rejected_by_encoder": ( + not fp["encoder"]["encoded"] + and fp["encoder"]["error_type"] == "UnsupportedInstructionError" + ), + "hardware_sqrt_emits_fsqrt_s": "fsqrt.s" in hardware["fp_mnemonics"], + "hardware_sqrt_rejected_by_encoder": ( + not hardware["encoder"]["encoded"] + and hardware["encoder"]["error_type"] + == "UnsupportedInstructionError" + ), + "integer_extended_asm_is_encodable": ( + integer["encoder"]["encoded"] + and integer["encoder"]["words"] > 0 + and not integer["fp_mnemonics"] + ), + "base_selector_rejects_extended_ops": ( + base_selector["rejected"] + and base_selector["error_type"] == "ValueError" + ), + "fp64_disabled_raises_explicitly": ( + no_fp64["rejected"] + and "enable_fp64" in (no_fp64["error"] or "") + ), + "llvm_backend_warns": ( + matrix["llvm_success"] + and any(row["id"] == "llvm_backend" + and row["status"] == "warning" + for row in matrix["rows"]) + ), + "dag_isel_precedence_warns": ( + matrix["dag_success"] + and any(row["id"] == "dag_isel_precedence" + and row["status"] == "warning" + for row in matrix["rows"]) + ), + "fp64_flags_without_extended_warn": ( + any(row["id"] in ("fp64_flag_without_extended", + "hardware_sqrt_without_extended") + and row["status"] == "warning" + for row in matrix["rows"]) + ), + } + failed = sorted(name for name, ok in hard_checks.items() if not ok) + + return { + "schema_version": SCHEMA_VERSION, + "topic": "topic28-extended-isel", + "generated_at": datetime.now(timezone.utc).isoformat(), + "case": str(case_path), + "config": { + "reg_alloc": REG_ALLOC, + "optimize_level": OPTIMIZE_LEVEL, + }, + "expected_dsl_result": EXPECTED_DSL_RESULT, + "required_fp_mnemonics": list(REQUIRED_FP_MNEMONICS), + "dsl_ab": dsl_ab, + "extended_probe": probe, + "behavior_matrix": matrix["rows"], + "execution": { + "dsl_ab": { + "status": "ok", + "backend": "rv32-emulator", + "expected_a0": EXPECTED_DSL_RESULT, + "off": off.get("execution"), + "on": on.get("execution"), + }, + "fp_probe": { + "status": "skipped", + "reason": ( + "RV32Emulator retires RV32IM only and the RV32IM " + "encoder rejects every F/D mnemonic (fail-loud), so no " + "F/D binary exists to execute; no execution-equivalence " + "claim is made." + ), + }, + }, + "hard_checks": hard_checks, + "hard_failures": failed, + "honesty": HONESTY, + } + + +def render_markdown(report: dict[str, Any]) -> str: + dsl_ab = report["dsl_ab"] + probe = report["extended_probe"] + off, on = dsl_ab["off"], dsl_ab["on"] + fp, integer = probe["fp"], probe["integer"] + hardware = probe["hardware_sqrt"] + off_exec = off.get("execution") or {} + on_exec = on.get("execution") or {} + + def mnemonics(items: list[str] | None) -> str: + return ", ".join(items) if items else "(none)" + + def gate(cell: dict[str, Any]) -> str: + if cell["encoded"]: + return f"accepted ({cell['words']} words)" + return f"rejected: {cell['error_type']}" + + hard_total = len(report["hard_checks"]) + hard_ok = hard_total - len(report["hard_failures"]) + lines = [ + "# Topic 28 Extended Instruction-Selection Feature Case", + "", + f"- Schema: `{report['schema_version']}`", + f"- Case: `{report['case']}`", + f"- Generated: {report['generated_at']}", + f"- Expected `a0` (DSL case): {report['expected_dsl_result']}", + f"- Hard checks: " + f"{'PASS' if not report['hard_failures'] else 'FAIL'} " + f"({hard_ok}/{hard_total})", + "", + "## DSL A/B (CompilerDriver, feature case)", + "", + "| Metric | extended_isel=False | extended_isel=True |", + "|--------|--------------------:|-------------------:|", + f"| compiled | {'yes' if off['success'] else 'no'} | " + f"{'yes' if on['success'] else 'no'} |", + f"| ASM instructions | {off['asm_instructions']} | " + f"{on['asm_instructions']} |", + f"| FP mnemonics in ASM | {mnemonics(off['fp_mnemonics'])} | " + f"{mnemonics(on['fp_mnemonics'])} |", + f"| two compiles identical | {off['deterministic']} | " + f"{on['deterministic']} |", + f"| `a0` (RV32 emulator) | {off_exec.get('a0')} | " + f"{on_exec.get('a0')} |", + f"| dynamic instructions | {off_exec.get('dynamic_instructions')} | " + f"{on_exec.get('dynamic_instructions')} |", + f"| ASM identical across configs | " + f"{'yes' if dsl_ab['identical_asm'] else 'no'} | - |", + "", + "The DSL frontend exposes a closed op table, so the A/B case covers " + "the opt-in wiring and execution path; the extended-only opcodes are " + "probed below at IR level.", + "", + "## Extended-only IR probe (FP mnemonics and encoder gate)", + "", + "| Probe | Config | FP mnemonics | `assemble_to_binary` |", + "|-------|--------|--------------|----------------------|", + f"| fp_feature (sqrt/min/max/abs/f64) | extended_isel=True | " + f"{mnemonics(fp['fp_mnemonics'])} | {gate(fp['encoder'])} |", + f"| hardware_sqrt | extended_isel=True, " + f"use_hardware_sqrt=True | {mnemonics(hardware['fp_mnemonics'])} | " + f"{gate(hardware['encoder'])} |", + f"| integer_extended (no F/D mnemonic) | extended_isel=True | " + f"{mnemonics(integer['fp_mnemonics'])} | {gate(integer['encoder'])} |", + f"| fp_feature | extended_isel=False | - | " + f"rejected: {probe['base_selector']['error_type']} |", + "", + "## Failure / degradation matrix", + "", + "| Config | Expected | Observed | Status |", + "|--------|----------|----------|--------|", + ] + for row in report["behavior_matrix"]: + observed = row["observed"].replace("\n", " ") + if len(observed) > 90: + observed = observed[:87] + "..." + lines.append( + f"| `{row['config']}` | {row['expected']} | {observed} | " + f"{row['status']} |" + ) + + lines += [ + "", + "## Execution", + "", + f"- DSL A/B: status `{report['execution']['dsl_ab']['status']}`, " + f"backend `{report['execution']['dsl_ab']['backend']}`, expected " + f"`a0 == {report['execution']['dsl_ab']['expected_a0']}`.", + f"- FP probe: status `{report['execution']['fp_probe']['status']}` — " + f"{report['execution']['fp_probe']['reason']}", + "", + "## Hard checks", + "", + ] + for name, ok in report["hard_checks"].items(): + lines.append(f"- [{'x' if ok else ' '}] {name}") + lines += [ + "", + "## Honesty", + "", + report["honesty"], + "", + ] + return "\n".join(lines) + + +def main(argv: list[str] | None = None) -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--case", type=Path, default=DEFAULT_CASE) + parser.add_argument("--json", type=Path, default=DEFAULT_JSON) + parser.add_argument("--markdown", type=Path, default=DEFAULT_MARKDOWN) + args = parser.parse_args(argv) + if not args.case.is_file(): + parser.error(f"feature case not found: {args.case}") + + report = evaluate(args.case) + args.json.parent.mkdir(parents=True, exist_ok=True) + args.json.write_text(json.dumps(report, indent=2) + "\n", encoding="utf-8") + args.markdown.parent.mkdir(parents=True, exist_ok=True) + args.markdown.write_text( + render_markdown(report) + "\n", encoding="utf-8") + print(render_markdown(report)) + if report["hard_failures"]: + print("HARD FAILURES: " + ", ".join(report["hard_failures"])) + return 1 + print(f"reports written: {args.json}, {args.markdown}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/test_topic28_extended_isel_case_report.py b/tests/test_topic28_extended_isel_case_report.py new file mode 100644 index 0000000..2aa850e --- /dev/null +++ b/tests/test_topic28_extended_isel_case_report.py @@ -0,0 +1,165 @@ +"""Tests for the Topic 28 extended instruction-selection feature case report. + +The report is the CI artifact that proves the ``--extended-isel`` opt-in is +wired through ``CompilerDriver``, that the extended-only opcodes select FP +mnemonics under ``extended_isel=True``, that the RV32IM encoder gate fails +loud on F/D assembly while accepting FP-mnemonic-free extended assembly, and +that the counterexample flag combinations are recorded explicitly. +""" + +from __future__ import annotations + +import json +from pathlib import Path + +import pytest + +from benchmarks.run_topic28_extended_isel_case import ( + EXPECTED_DSL_RESULT, + REQUIRED_FP_MNEMONICS, + SCHEMA_VERSION, + build_fp_feature_program, + build_integer_extended_program, + compile_ir_program, + evaluate, + main, + measure_dsl_ab, +) + +CASE = ( + Path(__file__).resolve().parents[1] + / "benchmarks" / "cases" / "topic28_extended_isel_feature.dsl" +) + + +def test_dsl_case_compiles_in_both_configurations(): + """Both A/B configurations compile and execute to the expected result.""" + ab = measure_dsl_ab(CASE) + assert ab["off"]["success"], ab["off"]["errors"] + assert ab["on"]["success"], ab["on"]["errors"] + assert ab["off"]["asm_instructions"] > 0 + assert ab["on"]["asm_instructions"] > 0 + assert ab["off"]["fp_mnemonics"] == [] + assert ab["on"]["fp_mnemonics"] == [] + assert ab["off"]["execution"]["a0"] == EXPECTED_DSL_RESULT + assert ab["on"]["execution"]["a0"] == EXPECTED_DSL_RESULT + assert ab["off"]["execution"]["dynamic_instructions"] > 0 + assert ( + ab["off"]["execution"]["registers"] + == ab["on"]["execution"]["registers"] + ) + + +def test_extended_probe_reports_fp_mnemonics(): + """Extended-only IR shapes must select the expected FP mnemonics.""" + probe = compile_ir_program(build_fp_feature_program()) + for mnemonic in REQUIRED_FP_MNEMONICS: + assert mnemonic in probe["fp_mnemonics"] + + +def test_encoder_gate_rejects_fp_and_accepts_integer_extended(): + """F/D assembly is rejected fail-loud; F/D-free extended asm encodes.""" + fp = compile_ir_program(build_fp_feature_program()) + assert fp["encoder"]["encoded"] is False + assert fp["encoder"]["error_type"] == "UnsupportedInstructionError" + + integer = compile_ir_program(build_integer_extended_program()) + assert integer["fp_mnemonics"] == [] + assert integer["encoder"]["encoded"] is True + assert integer["encoder"]["words"] > 0 + + +def test_evaluate_passes_hard_checks_and_records_matrix(): + report = evaluate(CASE) + assert report["schema_version"] == SCHEMA_VERSION + assert report["hard_failures"] == [] + assert all(report["hard_checks"].values()) + assert report["honesty"] + + rows = {row["id"]: row for row in report["behavior_matrix"]} + assert rows["extended_isel_off"]["status"] == "error" + assert rows["extended_isel_no_fp64"]["status"] == "error" + assert "enable_fp64" in rows["extended_isel_no_fp64"]["observed"] + assert rows["llvm_backend"]["status"] == "warning" + assert rows["dag_isel_precedence"]["status"] == "warning" + assert rows["fp64_flag_without_extended"]["status"] == "warning" + assert rows["hardware_sqrt_without_extended"]["status"] == "warning" + + fp_execution = report["execution"]["fp_probe"] + assert fp_execution["status"] == "skipped" + assert fp_execution["reason"] + + +def test_measurements_are_deterministic(): + first_ab = measure_dsl_ab(CASE) + second_ab = measure_dsl_ab(CASE) + assert first_ab["off"]["deterministic"] and first_ab["on"]["deterministic"] + assert ( + first_ab["off"]["asm_sha256"] == second_ab["off"]["asm_sha256"] + ) + assert ( + first_ab["on"]["asm_sha256"] == second_ab["on"]["asm_sha256"] + ) + assert first_ab["identical_asm"] is True + + first = compile_ir_program(build_fp_feature_program()) + second = compile_ir_program(build_fp_feature_program()) + assert first["asm_sha256"] == second["asm_sha256"] + + +def test_hard_check_gate_is_not_vacuous(monkeypatch): + """A degenerate extended probe must be reported as hard failures.""" + degenerate_asm = " add a0, x1, x2\n" + + def fake_compile_ir_program(_program, **_kwargs): + return { + "success": True, + "asm": degenerate_asm, + "asm_instructions": 1, + "fp_mnemonics": [], + "asm_sha256": "0" * 64, + "encoder": { + "encoded": True, + "bytes": 4, + "words": 1, + "error_type": None, + "error": None, + }, + } + + monkeypatch.setattr( + "benchmarks.run_topic28_extended_isel_case.compile_ir_program", + fake_compile_ir_program, + ) + report = evaluate(CASE) + assert report["hard_failures"] + assert "extended_probe_has_fp_mnemonics" in report["hard_failures"] + assert "fp_asm_rejected_by_encoder" in report["hard_failures"] + assert "base_selector_rejects_extended_ops" in report["hard_failures"] + + +def test_main_writes_json_and_markdown(tmp_path, capsys): + json_path = tmp_path / "report.json" + md_path = tmp_path / "report.md" + exit_code = main([ + "--case", str(CASE), + "--json", str(json_path), + "--markdown", str(md_path), + ]) + assert exit_code == 0 + data = json.loads(json_path.read_text()) + assert data["hard_failures"] == [] + assert data["topic"] == "topic28-extended-isel" + assert data["schema_version"] == SCHEMA_VERSION + assert data["honesty"] + markdown = md_path.read_text() + assert "Topic 28 Extended Instruction-Selection Feature Case" in markdown + assert "Failure / degradation matrix" in markdown + assert "Hard checks" in markdown + assert capsys.readouterr().out + + +def test_main_rejects_missing_case(tmp_path): + with pytest.raises(SystemExit) as exc: + main(["--case", str(tmp_path / "missing.dsl")]) + assert exc.value.code == 2