Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
27 commits
Select commit Hold shift + click to select a range
c8a8986
ML-DSA on x86-64: NTT and NTT⁻¹ in SSE2
claude Oct 1, 2026
e49e9bb
Merge remote-tracking branch 'origin/main' into pr1-merge
claude Oct 1, 2026
ae7eb79
ML-DSA on x86-64: polynomial arithmetic in SSE2
claude Oct 1, 2026
04e1ac6
Merge remote-tracking branch 'origin/claude/fervent-einstein-ukl7t7' …
claude Oct 1, 2026
6703c7a
ML-DSA on x86-64: key generation, signing and verification generic ov…
claude Oct 1, 2026
3e2aa65
Merge origin/main into claude/fervent-einstein-ukl7t7-mul
claude Oct 1, 2026
55761ef
Merge claude/fervent-einstein-ukl7t7-mul (with origin/main) into clau…
claude Oct 1, 2026
791054f
ML-DSA on x86-64: polynomial arithmetic in AVX2
claude Oct 1, 2026
b818911
Spec: ML-DSA's RejNTTPoly four times (vg_mldsa_rej_ntt_poly4)
claude Oct 1, 2026
779ca16
Spec: give vg_mldsa_rej_ntt_poly4 8 KiB of scratch
claude Oct 1, 2026
e365dde
CI: test ML-DSA on the CPUs of the x86-64 feature matrix
claude Oct 1, 2026
e04e23b
Merge branch 'claude/fervent-einstein-ukl7t7-rej4spec' into claude/fe…
claude Oct 1, 2026
be2f80c
ML-DSA on x86-64: sample  four entries at a time in verification
claude Oct 1, 2026
856f4bb
ML-DSA on x86-64: HighBits and LowBits with AVX2
claude Oct 1, 2026
4563920
Merge main into claude/fervent-einstein-ukl7t7-avx2
claude Oct 1, 2026
11c34b1
Merge branch 'claude/fervent-einstein-ukl7t7-avx2' into claude/ferven…
claude Oct 1, 2026
8160f69
Merge branch 'claude/fervent-einstein-ukl7t7-avx2' into claude/ferven…
claude Oct 1, 2026
6819894
Merge main into claude/fervent-einstein-ukl7t7-avx2
claude Oct 1, 2026
ac45439
Merge branch 'claude/fervent-einstein-ukl7t7-avx2' into claude/ferven…
claude Oct 1, 2026
3f51d68
Merge branch 'claude/fervent-einstein-ukl7t7-avx2' into claude/ferven…
claude Oct 1, 2026
8cc896c
Merge main into claude/fervent-einstein-ukl7t7-rej4
claude Oct 2, 2026
b482411
Merge main into claude/fervent-einstein-ukl7t7-ybits2
claude Oct 2, 2026
db72deb
Merge branch 'claude/fervent-einstein-ukl7t7-rej4' into claude/ferven…
claude Oct 2, 2026
c3926b2
Merge main into claude/fervent-einstein-ukl7t7-ybits2
claude Oct 2, 2026
199d9b3
Merge main into claude/fervent-einstein-ukl7t7-ybits2
claude Oct 2, 2026
88c4d9a
ML-DSA on x86-64: the norm check and MakeHint with AVX2 (sign −8% ins…
alex Oct 2, 2026
2e45fea
Merge #519 (already on the remote branch) with main
claude Oct 2, 2026
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
6 changes: 3 additions & 3 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -923,7 +923,7 @@ yours to keep:

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)</td>
<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

Expand All @@ -939,7 +939,7 @@ yours to keep:

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)</td>
<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

Expand All @@ -955,7 +955,7 @@ yours to keep:

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)</td>
<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

