Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 21 additions & 0 deletions scratchv/analysis/ir_verifier.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,11 @@ class OpcodeSpec:
OpCode.RETURN: OpcodeSpec(0, 1, False),
OpCode.FOR: OpcodeSpec(0, 0, True),
OpCode.ENDFOR: OpcodeSpec(0, 0, False),
# FWHT/WINOGRAD keep a single dtype across operands and result; SPMM_CSR
# mixes float values with integer CSR indices, so it is validated below.
OpCode.FWHT: OpcodeSpec(1, 1, True, "T"),
OpCode.WINOGRAD_CONV: OpcodeSpec(2, 3, True, "T"),
OpCode.SPMM_CSR: OpcodeSpec(4, 4, True, None),
})
_RULES = {name: i for i, name in enumerate((
"def-before-use", "label-existence", "block-termination", "type-consistency",
Expand Down Expand Up @@ -398,6 +403,22 @@ def error(reason, message):
error("gather-indices", "GATHER indices must be i32 or i64")
if inst.dest is not None and inst.dest.dtype != data.dtype:
error("gather-result", "GATHER result dtype must match data dtype")
if inst.opcode == OpCode.SPMM_CSR:
values, col, rowptr, b = inst.operands
if values.dtype != b.dtype:
error("spmm-value-types", "SPMM_CSR values and B must share a dtype")
if col.dtype not in _INTS:
error("spmm-index", "SPMM_CSR col must be i32 or i64")
if rowptr.dtype not in _INTS:
error("spmm-index", "SPMM_CSR rowptr must be i32 or i64")
if inst.dest is not None and inst.dest.dtype != values.dtype:
error("spmm-result", "SPMM_CSR result dtype must match values")
if inst.opcode == OpCode.WINOGRAD_CONV:
bias = inst.operands[2] if count == 3 else None
cout = inst.attrs.get("cout")
if (bias is not None and isinstance(cout, int) and not isinstance(cout, bool)
and cout > 0 and tuple(bias.shape) not in ((), (cout,))):
error("winograd-bias", "WINOGRAD_CONV bias shape must match cout or be scalar")
if inst.opcode == OpCode.BR_IF:
if count == 1:
if isinstance(inst.operands[0].dtype, DataType) and inst.operands[0].dtype not in _INTS:
Expand Down
12 changes: 12 additions & 0 deletions scratchv/backend/asm_emit.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,8 @@
MachineOp.DIV: "div",
MachineOp.MAX: "max",
MachineOp.SRAI: "srai",
MachineOp.SLLI: "slli",
MachineOp.SRLI: "srli",
MachineOp.XOR: "xor",
MachineOp.AND: "and",
MachineOp.SLT: "slt",
Expand All @@ -34,6 +36,7 @@
MachineOp.J: "j",
MachineOp.JAL: "jal",
MachineOp.JALR: "jalr",
MachineOp.RET: "ret",
MachineOp.BEQ: "beq",
MachineOp.BNE: "bne",
MachineOp.BLT: "blt",
Expand Down Expand Up @@ -71,6 +74,7 @@
MachineOp.FSW: "fsw",
MachineOp.FMV_S: "fmv.s",
MachineOp.FMV_S_X: "fmv.s.x",
MachineOp.FMV_W_X: "fmv.w.x",
}


Expand Down Expand Up @@ -162,6 +166,14 @@ def _format_instr(self, instr: MachineInstr) -> str:
f"{_comment_suffix(instr)}"
)

if instr.op == MachineOp.FLW and instr.dst and instr.src1:
address = _memory_address(instr.src1)
return f" flw {_fmt_op(instr.dst)}, {address}{_comment_suffix(instr)}"

if instr.op == MachineOp.FSW and instr.dst and instr.src1:
address = _memory_address(instr.src1)
return f" fsw {_fmt_op(instr.dst)}, {address}{_comment_suffix(instr)}"

# Branch/jump/call use the structured target resolved above.
if instr.op in (MachineOp.CALL, MachineOp.J, MachineOp.JAL,
MachineOp.BNEZ, MachineOp.BEQ, MachineOp.BNE,
Expand Down
12 changes: 12 additions & 0 deletions scratchv/backend/machine_semantics.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,12 @@ class MachineOpSemantics:
MachineOp.SRAI: MachineOpSemantics(
defs=(0,), uses=(1,), immediate_positions=(2,), n_phys=2
),
MachineOp.SLLI: MachineOpSemantics(
defs=(0,), uses=(1,), immediate_positions=(2,), n_phys=2
),
MachineOp.SRLI: MachineOpSemantics(
defs=(0,), uses=(1,), immediate_positions=(2,), n_phys=2
),
MachineOp.XOR: _DEF_USE_USE,
MachineOp.AND: _DEF_USE_USE,
MachineOp.SLT: _DEF_USE_USE,
Expand Down Expand Up @@ -121,6 +127,11 @@ class MachineOpSemantics:
is_terminator=True,
n_phys=2,
),
# ret -> jalr x0, 0(ra); emitted verbatim by the assembler.
MachineOp.RET: MachineOpSemantics(
is_terminator=True,
implicit_uses=frozenset({"ra"}),
),
MachineOp.JAL: MachineOpSemantics(
defs=(0,),
is_terminator=True,
Expand Down Expand Up @@ -207,6 +218,7 @@ class MachineOpSemantics:
defs=(0,), uses=(1,), n_phys=2, is_pseudo=True
),
MachineOp.FMV_S_X: _DEF_USE,
MachineOp.FMV_W_X: _DEF_USE,
}


Expand Down
4 changes: 4 additions & 0 deletions scratchv/backend/machine_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,8 @@ class MachineOp(enum.Enum):
DIV = "div"
MAX = "max" # pseudo: max rd, rs1, rs2
SRAI = "srai"
SLLI = "slli"
SRLI = "srli"
XOR = "xor"
AND = "and"
SLT = "slt"
Expand All @@ -47,6 +49,7 @@ class MachineOp(enum.Enum):
J = "j"
JAL = "jal"
JALR = "jalr"
RET = "ret" # pseudo: jalr x0, 0(ra)
BEQ = "beq"
BNE = "bne"
BLT = "blt"
Expand Down Expand Up @@ -93,6 +96,7 @@ class MachineOp(enum.Enum):
FSW = "fsw"
FMV_S = "fmv.s"
FMV_S_X = "fmv.s.x"
FMV_W_X = "fmv.w.x"


# ═══════════════════════════════════════════════════════════════════════════════
Expand Down
Loading
Loading