Skip to content
Merged
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
6 changes: 3 additions & 3 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -891,7 +891,7 @@ yours to keep:

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits and LowBits with AVX2; matrix sampled four entries at a time in 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 verification (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

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

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits and LowBits with AVX2; matrix sampled four entries at a time in 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 verification (four SHAKE128 instances at once with AVX2)</td>

<td>✅ SHA extensions</td>

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

<td>✅</td>

<td>✅ AVX2; SSE2 and AVX2 polynomial arithmetic; HighBits and LowBits with AVX2; matrix sampled four entries at a time in 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 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; HighBits and LowBits with AVX2; matrix sampled four entries at a time in 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 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; HighBits and LowBits with AVX2; matrix sampled four entries at a time in 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 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; HighBits and LowBits with AVX2; matrix sampled four entries at a time in 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 verification (four SHAKE128 instances at once with AVX2)" }
21 changes: 21 additions & 0 deletions lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ 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 @@ -78,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
2 changes: 1 addition & 1 deletion lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean
Original file line number Diff line number Diff line change
Expand Up @@ -201,6 +201,6 @@ def subAvx2 : Prog isa :=
/-- The AVX2 code. -/
def Backend.avx2 : Backend :=
⟨nttAvx2, nttInvAvx2, mulAvx2, mulAddAvx2, addAvx2, subAvx2, Round.highBitsAvx2,
Round.lowBitsAvx2, Sample.Rej4.rejNTT4Avx2, "_avx2"⟩
Round.lowBitsAvx2, Round.normLtAvx2, Round.makeHintAvx2, Sample.Rej4.rejNTT4Avx2, "_avx2"⟩

end VG.Impl.MlDsa.X86_64.Arith
12 changes: 8 additions & 4 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean
Original file line number Diff line number Diff line change
Expand Up @@ -10,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`, of `vg_mldsa_high_bits` and `vg_mldsa_low_bits`, 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 @@ -32,19 +33,22 @@ structure Backend where
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, Round.highBits,
Round.lowBits, Sample.Rej4.rejNTT4, ""⟩
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 [], .block [], .block [], ""⟩
⟨.block [], .block [], .block [], .block [], .block [], .block [], .block [], .block [], .block [],
.block [], .block [], ""⟩

end VG.Impl.MlDsa.X86_64.Arith
85 changes: 81 additions & 4 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/Round/Avx2.lean
Original file line number Diff line number Diff line change
Expand Up @@ -5,11 +5,13 @@ import VerifiedGarbage.Impl.MlKem.X86_64.Avx
/-!
# ML-DSA on x86-64: rounding with AVX2

`vg_mldsa_high_bits_avx2` and `vg_mldsa_low_bits_avx2` are
`vg_mldsa_high_bits` and `vg_mldsa_low_bits` (`Round.lean`) on eight
`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`), with `rdi` and `r10` at the eight
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`:
Expand All @@ -29,7 +31,7 @@ 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)
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]
Expand Down Expand Up @@ -78,4 +80,79 @@ 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
4 changes: 2 additions & 2 deletions lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean
Original file line number Diff line number Diff line change
Expand Up @@ -200,11 +200,11 @@ def lowBitsAt (r : Ptr) (gamma2 : Nat) (out : Ptr) : Prog isa :=

/-- `‖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
8 changes: 8 additions & 0 deletions lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,8 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Mul
import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.AddSub
import VerifiedGarbage.Proof.MlKem.X86_64.ArithOk
import VerifiedGarbage.Proof.MlDsa.X86_64.Round.Bits
import VerifiedGarbage.Proof.MlDsa.X86_64.Round.MakeHint
import VerifiedGarbage.Proof.MlDsa.X86_64.Round.NormLt
import VerifiedGarbage.Proof.MlDsa.X86_64.Sample.Rej4Verified

/-!
Expand Down Expand Up @@ -52,6 +54,8 @@ structure BackendOk (B : Backend) : Prop where
sub : FnOk (fun S => Spec.MlDsa.subContract X86_64.abi S) B.sub
highBits : FnOk (fun S => Spec.MlDsa.highBitsContract X86_64.abi S) B.highBits
lowBits : FnOk (fun S => Spec.MlDsa.lowBitsContract X86_64.abi S) B.lowBits
normLt : FnOk (fun S => Spec.MlDsa.normLtContract X86_64.abi S) B.normLt
makeHint : FnOk (fun S => Spec.MlDsa.makeHintContract X86_64.abi S) B.makeHint
rej4 : Rej4Ok B.rej4

/-- An implementation of the polynomial arithmetic on x86-64. -/
Expand Down Expand Up @@ -87,6 +91,10 @@ def ArithImpl.sse2 : ArithImpl where
(by decide +kernel)
lowBits := FnOk.of Round.lowBits_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
(by decide +kernel)
normLt := FnOk.of Round.normLt_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
(by decide +kernel)
makeHint := FnOk.of Round.makeHint_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
(by decide +kernel)
rej4 := ⟨Rej4.rejNTT4_verified, Proof.MlKem.X86_64.nosp_of (by decide +kernel), by decide +kernel,
by decide +kernel, Code.all_of_allInstrs (by decide +kernel)⟩ }
features := []
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Backend
import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.YNtt
import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.YMul
import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.YAddSub
import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YBits
import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YHint

/-!
# ML-DSA on x86-64: the polynomial arithmetic with AVX2, as an `ArithImpl`
Expand Down Expand Up @@ -36,6 +36,10 @@ def ArithImpl.avx2 : ArithImpl where
(by decide +kernel)
lowBits := FnOk.of Round.lowBitsY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
(by decide +kernel)
normLt := FnOk.of Round.normLtY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
(by decide +kernel)
makeHint := FnOk.of Round.makeHintY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
(by decide +kernel)
rej4 := ⟨Rej4.rejNTT4Avx2_verified, Proof.MlKem.X86_64.nosp_of (by decide +kernel), by decide +kernel,
by decide +kernel, Code.all_of_allInstrs (by decide +kernel)⟩ }
features := ["avx", "avx2"]
Expand Down
Loading
Loading