Expand Down
2 changes: 1 addition & 1 deletion docs/algorithms/ml-dsa-44.toml
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,4 @@ family = "Signatures"
specs = ["MlDsa"]
modules = ["src/mldsa44.rs"]
asm = ["mldsa44", "mldsa"]
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" }
2 changes: 1 addition & 1 deletion docs/algorithms/ml-dsa-65.toml
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,4 @@ family = "Signatures"
specs = ["MlDsa"]
modules = ["src/mldsa65.rs"]
asm = ["mldsa65", "mldsa"]
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" }
2 changes: 1 addition & 1 deletion docs/algorithms/ml-dsa-87.toml
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,4 @@ family = "Signatures"
specs = ["MlDsa"]
modules = ["src/mldsa87.rs"]
asm = ["mldsa87", "mldsa"]
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" }
optimized = { x86_64 = "SSE2 and AVX2 polynomial arithmetic; HighBits, LowBits, MakeHint and the norm check with AVX2; matrix sampled four entries at a time in key generation and verification (four SHAKE128 instances at once with AVX2)" }
42 changes: 42 additions & 0 deletions lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,8 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Round.Bits
import VerifiedGarbage.Proof.MlDsa.X86_64.Round.NormLt
import VerifiedGarbage.Proof.MlDsa.X86_64.Round.UseHint
import VerifiedGarbage.Proof.MlDsa.X86_64.Round.MakeHint
import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YBits
import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YHint

/-!
# ML-DSA (FIPS 204) on x86-64: rounding and hints
Expand Down Expand Up @@ -43,6 +45,26 @@ def artifacts : List Artifact := [
contract := Spec.MlDsa.lowBitsContract X86_64.abi
verified := lowBits_verified
spSafe := Code.all_of_allInstrs (by decide +kernel) },
{ Spec.MlDsa.highBitsApi with
name := Spec.MlDsa.highBitsApi.name ++ "_avx2"
target := X86_64.target
doc := Spec.MlDsa.highBitsApi.doc (notes := ["The function computes on eight coefficients at a time in AVX2 \
registers, multiplying by shifts and additions; it needs AVX and AVX2."])
code := Impl.MlDsa.X86_64.Round.highBitsAvx2
contract := Spec.MlDsa.highBitsContract X86_64.abi
verified := highBitsY_verified
spSafe := Code.all_of_allInstrs (by decide +kernel)
features := ["avx", "avx2"] },
{ Spec.MlDsa.lowBitsApi with
name := Spec.MlDsa.lowBitsApi.name ++ "_avx2"
target := X86_64.target
doc := Spec.MlDsa.lowBitsApi.doc (notes := ["The function computes on eight coefficients at a time in AVX2 \
registers, multiplying by shifts and additions; it needs AVX and AVX2."])
code := Impl.MlDsa.X86_64.Round.lowBitsAvx2
contract := Spec.MlDsa.lowBitsContract X86_64.abi
verified := lowBitsY_verified
spSafe := Code.all_of_allInstrs (by decide +kernel)
features := ["avx", "avx2"] },
{ Spec.MlDsa.normLtApi with
target := X86_64.target
doc := Spec.MlDsa.normLtApi.doc
Expand All @@ -57,6 +79,26 @@ def artifacts : List Artifact := [
contract := Spec.MlDsa.makeHintContract X86_64.abi
verified := makeHint_verified
spSafe := Code.all_of_allInstrs (by decide +kernel) },
{ Spec.MlDsa.normLtApi with
name := Spec.MlDsa.normLtApi.name ++ "_avx2"
target := X86_64.target
doc := Spec.MlDsa.normLtApi.doc (notes := ["The function compares eight coefficients at a time in AVX2 \
registers; it needs AVX and AVX2."])
code := Impl.MlDsa.X86_64.Round.normLtAvx2
contract := Spec.MlDsa.normLtContract X86_64.abi
verified := normLtY_verified
spSafe := Code.all_of_allInstrs (by decide +kernel)
features := ["avx", "avx2"] },
{ Spec.MlDsa.makeHintApi with
name := Spec.MlDsa.makeHintApi.name ++ "_avx2"
target := X86_64.target
doc := Spec.MlDsa.makeHintApi.doc (notes := ["The function computes eight hints at a time in AVX2 \
registers, multiplying by shifts and additions; it needs AVX and AVX2."])
code := Impl.MlDsa.X86_64.Round.makeHintAvx2
contract := Spec.MlDsa.makeHintContract X86_64.abi
verified := makeHintY_verified
spSafe := Code.all_of_allInstrs (by decide +kernel)
features := ["avx", "avx2"] },
{ Spec.MlDsa.useHintApi with
target := X86_64.target
doc := Spec.MlDsa.useHintApi.doc
Expand Down
4 changes: 3 additions & 1 deletion lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Backend
import VerifiedGarbage.Impl.MlKem.X86_64.Avx
import VerifiedGarbage.Impl.MlDsa.X86_64.Round.Avx2

/-!
# ML-DSA on x86-64: the polynomial arithmetic with AVX2
Expand Down Expand Up @@ -199,6 +200,7 @@ def subAvx2 : Prog isa :=

/-- The AVX2 code. -/
def Backend.avx2 : Backend :=
⟨nttAvx2, nttInvAvx2, mulAvx2, mulAddAvx2, addAvx2, subAvx2, Sample.Rej4.rejNTT4Avx2, "_avx2"⟩
⟨nttAvx2, nttInvAvx2, mulAvx2, mulAddAvx2, addAvx2, subAvx2, Round.highBitsAvx2,
Round.lowBitsAvx2, Round.normLtAvx2, Round.makeHintAvx2, Sample.Rej4.rejNTT4Avx2, "_avx2"⟩

end VG.Impl.MlDsa.X86_64.Arith
17 changes: 13 additions & 4 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Ntt
import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Mul
import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.AddSub
import VerifiedGarbage.Impl.MlDsa.X86_64.Round.Round
import VerifiedGarbage.Impl.MlDsa.X86_64.Sample.RejNtt4

/-!
Expand All @@ -9,8 +10,9 @@ import VerifiedGarbage.Impl.MlDsa.X86_64.Sample.RejNtt4
Key generation, signing and verification call the polynomial arithmetic of
one implementation, a `Backend`: the code of `vg_mldsa_ntt`,
`vg_mldsa_inv_ntt`, `vg_mldsa_multiply_ntt`, `vg_mldsa_multiply_add_ntt`,
`vg_mldsa_add` and `vg_mldsa_sub`, and of `vg_mldsa_rej_ntt_poly4`
(which samples four entries of the matrix `Â` at once), whose names end with `sfx` (e.g.
`vg_mldsa_add` and `vg_mldsa_sub`, of `vg_mldsa_high_bits`, `vg_mldsa_low_bits`,
`vg_mldsa_norm_lt` and `vg_mldsa_make_hint`, and of `vg_mldsa_rej_ntt_poly4` (which samples four
entries of the matrix `Â` at once), whose names end with `sfx` (e.g.
`_avx2`; nothing for the SSE2 code, `sse2`). Each is a variant of the
interface `MlDsaArith` on x86-64 (`Variants/MlDsaArith/X86_64/`), and the
functions that call them are emitted once for each
Expand All @@ -29,17 +31,24 @@ structure Backend where
mulAdd : Prog isa
add : Prog isa
sub : Prog isa
highBits : Prog isa
lowBits : Prog isa
normLt : Prog isa
makeHint : Prog isa
rej4 : Prog isa
/-- What the names of its functions, and of those calling them, end with. -/
sfx : String

/-- The SSE2 code. -/
def Backend.sse2 : Backend :=
⟨Arith.ntt, Arith.nttInv, Arith.mul, Arith.mulAdd, Arith.add, Arith.sub, Sample.Rej4.rejNTT4, ""⟩
⟨Arith.ntt, Arith.nttInv, Arith.mul, Arith.mulAdd, Arith.add, Arith.sub, Round.highBits,
Round.lowBits, Round.normLt, Round.makeHint, Sample.Rej4.rejNTT4, ""⟩

/-- Every function empty, which the proofs that the functions calling a
backend never write `rsp` (and load MXCSR only to restore it) evaluate in
its place. -/
def Backend.empty : Backend := ⟨.block [], .block [], .block [], .block [], .block [], .block [], .block [], ""⟩
def Backend.empty : Backend :=
⟨.block [], .block [], .block [], .block [], .block [], .block [], .block [], .block [], .block [],
.block [], .block [], ""⟩

end VG.Impl.MlDsa.X86_64.Arith
158 changes: 158 additions & 0 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/Round/Avx2.lean
Original file line number Diff line number Diff line change
@@ -0,0 +1,158 @@
import VerifiedGarbage.Impl.MlDsa.X86_64.Round.Round
import VerifiedGarbage.Impl.MlDsa.X86_64.Arith.Vec
import VerifiedGarbage.Impl.MlKem.X86_64.Avx

/-!
# ML-DSA on x86-64: rounding with AVX2

`vg_mldsa_high_bits_avx2`, `vg_mldsa_low_bits_avx2`,
`vg_mldsa_norm_lt_avx2` and `vg_mldsa_make_hint_avx2` are
`vg_mldsa_high_bits`, `vg_mldsa_low_bits`, `vg_mldsa_norm_lt` and
`vg_mldsa_make_hint` (`Round.lean`) on eight
coefficients at a time, in the doublewords of `ymm` registers: in each
128-bit lane, the VEX.256 form (`toY`) of SSE2 code on the four doublewords
of an `xmm` register (`hbX`, `lbX`, `nlX`, `mhX`), with `rdi` and `r10` at the eight
coefficients of `r` and `out` and `rcx` counting down the 32 vectors.

`Decompose` is the reference implementation's, as in `Round.lean`:
`f = ⌊(⌊(a + 127)/2⁷⌋ · M + 2^(S-1)) / 2^S⌋` and `r₁ = f mod m`, but in 32
bits, which hold every intermediate value, with the multiplication by `M` a
sum of shifts (`mulX`: `1025 = 2¹⁰ + 1` and `11275 = 2¹³ + 2¹¹ + 2¹⁰ + 2³ +
2 + 1`), and, as `f ≤ m`, `r₁` is `f` ANDed with the sign of `f - m`
(`psrad` by 31). `r₀ = a - r₁ · 2γ₂` likewise multiplies by shifts
(`2γ₂ = 2¹⁹ - 2⁹` or `2¹⁷ + 2¹⁵ + 2¹⁴ + 2¹³ + 2¹¹`), plus `q` if negative
(`vcadd`). There are no multiplication instructions, and no branch on data:
the functions branch once on the public `γ₂`, and every address depends
only on the pointers. Each clears the upper halves of the vector registers
before returning (`vzeroupper`).
-/

namespace VG.Impl.MlDsa.X86_64.Round

open VG.X86_64
open VG.Impl.MlKem.X86_64 (xb xmov rcxLoop toY yconst at_)
open VG.Impl.MlDsa.X86_64.Arith (vcadd vcsub)

/-- The shifts whose sum, with 1, is `M`. -/
def dSh (g : Nat) : List Nat := if g = 261888 then [10] else [1, 3, 10, 11, 13]

/-- `xmm0 ← xmm0 · (1 + Σ 2^k)`, through `xmm1` and `xmm2`. -/
def mulX (sh : List Nat) : List Instr :=
xmov .xmm1 .xmm0 :: sh.flatMap fun k =>
[xmov .xmm2 .xmm1, .xop (.shift .pslld .xmm2 (BitVec.ofNat 8 k)), xb .paddd .xmm0 .xmm2]

/-- `xmm0 ← r₁` of `xmm0`, with `127`, `2^(S-1)` and `m` in `xmm8`, `xmm9` and `xmm10`. -/
def hbX (g : Nat) : List Instr :=
[xb .paddd .xmm0 .xmm8, .xop (.shift .psrld .xmm0 7)] ++ mulX (dSh g) ++
[xb .paddd .xmm0 .xmm9, .xop (.shift .psrld .xmm0 (BitVec.ofNat 8 (dShift g))), xmov .xmm1 .xmm0,
xb .psubd .xmm1 .xmm10, .xop (.shift .psrad .xmm1 31), xb .pand .xmm0 .xmm1]

/-- `xmm1 ← xmm0 · 2γ₂`, through `xmm2`. -/
def mul2X (g : Nat) : List Instr :=
if g = 261888 then
[xmov .xmm1 .xmm0, .xop (.shift .pslld .xmm1 19), xmov .xmm2 .xmm0, .xop (.shift .pslld .xmm2 9),
xb .psubd .xmm1 .xmm2]
else
[xmov .xmm1 .xmm0, .xop (.shift .pslld .xmm1 11)] ++ [13, 14, 15, 17].flatMap fun k =>
[xmov .xmm2 .xmm0, .xop (.shift .pslld .xmm2 (BitVec.ofNat 8 k)), xb .paddd .xmm1 .xmm2]

/-- `xmm3 ← r₀` of `xmm0`, with `q` also in `xmm15`. -/
def lbX (g : Nat) : List Instr :=
xmov .xmm3 .xmm0 :: hbX g ++ mul2X g ++ xb .psubd .xmm3 .xmm1 :: vcadd .xmm3 .xmm1

/-- The constants of `hbX` and `lbX`. -/
def yC (g : Nat) : List Instr :=
yconst .xmm8 127 ++ yconst .xmm9 (BitVec.ofNat 32 (dAdd g)) ++ yconst .xmm10 (BitVec.ofNat 32 (dMod g)) ++
yconst .xmm15 8380417

/-- Eight coefficients of `r` (at `rdi`) to `out` (at `r10`), through `x`, the result left in `ymm d`. -/
def bitsBodyY (x : List Instr) (d : XReg) : List Instr :=
[.vmovdquLoad .l256 .xmm0 (at_ .rdi 0)] ++ toY x ++
[.vmovdquStore .l256 (at_ .r10 0) d, .alu .add .rdi (.imm 32), .alu .add .r10 (.imm 32)]

/-- `γ₂` compared, the output pointer to `r10`, and the loop of `γ₂`, `x g` leaving its result in `ymm d`. -/
def bitsY (x : Nat → List Instr) (d : XReg) : Prog isa :=
.seq (.block (gammaCmp .rsi ++ [.mov .r10 (.reg .rdx)]))
(.seq (.ite .e (.seq (.block (yC g32)) (rcxLoop 32 (bitsBodyY (x g32) d)))
(.seq (.block (yC g88)) (rcxLoop 32 (bitsBodyY (x g88) d)))) (.block [.vop .vzeroupper]))

def highBitsAvx2 : Prog isa := bitsY hbX .xmm0

def lowBitsAvx2 : Prog isa := bitsY lbX .xmm3

/-! ## `vg_mldsa_norm_lt_avx2`

The bound is first clamped to `q` (a branch on the public bound), which
changes no result, as every reduced coefficient is less than `q`; then the
differences `a - b` and `(q - b) - a` fit in 32 bits, and one of them is
negative exactly when `a < b` or `q - a < b`. Each doubleword of `ymm10`
ANDs the ORs of the differences of its coefficients; its top bit stays set
while all of them are good. At the end the top bits are spread over their
doublewords (`vpsrad` by 31), and `vpmovmskb` gathers the top bits of the
32 bytes: the result is `(mask + 1) >> 32`, 1 exactly when every bit is set. -/

/-- `ymm10 ← ymm10 & ((a - b) | ((q - b) - a))` for `a` in `xmm0`, `b` in `xmm8`, `q - b` in `xmm9`. -/
def nlX : List Instr :=
[xmov .xmm1 .xmm0, xb .psubd .xmm1 .xmm8, xmov .xmm2 .xmm9, xb .psubd .xmm2 .xmm0, xb .por .xmm1 .xmm2,
xb .pand .xmm10 .xmm1]

def nlBodyY : List Instr := [.vmovdquLoad .l256 .xmm0 (at_ .rdi 0)] ++ toY nlX ++ [.alu .add .rdi (.imm 32)]

/-- The low doubleword of `rax` in each doubleword of `ymm r`. -/
def ybcast (r : XReg) : List Instr := [.vop (.vmovq r .rax), .vop (.vpbroadcastd .l256 r r)]

/-- `b` in `ymm8`, `q - b` in `ymm9` and all ones in `ymm10`. -/
def nlConsts : List Instr :=
[.mov32 .rax (.reg .rsi)] ++ ybcast .xmm8 ++ [.mov32 .rax (.imm qImm), .alu32 .sub .rax (.reg .rsi)] ++
ybcast .xmm9 ++ yconst .xmm10 0xFFFFFFFF

def nlEnd : List Instr :=
toY [.xop (.shift .psrad .xmm10 31)] ++ [.vpmovmskb .l256 .rax .xmm10, .vop .vzeroupper] ++
[.alu .add .rax (.imm 1), .shift .shr .rax 32]

/-- The bound, clamped to `q`. -/
def nlPro : Prog isa :=
.seq (.block [.mov32 .rsi (.reg .rsi), .alu32 .cmp .rsi (.imm qImm)]) (.ite .b (.block []) (.block [.mov32 .rsi (.imm qImm)]))

def normLtAvx2 : Prog isa := .seq nlPro (.seq (.block nlConsts) (.seq (rcxLoop 32 nlBodyY) (.block nlEnd)))

/-! ## `vg_mldsa_make_hint_avx2`

In each lane, `r₁` of `r` and of `(r + z) mod q` (`hbX` twice, with `vcsub`
between), XORed, is nonzero (`(x + 63) >> 6`, as `x < 64`) exactly when the
hint is 1 (`mhX`); the eight hints are stored. Their count is the sum of the
nibbles of the byte mask of the hints shifted to the top bit of their low
bytes (`vpmovmskb`, which has hint `i` at bit `4i`): three shifts and
additions put it in the low nibble (`cntH`), which is added to `r9`. -/

/-- `r₁` of `r` to `xmm3`, and `(r + z) mod q` to `xmm0`, with `r` in `xmm4` and `z` in `xmm5`. -/
def mhMid : List Instr := [xmov .xmm3 .xmm0, xmov .xmm0 .xmm4, xb .paddd .xmm0 .xmm5] ++ vcsub .xmm0 .xmm1

/-- The hint from the two `r₁`, to `xmm0`, and shifted to bit 7, to `xmm1`. -/
def mhTail : List Instr :=
[xb .pxor .xmm0 .xmm3, xb .paddd .xmm0 .xmm11, .xop (.shift .psrld .xmm0 6), xmov .xmm1 .xmm0,
.xop (.shift .pslld .xmm1 7)]

/-- `xmm0 ← the hints` of `r` in `xmm0` and `z` in `xmm5`, and `xmm1 ← xmm0 << 7`, with `63` in `xmm11`. -/
def mhX (g : Nat) : List Instr := [xmov .xmm4 .xmm0] ++ hbX g ++ mhMid ++ hbX g ++ mhTail

/-- `r9 ← r9 +` the sum of the nibbles of `eax`, through `rdx`. -/
def cntH : List Instr :=
[.mov32 .rdx (.reg .rax), .shift32 .shr .rdx 4, .alu32 .add .rax (.reg .rdx),
.mov32 .rdx (.reg .rax), .shift32 .shr .rdx 8, .alu32 .add .rax (.reg .rdx),
.mov32 .rdx (.reg .rax), .shift32 .shr .rdx 16, .alu32 .add .rax (.reg .rdx),
.alu32 .and .rax (.imm 15), .alu .add .r9 (.reg .rax)]

/-- Eight coefficients of `r` (at `rsi`) and `z` (at `rdi`) to the hints at `r10`, and their count added to `r9`. -/
def mhBodyY (g : Nat) : List Instr :=
[.vmovdquLoad .l256 .xmm0 (at_ .rsi 0), .vmovdquLoad .l256 .xmm5 (at_ .rdi 0)] ++ toY (mhX g) ++
[.vmovdquStore .l256 (at_ .r10 0) .xmm0, .vpmovmskb .l256 .rax .xmm1] ++ cntH ++
[.alu .add .rdi (.imm 32), .alu .add .rsi (.imm 32), .alu .add .r10 (.imm 32)]

def mhY (g : Nat) : Prog isa := .seq (.block (yC g ++ yconst .xmm11 63)) (rcxLoop 32 (mhBodyY g))

def makeHintAvx2 : Prog isa :=
.seq (.block (gammaCmp .rdx ++ [.mov .r10 (.reg .rcx), .mov32 .r9 (.imm 0)]))
(.seq (.ite .e (mhY g32) (mhY g88)) (.block [.mov .rax (.reg .r9), .vop .vzeroupper]))

end VG.Impl.MlDsa.X86_64.Round
8 changes: 4 additions & 4 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean
Original file line number Diff line number Diff line change
Expand Up @@ -193,18 +193,18 @@ def ballAt (len tau : Nat) (c : Ptr) : Prog isa :=
callP "vg_mldsa_sample_in_ball" P.ball [.ptr (sc oCT), .imm len, .imm tau, .ptr c, .ptr (sc oPS)]

def highBitsAt (r : Ptr) (gamma2 : Nat) (out : Ptr) : Prog isa :=
callP "vg_mldsa_high_bits" P.highBits [.ptr r, .imm gamma2, .ptr out]
callP ("vg_mldsa_high_bits" ++ P.sfx) P.highBits [.ptr r, .imm gamma2, .ptr out]

def lowBitsAt (r : Ptr) (gamma2 : Nat) (out : Ptr) : Prog isa :=
callP "vg_mldsa_low_bits" P.lowBits [.ptr r, .imm gamma2, .ptr out]
callP ("vg_mldsa_low_bits" ++ P.sfx) P.lowBits [.ptr r, .imm gamma2, .ptr out]

/-- `‖f‖∞ < bound`, and `r15 ← r15 ∧ result`. -/
def normAt (f : Ptr) (bound : Nat) : Prog isa :=
.seq (callP "vg_mldsa_norm_lt" P.normLt [.ptr f, .imm bound]) (.block [.alu32 .and .r15 (.reg .rax)])
.seq (callP ("vg_mldsa_norm_lt" ++ P.sfx) P.normLt [.ptr f, .imm bound]) (.block [.alu32 .and .r15 (.reg .rax)])

/-- `MakeHint` of `z` and `r` to `h`, and the number of 1s added to `ONES`. -/
def makeHintAt (z r : Ptr) (gamma2 : Nat) (h : Ptr) : Prog isa :=
.seq (callP "vg_mldsa_make_hint" P.makeHint [.ptr z, .ptr r, .imm gamma2, .ptr h])
.seq (callP ("vg_mldsa_make_hint" ++ P.sfx) P.makeHint [.ptr z, .ptr r, .imm gamma2, .ptr h])
(.block [.mov32 .rcx (.mem (at_ .rbx oONES)), .alu32 .add .rcx (.reg .rax), .store (at_ .rbx oONES) .rcx])

def simpleBitPackAt (f : Ptr) (b : Nat) (out : Ptr) (len : Nat) : Prog isa :=
Expand Down
2 changes: 1 addition & 1 deletion lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean
Original file line number Diff line number Diff line change
Expand Up @@ -189,7 +189,7 @@ def hintUnpackAt (y : Ptr) (len omega : Nat) (h : Ptr) (hlen : Nat) : Prog isa :
[(.rdi, .ptr y), (.rsi, .imm len), (.rdx, .imm omega), (.rcx, .ptr h), (.r8, .imm hlen)]

def normLtAt (f : Ptr) (bound : Nat) : Prog isa :=
callAt "vg_mldsa_norm_lt" P.normLt [(.rdi, .ptr f), (.rsi, .imm bound)]
callAt ("vg_mldsa_norm_lt" ++ P.sfx) P.normLt [(.rdi, .ptr f), (.rsi, .imm bound)]

end

Expand Down
Loading
Loading