diff --git a/README.md b/README.md
index 25a3fa678..1a8ace9e7 100644
--- a/README.md
+++ b/README.md
@@ -923,7 +923,7 @@ yours to keep:
✅ |
-✅ 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) |
+✅ 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) |
✅ SHA extensions |
@@ -939,7 +939,7 @@ yours to keep:
✅ |
-✅ 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) |
+✅ 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) |
✅ SHA extensions |
@@ -955,7 +955,7 @@ yours to keep:
✅ |
-✅ 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) |
+✅ 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) |
✅ SHA extensions |
diff --git a/docs/algorithms/ml-dsa-44.toml b/docs/algorithms/ml-dsa-44.toml
index 23c04ea14..77df162de 100644
--- a/docs/algorithms/ml-dsa-44.toml
+++ b/docs/algorithms/ml-dsa-44.toml
@@ -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)" }
diff --git a/docs/algorithms/ml-dsa-65.toml b/docs/algorithms/ml-dsa-65.toml
index 04ce839cb..3d5f174db 100644
--- a/docs/algorithms/ml-dsa-65.toml
+++ b/docs/algorithms/ml-dsa-65.toml
@@ -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)" }
diff --git a/docs/algorithms/ml-dsa-87.toml b/docs/algorithms/ml-dsa-87.toml
index 88e39d47a..75d9dc3d0 100644
--- a/docs/algorithms/ml-dsa-87.toml
+++ b/docs/algorithms/ml-dsa-87.toml
@@ -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)" }
diff --git a/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean b/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean
index 09ad7c11e..69e4ffd2e 100644
--- a/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean
+++ b/lean/VerifiedGarbage/Artifacts/MlDsaRound/X86_64.lean
@@ -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
@@ -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
@@ -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
diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean
index 3c987cbb6..2d76aa1a0 100644
--- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean
+++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Avx2.lean
@@ -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
@@ -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
diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean
index 65d34d28d..b170d109b 100644
--- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean
+++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Arith/Backend.lean
@@ -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
/-!
@@ -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
@@ -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
diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Round/Avx2.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Round/Avx2.lean
new file mode 100644
index 000000000..79292bcbc
--- /dev/null
+++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Round/Avx2.lean
@@ -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
diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean
index 4fa71ab43..90f290c61 100644
--- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean
+++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Sign/Frag.lean
@@ -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 :=
diff --git a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean
index 72ef6c175..b156b4e60 100644
--- a/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean
+++ b/lean/VerifiedGarbage/Impl/MlDsa/X86_64/Verify/Frag.lean
@@ -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
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean
index d8ad5b2bc..759bdcb13 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/Backend.lean
@@ -4,6 +4,9 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.NttInv
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
/-!
@@ -49,6 +52,10 @@ structure BackendOk (B : Backend) : Prop where
mulAdd : FnOk (fun S => Spec.MlDsa.mulAddContract X86_64.abi S) B.mulAdd
add : FnOk (fun S => Spec.MlDsa.addContract X86_64.abi S) B.add
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. -/
@@ -80,6 +87,14 @@ def ArithImpl.sse2 : ArithImpl where
(by decide +kernel)
sub := FnOk.of Arith.sub_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
(by decide +kernel)
+ highBits := FnOk.of Round.highBits_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
+ (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 := []
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean
index 0036f4cdb..f91cfc25f 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Arith/BackendAvx2.lean
@@ -2,6 +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.YHint
/-!
# ML-DSA on x86-64: the polynomial arithmetic with AVX2, as an `ArithImpl`
@@ -31,6 +32,14 @@ def ArithImpl.avx2 : ArithImpl where
(by decide +kernel)
sub := FnOk.of Arith.subY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
(by decide +kernel)
+ highBits := FnOk.of Round.highBitsY_verified (by decide +kernel) (by decide +kernel) (by decide +kernel)
+ (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"]
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBits.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBits.lean
new file mode 100644
index 000000000..10064a572
--- /dev/null
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBits.lean
@@ -0,0 +1,273 @@
+import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YBlock
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.YBase
+import VerifiedGarbage.Proof.MlDsa.X86_64.Round.Bits
+
+/-!
+# ML-DSA on x86-64: `vg_mldsa_high_bits_avx2` and `vg_mldsa_low_bits_avx2`
+
+Untrusted: everything here is checked by Lean. After the constants
+(`yC_ok`), each iteration of the loop loads eight coefficients of `r`, does
+in each lane what `hbX` or `lbX` does to a register (`YBlock.lean`), whose
+proof holds of each lane (`ylanes`), and stores the eight results to `out`
+(`YMap.step`); the loop leaves `out` with all 256 (`YMap.loop_ok`), for the
+`γ₂` the function compared (`ybits_ok`).
+-/
+
+namespace VG.Proof.MlDsa.X86_64.Arith
+
+open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round
+open VG.Proof.MlDsa.Arith
+open VG.Proof.MlDsa.X86_64.Round (HbC bc_ofDwords)
+open VG.Proof.MlKem.X86_64 (Keep XOnly YOnly ylanes yld_ok yconst_ok WP.keep ifp ifn ptr_step GOnly wp_rcxLoopY
+ add_ofNat_zero lane_setReg lane_setFlags sx32 State.setMem_ymm)
+open VG.Impl.MlKem.X86_64 (xb xmov toY yconst)
+open VG.Spec.MlDsa (q n coeffAt)
+
+/-- Each doubleword of `v`. -/
+def bc (v : BitVec 32) : BitVec 128 := ofDwords v v v v
+
+/-- The constants of `hbX` and `lbX` in both lanes. -/
+def YC (g : Nat) (s : State) : Prop :=
+ ∀ l < 2, s.lane .xmm8 l = bc 127 ∧ s.lane .xmm9 l = bc (BitVec.ofNat 32 (dAdd g)) ∧
+ s.lane .xmm10 l = bc (BitVec.ofNat 32 (dMod g)) ∧ s.lane .xmm15 l = qV
+
+theorem YC.hbc {g : Nat} {s : State} (h : YC g s) {l : Nat} (hl : l < 2) : HbC g (s.proj l) :=
+ ⟨by rw [State.proj_xmm, (h l hl).1]; exact bc_ofDwords _, by rw [State.proj_xmm, (h l hl).2.1]; exact bc_ofDwords _,
+ by rw [State.proj_xmm, (h l hl).2.2.1]; exact bc_ofDwords _⟩
+
+theorem YC.q {g : Nat} {s : State} (h : YC g s) {l : Nat} (hl : l < 2) : (s.proj l).xmm .xmm15 = qV := by
+ rw [State.proj_xmm, (h l hl).2.2.2]
+
+/-- The constants are kept by code that writes only `xmm0` to `xmm3`. -/
+theorem YC.keep {g : Nat} {s s' : State} (h : YC g s) {rs : List XReg} (hr : ∀ r ∈ rs, r ∈ [.xmm0, .xmm1, .xmm2, .xmm3])
+ (hl : ∀ r ∉ rs, ∀ l < 2, s'.lane r l = s.lane r l) : YC g s' := fun l hl' => by
+ have n : ∀ r ∈ [XReg.xmm8, .xmm9, .xmm10, .xmm15], r ∉ rs := fun r hr' hr'' => by
+ have := hr r hr''; simp only [List.mem_cons, List.not_mem_nil, or_false] at hr' this
+ rcases hr' with rfl | rfl | rfl | rfl <;> rcases this with h | h | h | h <;> cases h
+ rw [hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl']
+ exact h l hl'
+
+theorem yC_ok (g : Nat) (s : State) :
+ WP isa (.block (yC g)) s fun s' => YC g s' ∧ Keep [.rax] s s' ∧ s'.mem = s.mem ∧ s'.mxcsr = s.mxcsr ∧
+ ∀ r ∉ [XReg.xmm8, .xmm9, .xmm10, .xmm15], ∀ l < 2, s'.lane r l = s.lane r l := by
+ simp only [yC, List.append_assoc]
+ rw [WP.block_append_iff]
+ refine WP.mono (yconst_ok .xmm8 _ s) fun s1 ⟨l1, k1, m1, x1, o1⟩ => ?_
+ rw [WP.block_append_iff]
+ refine WP.mono (yconst_ok .xmm9 _ s1) fun s2 ⟨l2, k2, m2, x2, o2⟩ => ?_
+ rw [WP.block_append_iff]
+ refine WP.mono (yconst_ok .xmm10 _ s2) fun s3 ⟨l3, k3, m3, x3, o3⟩ => ?_
+ refine WP.mono (yconst_ok .xmm15 _ s3) fun s4 ⟨l4, k4, m4, x4, o4⟩ =>
+ ⟨fun l hl => ⟨?_, ?_, ?_, ?_⟩, (((k1.trans k2).trans k3).trans k4).mono (by simp), m4.trans (m3.trans (m2.trans m1)),
+ x4.trans (x3.trans (x2.trans x1)), fun r hr l hl => ?_⟩
+ · rw [o4 _ (by decide) l hl, o3 _ (by decide) l hl, o2 _ (by decide) l hl, l1 l hl]; rfl
+ · rw [o4 _ (by decide) l hl, o3 _ (by decide) l hl, l2 l hl]; rfl
+ · rw [o4 _ (by decide) l hl, l3 l hl]; rfl
+ · rw [l4 l hl]; decide
+ · simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr
+ rw [o4 r hr.2.2.2 l hl, o3 r hr.2.2.1 l hl, o2 r hr.2.1 l hl, o1 r hr.1 l hl]
+
+namespace YMap
+
+/-- After `i` vectors of eight, each coefficient of `out` before `8i` is `v k`. -/
+structure Inv (s₀ : State) (g : Nat) (v : Nat → BitVec 32) (i : Nat) (s : State) : Prop where
+ rdi : s.gpr .rdi = s₀.gpr .rdi + BitVec.ofNat 64 (32 * i)
+ r10 : s.gpr .r10 = s₀.gpr .rdx + BitVec.ofNat 64 (32 * i)
+ rd : s.rd = s₀.rd
+ wr : s.wr = s₀.wr
+ yc : YC g s
+ frame : Frame [pR (s₀.gpr .rdx)] s₀.mem s.mem
+ coeff : ∀ k < 256, coeffAt s.mem (s₀.gpr .rdx) k = if k < 8 * i then v k else coeffAt s₀.mem (s₀.gpr .rdx) k
+
+section
+variable {s₀ : State} (hrd : s₀.rd = [pR (s₀.gpr .rdi)]) (hwr : s₀.wr = [pR (s₀.gpr .rdx)])
+ (hdis : (pR (s₀.gpr .rdi)).Disjoint (pR (s₀.gpr .rdx)))
+ {g : Nat} {x : List Instr} {d : XReg} {L : BitVec 32 → BitVec 32}
+ (hX : ∀ t : State, HbC g t → t.xmm .xmm15 = qV →
+ WP isa (.block x) t fun t' => (∀ e < 4, dword (t'.xmm d) e = L (dword (t.xmm .xmm0) e)) ∧
+ XOnly [.xmm0, .xmm1, .xmm2, .xmm3] t t')
+ (hY : laneSseBlock (toY x) = some x)
+include hrd hwr hdis hX hY
+
+/-- An iteration, which stores `L` of the eight coefficients of `r` to `out`. -/
+theorem step {i : Nat} (hi : i < 32) {s : State}
+ (hI : Inv s₀ g (fun k => L (coeffAt s₀.mem (s₀.gpr .rdi) k)) i s) :
+ WP isa (.block (bitsBodyY x d ++ ([.alu .sub .rcx (.imm 1)] : List Instr))) s fun s' =>
+ Inv s₀ g (fun k => L (coeffAt s₀.mem (s₀.gpr .rdi) k)) (i + 1) s' ∧
+ s'.gpr .rcx = s.gpr .rcx - 1 ∧ s'.zf = some (s.gpr .rcx - 1 == 0) := by
+ have j0 : 8 * i + 8 ≤ 256 := by omega
+ have hw : pR (s₀.gpr .rdx) ∈ s.wr := by rw [hI.wr, hwr]; simp
+ have hr : pR (s₀.gpr .rdi) ∈ s.rd ++ s.wr := by rw [hI.rd, hrd]; simp
+ have e1 : s.gpr .rdi + BitVec.ofNat 64 0 = coeffAddr (s₀.gpr .rdi) (8 * i) := by
+ rw [add_ofNat_zero, hI.rdi]; congr 2; omega
+ have e2 : s.gpr .r10 = coeffAddr (s₀.gpr .rdx) (8 * i) := by
+ rw [hI.r10]; congr 2; omega
+ rw [bitsBodyY, List.append_assoc, List.append_assoc, WP.block_append_iff]
+ refine WP.mono (yld_ok (by rw [e1]; exact f_in32 hr j0)) fun s1 ⟨L1, o1⟩ => ?_
+ have yc1 : YC g s1 := hI.yc.keep (by simp) o1.lane
+ rw [WP.block_append_iff]
+ refine WP.mono (ylanes hY (P := fun l t => ∀ e < 4, dword (t.xmm d) e = L (dword ((s1.proj l).xmm .xmm0) e))
+ fun l hl => hX _ (yc1.hbc hl) (yc1.q hl)) fun s3 ⟨B3, o3⟩ => ?_
+ have o13 := o1.trans o3
+ have yc3 : YC g s3 := yc1.keep (by simp) o3.lane
+ have g3 : s3.gpr .r10 = coeffAddr (s₀.gpr .rdx) (8 * i) := by rw [o13.gpr, e2]
+ have g3' : s3.gpr .rdi = coeffAddr (s₀.gpr .rdi) (8 * i) := by rw [o13.gpr, ← e1, add_ofNat_zero]
+ have w0 : InRegions s3.wr (s3.gpr .r10) 32 := by
+ rw [o13.wr, g3]; exact f_in32 hw j0
+ vrund [State.store256_eq, State.setMem_gpr, State.setMem_wr, State.setMem_mem, State.setMem_rd,
+ State.setMem_ymm, w0, sx32]
+ refine ⟨⟨?_, ?_, ?_, ?_, ?_, ?_, fun k hk => ?_⟩, ?_⟩
+ · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq, State.setMem_gpr]
+ rw [o13.gpr, hI.rdi]; exact ptr_step _ i 32
+ · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq, State.setMem_gpr]
+ rw [o13.gpr, hI.r10]; exact ptr_step _ i 32
+ · simp only [RegUpd.rd_setReg, RegUpd.rd_setFlags, State.setMem_rd]; rw [o13.rd, hI.rd]
+ · simp only [RegUpd.wr_setReg, RegUpd.wr_setFlags, State.setMem_wr]; rw [o13.wr, hI.wr]
+ · exact yc3.keep (rs := []) (by simp) fun r _ l _ => by simp only [lane_setReg, lane_setFlags, State.setMem_lane]
+ · simp only [RegUpd.mem_setReg, RegUpd.mem_setFlags, State.setMem_mem]
+ rw [g3, o13.mem]; exact hI.frame.writeW (List.mem_singleton_self _) _ (pR_contains32 _ j0)
+ · simp only [RegUpd.mem_setReg, RegUpd.mem_setFlags, State.setMem_mem]
+ rw [g3, o13.mem, coeffAt_write256 _ _ j0 _ hk]
+ split
+ · rename_i h
+ rw [ifp (show k < 8 * (i + 1) by omega), State.ymm, extract_ymm _ _ (by omega)]
+ have hc : ∀ l < 2, ∀ e < 4, dword (s3.lane d l) e = L (coeffAt s₀.mem (s₀.gpr .rdi) (8 * i + 4 * l + e)) :=
+ fun l hl e he => by
+ rw [← State.proj_xmm, B3 l hl e he, State.proj_xmm, L1 l hl, e1, dword_readW _ _ he, lane_load,
+ coeffAddr_add, ← coeffAt_eq, coeffAt_frame hI.frame (by simpa using hdis) (by rw [n_eq]; omega)]
+ split
+ · rename_i h4
+ have := hc 0 (by decide) (k - 8 * i) h4
+ rw [show 8 * i + 4 * 0 + (k - 8 * i) = k by omega] at this
+ exact this
+ · rename_i h4
+ have := hc 1 (by decide) (k - 8 * i - 4) (by omega)
+ rw [show 8 * i + 4 * 1 + (k - 8 * i - 4) = k by omega] at this
+ exact this
+ · rename_i h
+ rw [hI.coeff k hk]
+ by_cases h' : k < 8 * i
+ · rw [ifp h', ifp (by omega)]
+ · rw [ifn h', ifn (by omega)]
+ · exact ⟨by rw [o13.gpr], by rw [o13.gpr]⟩
+
+/-- The constants and the loop. -/
+theorem loop_ok {s : State} (hs0 : s.gpr .rdi = s₀.gpr .rdi) (hs10 : s.gpr .r10 = s₀.gpr .rdx)
+ (hsrd : s.rd = s₀.rd) (hswr : s.wr = s₀.wr) (hsm : s.mem = s₀.mem) :
+ WP isa (.seq (.block (yC g)) (VG.Impl.MlKem.X86_64.rcxLoop 32 (bitsBodyY x d))) s fun s' =>
+ Frame [pR (s₀.gpr .rdx)] s₀.mem s'.mem ∧
+ ∀ k < 256, coeffAt s'.mem (s₀.gpr .rdx) k = L (coeffAt s₀.mem (s₀.gpr .rdi) k) := by
+ refine WP.seq (WP.mono (yC_ok g s) fun w ⟨yc, k1, m1, _, _⟩ => ?_)
+ refine WP.mono (wp_rcxLoopY (N := 32) (by decide) (by decide) _ (fun u o hy _ =>
+ ⟨by rw [o.keep.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero, hs0],
+ by rw [o.keep.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero, hs10],
+ by rw [o.keep.2.1, k1.2.1, hsrd], by rw [o.keep.2.2, k1.2.2, hswr],
+ fun l hl => by
+ simp only [State.lane]; rw [o.xmm, hy]; exact yc l hl,
+ by rw [o.mem, m1, hsm]; exact Frame.refl _ _, fun k _ => by rw [o.mem, m1, hsm, ifn (by omega)]⟩)
+ fun i hi u hI => step hrd hwr hdis hX hY hi hI) fun u hI => ⟨hI.frame, fun k hk => by
+ rw [hI.coeff k hk, ifp (by omega)]⟩
+
+end
+
+end YMap
+
+end VG.Proof.MlDsa.X86_64.Arith
+
+namespace VG.Proof.MlDsa.X86_64.Round
+
+open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round VG.Proof.MlDsa.Round
+open VG.Spec.MlDsa
+open VG.Proof.MlKem.X86_64 (Keep Keep.gpr WP.keep gprPreserved_of XOnly)
+open VG.Impl.MlKem.X86_64 (toY)
+open VG.Proof.MlDsa.X86_64.Arith (qV)
+
+section
+variable {post : State → State → Prop} {s₀ : State} (hp : (bitsK post).pre s₀)
+include hp
+
+/-- The prologue, the loop of the `γ₂` it compared, and the epilogue: `out` holds `L γ₂` of each
+coefficient of `r`. -/
+theorem ybits_ok {x : Nat → List Instr} {d : XReg} {L : Nat → BitVec 32 → BitVec 32}
+ (hX : ∀ g, (g = g32 ∨ g = g88) → ∀ t : State, HbC g t → t.xmm .xmm15 = qV →
+ WP isa (.block (x g)) t fun t' => (∀ e < 4, dword (t'.xmm d) e = L g (dword (t.xmm .xmm0) e)) ∧
+ XOnly [.xmm0, .xmm1, .xmm2, .xmm3] t t')
+ (hY : ∀ g, (g = g32 ∨ g = g88) → laneSseBlock (toY (x g)) = some (x g)) :
+ WP isa (bitsY x d) s₀ fun s' =>
+ (∀ k < 256, coeffAt s'.mem (s₀.gpr .rdx) k = L (arg32 s₀ .rsi) (coeffAt s₀.mem (s₀.gpr .rdi) k)) ∧
+ Frame [pR (s₀.gpr .rdx)] s₀.mem s'.mem := by
+ unfold bitsY
+ refine WP.seq (WP.mono (prologue_rsi_rdx s₀) fun s₁ ⟨⟨h10, hz, hm⟩, hk⟩ => ?_)
+ have h0 : s₁.gpr .rdi = s₀.gpr .rdi := hk.gpr (by decide)
+ have go : ∀ g, arg32 s₀ .rsi = g → (g = g32 ∨ g = g88) →
+ WP isa (.seq (.block (yC g)) (VG.Impl.MlKem.X86_64.rcxLoop 32 (bitsBodyY (x g) d))) s₁ fun s' =>
+ (∀ k < 256, coeffAt s'.mem (s₀.gpr .rdx) k = L (arg32 s₀ .rsi) (coeffAt s₀.mem (s₀.gpr .rdi) k)) ∧
+ Frame [pR (s₀.gpr .rdx)] s₀.mem s'.mem := fun g hge hg =>
+ WP.mono (Arith.YMap.loop_ok hp.1 hp.2.1 hp.2.2.1 (hX g hg) (hY g hg) h0 h10 hk.2.1 hk.2.2 hm)
+ fun _ ⟨hf, hc⟩ => ⟨fun k hk => by rw [hc k hk, hge], hf⟩
+ have hite : WP isa (.ite .e (.seq (.block (yC g32)) (VG.Impl.MlKem.X86_64.rcxLoop 32 (bitsBodyY (x g32) d)))
+ (.seq (.block (yC g88)) (VG.Impl.MlKem.X86_64.rcxLoop 32 (bitsBodyY (x g88) d)))) s₁ fun s' =>
+ (∀ k < 256, coeffAt s'.mem (s₀.gpr .rdx) k = L (arg32 s₀ .rsi) (coeffAt s₀.mem (s₀.gpr .rdi) k)) ∧
+ Frame [pR (s₀.gpr .rdx)] s₀.mem s'.mem := by
+ refine WP.ite (M := isa) _ (show isa.eval .e s₁ = _ from hz) (fun h => ?_) (fun h => ?_)
+ · rw [sub_beq_zero32, decide_eq_true_eq] at h
+ exact go _ ((gamma_cases hp.2.2.2.2.2.1).1 h) (.inl rfl)
+ · rw [sub_beq_zero32, decide_eq_false_iff_not] at h
+ exact go _ ((gamma_cases hp.2.2.2.2.2.1).2 h) (.inr rfl)
+ refine WP.seq (WP.mono hite fun u ⟨hc, hf⟩ => ?_)
+ refine WP.mono (Q := fun (u' : State) => u'.mem = u.mem) (by vrund; rfl)
+ fun u' hm' => ?_
+ rw [hm']
+ exact ⟨hc, hf⟩
+
+end
+
+theorem hbX_okG (g : Nat) (hg : g = g32 ∨ g = g88) (t : State) (hc : HbC g t) (_ : t.xmm .xmm15 = qV) :
+ WP isa (.block (hbX g)) t fun t' => (∀ e < 4, dword (t'.xmm .xmm0) e = hbL g (dword (t.xmm .xmm0) e)) ∧
+ XOnly [.xmm0, .xmm1, .xmm2, .xmm3] t t' :=
+ WP.mono (hbX_ok hg t hc) fun _ ⟨h1, h2⟩ => ⟨h1, h2.mono (by simp)⟩
+
+theorem lane_hbX : ∀ g, (g = g32 ∨ g = g88) → laneSseBlock (toY (hbX g)) = some (hbX g) := by
+ intro g hg; rcases hg with rfl | rfl <;> decide +kernel
+
+theorem lane_lbX : ∀ g, (g = g32 ∨ g = g88) → laneSseBlock (toY (lbX g)) = some (lbX g) := by
+ intro g hg; rcases hg with rfl | rfl <;> decide +kernel
+
+theorem highBitsY_correct (s₀ : State) (hp : highBitsK.pre s₀) :
+ ∃ t s', Exec isa highBitsAvx2 s₀ t s' ∧ abiPreserved s₀ s' ∧ highBitsK.post s₀ s' := by
+ obtain ⟨t, s', he, ⟨hv, hf⟩, hk⟩ := WP.keep [.rax, .rcx, .rsi, .rdi, .r10]
+ (ybits_ok hp (L := hbL) hbX_okG lane_hbX) (by decide +kernel)
+ have hr : Reduced s₀.mem (s₀.gpr .rdi) := hp.2.2.2.2.2.2
+ refine ⟨t, s', he, abiPreserved_of_exec (by decide +kernel) he (gprPreserved_of hk (by decide) hf ?_), ?_⟩
+ · simpa using hp.2.2.2.2.1
+ · refine natPolyIs_of_toNat fun k hk => ?_
+ rw [hv k hk, map_get _ _ hk, hbL_toNat hp.2.2.2.2.2.1 (hr k hk), highBits_eq hp.2.2.2.2.2.1, polyAt_val hr hk]
+ exact (Int.toNat_natCast _).symm
+
+theorem lowBitsY_correct (s₀ : State) (hp : lowBitsK.pre s₀) :
+ ∃ t s', Exec isa lowBitsAvx2 s₀ t s' ∧ abiPreserved s₀ s' ∧ lowBitsK.post s₀ s' := by
+ obtain ⟨t, s', he, ⟨hv, hf⟩, hk⟩ := WP.keep [.rax, .rcx, .rsi, .rdi, .r10]
+ (ybits_ok hp (L := lbL) (fun g hg t hc hq => lbX_ok hg t hc hq) lane_lbX) (by decide +kernel)
+ have hr : Reduced s₀.mem (s₀.gpr .rdi) := hp.2.2.2.2.2.2
+ refine ⟨t, s', he, abiPreserved_of_exec (by decide +kernel) he (gprPreserved_of hk (by decide) hf ?_), ?_⟩
+ · simpa using hp.2.2.2.2.1
+ · refine polyIs_of_toNat fun k hk => ?_
+ rw [hv k hk, map_get _ _ hk, lbL_toNat hp.2.2.2.2.2.1 (hr k hk), polyAt_get _ _ hk]
+
+theorem highBitsY_ct : ConstantTime isa highBitsK.pre highBitsK.pub highBitsAvx2 :=
+ VG.Taint.constantTime (A := X86_64.taint) bitsτ bits_agree (by taint_decide)
+
+theorem lowBitsY_ct : ConstantTime isa lowBitsK.pre lowBitsK.pub lowBitsAvx2 :=
+ VG.Taint.constantTime (A := X86_64.taint) bitsτ bits_agree (by taint_decide)
+
+theorem highBitsY_verified : Verified X86_64.target highBitsAvx2 (highBitsContract X86_64.abi) :=
+ Verified.of_correct highBitsY_correct highBitsY_ct (by
+ round_implies [highBitsContract, bitsSig, highBitsK, bitsK, X86_64.abi, X86_64.argRegs] [bitsSat]
+ using bitsSat)
+
+theorem lowBitsY_verified : Verified X86_64.target lowBitsAvx2 (lowBitsContract X86_64.abi) :=
+ Verified.of_correct lowBitsY_correct lowBitsY_ct (by
+ round_implies [lowBitsContract, bitsSig, lowBitsK, bitsK, X86_64.abi, X86_64.argRegs] [bitsSat]
+ using bitsSat)
+
+end VG.Proof.MlDsa.X86_64.Round
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBlock.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBlock.lean
new file mode 100644
index 000000000..bc93dce0c
--- /dev/null
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YBlock.lean
@@ -0,0 +1,69 @@
+import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YLane
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.VLanes
+
+/-!
+# ML-DSA on x86-64: the SSE2 code of the AVX2 rounding, on a register
+
+Untrusted: everything here is checked by Lean. `hbX` and `lbX` compute
+`hbL` and `lbL` in each doubleword of `xmm0` (`hbX_ok`, `lbX_ok`), with
+the constants in `xmm8`, `xmm9`, `xmm10` (and `q` in `xmm15`).
+-/
+
+namespace VG.Proof.MlDsa.X86_64.Round
+
+open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round
+open VG.Proof.MlDsa.X86_64.Arith (dword_psubd dword_pand dword_psrad caddL caddV dword_caddV qV
+ dword_qV)
+open VG.Proof.MlKem.X86_64 (XOnly)
+open VG.Impl.MlKem.X86_64 (xb xmov)
+open VG.Spec.MlDsa (gamma2s)
+
+/-- Each doubleword of `x` is `v`. -/
+def Bc (x : BitVec 128) (v : BitVec 32) : Prop := ∀ e < 4, dword x e = v
+
+theorem bc_ofDwords (v : BitVec 32) : Bc (ofDwords v v v v) v := fun e he => by
+ rcases cases4 he with rfl | rfl | rfl | rfl <;> simp
+
+/-- The constants of `hbX`. -/
+structure HbC (g : Nat) (s : State) : Prop where
+ c8 : Bc (s.xmm .xmm8) (BitVec.ofNat 32 127)
+ c9 : Bc (s.xmm .xmm9) (BitVec.ofNat 32 (dAdd g))
+ c10 : Bc (s.xmm .xmm10) (BitVec.ofNat 32 (dMod g))
+
+theorem dSh_32 : dSh g32 = [10] := rfl
+theorem dSh_88 : dSh g88 = [1, 3, 10, 11, 13] := rfl
+theorem dShift_32 : dShift g32 = 22 := rfl
+theorem dShift_88 : dShift g88 = 24 := rfl
+
+theorem hbX_ok {g : Nat} (hg : g = g32 ∨ g = g88) (s : State) (hc : HbC g s) :
+ WP isa (.block (hbX g)) s fun s' => (∀ e < 4, dword (s'.xmm .xmm0) e = hbL g (dword (s.xmm .xmm0) e)) ∧
+ XOnly [.xmm0, .xmm1, .xmm2] s s' := by
+ rcases hg with rfl | rfl <;>
+ · simp only [hbX, mulX, dSh_32, dSh_88, dShift_32, dShift_88, xmov, xb, List.cons_append, List.nil_append,
+ List.flatMap_cons, List.flatMap_nil, List.append_nil]
+ vrun [VG.X86_64.eval_movdqa]
+ refine ⟨fun e he => ?_, by xonly⟩
+ simp (disch := first | decide | assumption) only [dword_pand, dword_psrad, dword_psubd, dword_psrld,
+ dword_pslld, dword_paddd, hc.c8 e he, hc.c9 e he, hc.c10 e he, BitVec.toNat_ofNat]
+ rfl
+
+theorem mul2X_32 : mul2X g32 = [xmov .xmm1 .xmm0, .xop (.shift .pslld .xmm1 19), xmov .xmm2 .xmm0,
+ .xop (.shift .pslld .xmm2 9), xb .psubd .xmm1 .xmm2] := rfl
+
+theorem mul2X_88 : mul2X g88 = [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] := rfl
+
+theorem lbX_ok {g : Nat} (hg : g = g32 ∨ g = g88) (s : State) (hc : HbC g s) (hq : s.xmm .xmm15 = qV) :
+ WP isa (.block (lbX g)) s fun s' => (∀ e < 4, dword (s'.xmm .xmm3) e = lbL g (dword (s.xmm .xmm0) e)) ∧
+ XOnly [.xmm0, .xmm1, .xmm2, .xmm3] s s' := by
+ rcases hg with rfl | rfl <;>
+ · simp only [lbX, hbX, mulX, mul2X_32, mul2X_88, dSh_32, dSh_88, dShift_32, dShift_88,
+ VG.Impl.MlDsa.X86_64.Arith.vcadd, xmov, xb, List.cons_append, List.nil_append, List.flatMap_cons,
+ List.flatMap_nil, List.append_nil]
+ vrun [VG.X86_64.eval_movdqa]
+ refine ⟨fun e he => ?_, by xonly⟩
+ simp (disch := first | decide | assumption) only [dword_pand, dword_psrad, dword_psubd, dword_psrld,
+ dword_pslld, dword_paddd, hc.c8 e he, hc.c9 e he, hc.c10 e he, BitVec.toNat_ofNat, hq, dword_qV he]
+ rfl
+
+end VG.Proof.MlDsa.X86_64.Round
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YHint.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YHint.lean
new file mode 100644
index 000000000..9a5d3d831
--- /dev/null
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YHint.lean
@@ -0,0 +1,508 @@
+import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YNorm
+import VerifiedGarbage.Proof.MlDsa.X86_64.Round.MakeHint
+
+/-!
+# ML-DSA on x86-64: `vg_mldsa_make_hint_avx2`
+
+Untrusted: everything here is checked by Lean. In each doubleword, `mhX`
+computes the hint bit (`mhL`, `mhL_toNat`); the loop stores the eight
+hints of an iteration and adds their count, the sum of the nibbles of the
+byte mask of the hints shifted to bit 7 (`nib_count`), to `r9`.
+-/
+
+namespace VG.Proof.MlDsa.X86_64.Round
+
+open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round VG.Proof.MlDsa.Round
+open VG.Spec.MlDsa
+open VG.Proof.MlKem.X86_64 (Keep XOnly YOnly ylanes ifp ifn WP.keep)
+open VG.Impl.MlKem.X86_64 (xb xmov toY)
+open VG.Proof.MlDsa.X86_64.Arith (qV csubL csubL_toNat dword_csubV bc)
+
+/-! ## A doubleword -/
+
+/-- The hint bit `mhX` computes from `z` and `r`. -/
+def mhL (g : Nat) (z r : BitVec 32) : BitVec 32 :=
+ ((hbL g (csubL (r + z)) ^^^ hbL g r) + BitVec.ofNat 32 63) >>> 6
+
+theorem mhL_toNat {g : Nat} (h : g ∈ gamma2s) {z r : BitVec 32} (hz : z.toNat < q) (hr : r.toNat < q) :
+ (mhL g z r).toNat = (makeHint g (Fin.ofNat q z.toNat) (Fin.ofNat q r.toNat)).toNat := by
+ have e1 : (r + z).toNat = r.toNat + z.toNat := by
+ rw [BitVec.toNat_add]; rw [q_eq] at hz hr; omega
+ have e2 : (csubL (r + z)).toNat = (r.toNat + z.toNat) % q := by
+ rw [csubL_toNat (by rw [e1]; omega), e1, VG.Proof.MlDsa.Arith.condSub]
+ split <;> rw [q_eq] at * <;> omega
+ have hs : (csubL (r + z)).toNat < q := by rw [e2]; exact Nat.mod_lt _ (by decide)
+ have hM : hbM g ≤ 44 ∧ 0 < hbM g := by rcases mem_gamma2s h with rfl | rfl <;> decide
+ have hu := Nat.mod_lt (hbF g ((r.toNat + z.toNat) % q)) hM.2
+ have hv := Nat.mod_lt (hbF g r.toNat) hM.2
+ have hx : hbF g ((r.toNat + z.toNat) % q) % hbM g ^^^ hbF g r.toNat % hbM g < 2 ^ 6 :=
+ Nat.xor_lt_two_pow (by omega) (by omega)
+ have hvz : (Fin.ofNat q z.toNat).val = z.toNat := Nat.mod_eq_of_lt hz
+ have hvr : (Fin.ofNat q r.toNat).val = r.toNat := Nat.mod_eq_of_lt hr
+ rw [makeHint_eq h, hvz, hvr, mhL, BitVec.toNat_ushiftRight, BitVec.toNat_add, BitVec.toNat_xor,
+ hbL_toNat h hs, hbL_toNat h hr, e2, BitVec.toNat_ofNat, Nat.shiftRight_eq_div_pow]
+ by_cases e : hbF g r.toNat % hbM g = hbF g ((r.toNat + z.toNat) % q) % hbM g
+ · rw [decide_eq_false (fun h' => h' e), e, Nat.xor_self]; rfl
+ · rw [decide_eq_true e]
+ have : hbF g ((r.toNat + z.toNat) % q) % hbM g ^^^ hbF g r.toNat % hbM g ≠ 0 :=
+ fun h' => e (xor_eq_zero h').symm
+ show _ = 1
+ omega
+
+theorem mhL_le {g : Nat} (h : g ∈ gamma2s) {z r : BitVec 32} (hz : z.toNat < q) (hr : r.toNat < q) :
+ (mhL g z r).toNat ≤ 1 := by
+ rw [mhL_toNat h hz hr]; exact Bool.toNat_le _
+
+/-! ## The code on a register -/
+
+theorem mhMid_ok (s : State) (hq : s.xmm .xmm15 = qV) :
+ WP isa (.block mhMid) s fun s' => s'.xmm .xmm3 = s.xmm .xmm0 ∧
+ (∀ e < 4, dword (s'.xmm .xmm0) e = csubL (dword (s.xmm .xmm4) e + dword (s.xmm .xmm5) e)) ∧
+ XOnly [.xmm0, .xmm1, .xmm3] s s' := by
+ simp only [mhMid, VG.Impl.MlDsa.X86_64.Arith.vcsub, VG.Impl.MlDsa.X86_64.Arith.vcadd, xmov, xb, List.cons_append,
+ List.nil_append]
+ vrun [VG.X86_64.eval_movdqa]
+ refine ⟨trivial, fun e he => ?_, by xonly⟩
+ rw [hq, ← dword_paddd _ _ he, ← dword_csubV _ he]
+ rfl
+
+theorem mhTail_ok (s : State) (h11 : Bc (s.xmm .xmm11) (BitVec.ofNat 32 63)) :
+ WP isa (.block mhTail) s fun s' => (∀ e < 4,
+ dword (s'.xmm .xmm0) e = ((dword (s.xmm .xmm0) e ^^^ dword (s.xmm .xmm3) e) + BitVec.ofNat 32 63) >>> 6 ∧
+ dword (s'.xmm .xmm1) e = (((dword (s.xmm .xmm0) e ^^^ dword (s.xmm .xmm3) e) + BitVec.ofNat 32 63) >>> 6) <<< 7) ∧
+ XOnly [.xmm0, .xmm1] s s' := by
+ simp only [mhTail, xmov, xb]
+ vrun [VG.X86_64.eval_movdqa]
+ refine ⟨fun e he => ?_, by xonly⟩
+ simp (disch := first | decide | assumption) only [dword_pxor, dword_paddd, dword_psrld, dword_pslld, h11 e he]
+ exact ⟨rfl, rfl⟩
+
+theorem mhX_ok {g : Nat} (hg : g = g32 ∨ g = g88) (s : State) (hc : HbC g s) (hq : s.xmm .xmm15 = qV)
+ (h11 : Bc (s.xmm .xmm11) (BitVec.ofNat 32 63)) :
+ WP isa (.block (mhX g)) s fun s' => (∀ e < 4,
+ dword (s'.xmm .xmm0) e = mhL g (dword (s.xmm .xmm5) e) (dword (s.xmm .xmm0) e) ∧
+ dword (s'.xmm .xmm1) e = mhL g (dword (s.xmm .xmm5) e) (dword (s.xmm .xmm0) e) <<< 7) ∧
+ XOnly [.xmm0, .xmm1, .xmm2, .xmm3, .xmm4] s s' := by
+ rw [mhX, List.append_assoc, List.append_assoc, List.append_assoc, WP.block_append_iff]
+ refine WP.mono (Q := fun (s1 : State) => s1.xmm .xmm4 = s.xmm .xmm0 ∧ XOnly [.xmm4] s s1)
+ (by simp only [xmov, xb]; vrun [VG.X86_64.eval_movdqa]; exact ⟨trivial, by xonly⟩) fun s1 ⟨a1, o1⟩ => ?_
+ have hc1 : HbC g s1 := ⟨by rw [o1.xmm _ (by decide)]; exact hc.c8, by rw [o1.xmm _ (by decide)]; exact hc.c9,
+ by rw [o1.xmm _ (by decide)]; exact hc.c10⟩
+ rw [WP.block_append_iff]
+ refine WP.mono (hbX_ok hg s1 hc1) fun s2 ⟨a2, o2⟩ => ?_
+ rw [WP.block_append_iff]
+ refine WP.mono (mhMid_ok s2 (by rw [o2.xmm _ (by decide), o1.xmm _ (by decide), hq])) fun s3 ⟨b3, a3, o3⟩ => ?_
+ have hc3 : HbC g s3 := ⟨by rw [o3.xmm _ (by decide), o2.xmm _ (by decide)]; exact hc1.c8,
+ by rw [o3.xmm _ (by decide), o2.xmm _ (by decide)]; exact hc1.c9,
+ by rw [o3.xmm _ (by decide), o2.xmm _ (by decide)]; exact hc1.c10⟩
+ rw [WP.block_append_iff]
+ refine WP.mono (hbX_ok hg s3 hc3) fun s4 ⟨a4, o4⟩ => ?_
+ refine WP.mono (mhTail_ok s4 (by
+ rw [o4.xmm _ (by decide), o3.xmm _ (by decide), o2.xmm _ (by decide), o1.xmm _ (by decide)]; exact h11))
+ fun s5 ⟨a5, o5⟩ => ⟨fun e he => ?_, ?_⟩
+ · have hz : dword (s4.xmm .xmm0) e = hbL g (csubL (dword (s.xmm .xmm0) e + dword (s.xmm .xmm5) e)) := by
+ rw [a4 e he, a3 e he, o2.xmm _ (by decide), a1, o2.xmm _ (by decide), o1.xmm _ (by decide)]
+ have hr1 : dword (s4.xmm .xmm3) e = hbL g (dword (s.xmm .xmm0) e) := by
+ rw [o4.xmm _ (by decide), b3, a2 e he, o1.xmm _ (by decide)]
+ rw [(a5 e he).1, (a5 e he).2, hz, hr1]
+ exact ⟨rfl, rfl⟩
+ · exact ((((o1.trans o2).trans o3).trans o4).trans o5).mono (by simp)
+
+/-! ## The count -/
+
+/-- What `cntH` leaves in `eax`: the sum of the nibbles of `x`, if it fits in one. -/
+def nib (x : BitVec 32) : BitVec 32 :=
+ let x1 := x + x >>> 4
+ let x2 := x1 + x1 >>> 8
+ (x2 + x2 >>> 16) &&& 15
+
+/-- The eight bits `b d` at bits `4d`. -/
+def sp (b : Nat → Bool) : Nat → Nat
+ | 0 => 0
+ | k + 1 => sp b k + (b k).toNat * 16 ^ k
+
+/-- The number of the first `k` bits set. -/
+def cnt8 (b : Nat → Bool) : Nat → Nat
+ | 0 => 0
+ | k + 1 => cnt8 b k + (b k).toNat
+
+theorem nib_bools : ∀ b0 b1 b2 b3 b4 b5 b6 b7 : Bool,
+ (nib (BitVec.ofNat 32 (b0.toNat + b1.toNat * 16 + b2.toNat * 16 ^ 2 + b3.toNat * 16 ^ 3 + b4.toNat * 16 ^ 4 +
+ b5.toNat * 16 ^ 5 + b6.toNat * 16 ^ 6 + b7.toNat * 16 ^ 7))).toNat =
+ b0.toNat + b1.toNat + b2.toNat + b3.toNat + b4.toNat + b5.toNat + b6.toNat + b7.toNat := by
+ decide +kernel
+
+theorem nib_sp (b : Nat → Bool) : (nib (BitVec.ofNat 32 (sp b 8))).toNat = cnt8 b 8 := by
+ have := nib_bools (b 0) (b 1) (b 2) (b 3) (b 4) (b 5) (b 6) (b 7)
+ simp only [sp, cnt8, Nat.pow_zero, Nat.mul_one, Nat.zero_add, Nat.pow_one] at this ⊢
+ exact this
+
+/-- A byte mask with bit `4d` the bit `b d`, and the others clear. -/
+theorem bsum_sp {f : Nat → Bool} {b : Nat → Bool} : ∀ k, (∀ j < 4 * k, f j = (decide (j % 4 = 0) && b (j / 4))) →
+ VG.Proof.MlKem.X86_64.S4.bsum f (4 * k) = sp b k
+ | 0, _ => rfl
+ | k + 1, h => by
+ have ih := bsum_sp k fun j hj => h j (by omega)
+ rw [show 4 * (k + 1) = 4 * k + 1 + 1 + 1 + 1 by omega]
+ simp only [VG.Proof.MlKem.X86_64.S4.bsum, sp]
+ rw [ih, h _ (by omega), h _ (by omega), h _ (by omega), h _ (by omega)]
+ simp only [show (4 * k) % 4 = 0 by omega, show (4 * k + 1) % 4 = 1 by omega, show (4 * k + 2) % 4 = 2 by omega,
+ show (4 * k + 1 + 1 + 1) % 4 = 3 by omega, show 4 * k / 4 = k by omega,
+ decide_true, decide_false, Bool.true_and, Bool.false_and, Bool.toNat_false, Nat.zero_mul, Nat.add_zero,
+ show (1 : Nat) ≠ 0 by decide, show (2 : Nat) ≠ 0 by decide, show (3 : Nat) ≠ 0 by decide]
+ rw [show 2 ^ (4 * k) = 16 ^ k by rw [Nat.pow_mul]]
+
+theorem cntH_ok (s : State) :
+ WP isa (.block cntH) s fun s' =>
+ (s'.gpr .r9 = s.gpr .r9 + BitVec.setWidth 64 (nib ((s.gpr .rax).setWidth 32)) ∧ s'.mem = s.mem ∧
+ ∀ r l, s'.lane r l = s.lane r l) ∧ Keep [.rax, .rdx, .r9] s s' := by
+ refine WP.keep _ ?_ (by decide)
+ simp only [cntH]
+ xrun
+ exact ⟨rfl, fun _ _ => rfl⟩
+
+end VG.Proof.MlDsa.X86_64.Round
+
+namespace VG.Proof.MlDsa.X86_64.Arith
+
+open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round
+open VG.Proof.MlDsa.Arith
+open VG.Proof.MlDsa.X86_64.Round (HbC Bc bc_ofDwords mhL mhL_le mhX_ok nib sp cnt8 nib_sp bsum_sp ymm_bit cntH_ok
+ sw3264)
+open VG.Proof.MlKem.X86_64 (Keep XOnly YOnly ylanes yld_ok yconst_ok wp_rcxLoopY ifp ifn ptr_step add_ofNat_zero
+ lane_setReg lane_setFlags sx32 State.setMem_ymm)
+open VG.Impl.MlKem.X86_64 (toY)
+open VG.Spec.MlDsa (q n coeffAt Reduced gamma2s)
+
+/-- `Σ k < m, f k`. -/
+def csum (f : Nat → Nat) : Nat → Nat
+ | 0 => 0
+ | m + 1 => csum f m + f m
+
+namespace YHintL
+
+/-- After `i` vectors of eight. -/
+structure Inv (s₀ : State) (g : Nat) (i : Nat) (s : State) : Prop where
+ rdi : s.gpr .rdi = s₀.gpr .rdi + BitVec.ofNat 64 (32 * i)
+ rsi : s.gpr .rsi = s₀.gpr .rsi + BitVec.ofNat 64 (32 * i)
+ r10 : s.gpr .r10 = s₀.gpr .rcx + BitVec.ofNat 64 (32 * i)
+ rd : s.rd = s₀.rd
+ wr : s.wr = s₀.wr
+ yc : YC g s
+ c11 : ∀ l < 2, s.lane .xmm11 l = bc (BitVec.ofNat 32 63)
+ frame : Frame [pR (s₀.gpr .rcx)] s₀.mem s.mem
+ coeff : ∀ k < 256, coeffAt s.mem (s₀.gpr .rcx) k = if k < 8 * i then
+ mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k) else coeffAt s₀.mem (s₀.gpr .rcx) k
+ r9 : (s.gpr .r9).toNat =
+ csum (fun k => (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat) (8 * i)
+
+theorem lane_mhX (g : Nat) (hg : g = g32 ∨ g = g88) : laneSseBlock (toY (mhX g)) = some (mhX g) := by
+ rcases hg with rfl | rfl <;> decide +kernel
+
+/-- The constants are kept by code that writes only `xmm0` to `xmm5`. -/
+theorem keep_consts {g : Nat} {s s' : State} (hyc : YC g s) (h11 : ∀ l < 2, s.lane .xmm11 l = bc (BitVec.ofNat 32 63))
+ {rs : List XReg} (hr : ∀ r ∈ rs, r ∈ [.xmm0, .xmm1, .xmm2, .xmm3, .xmm4, .xmm5])
+ (hl : ∀ r ∉ rs, ∀ l < 2, s'.lane r l = s.lane r l) :
+ YC g s' ∧ ∀ l < 2, s'.lane .xmm11 l = bc (BitVec.ofNat 32 63) := by
+ have n : ∀ r ∈ [XReg.xmm8, .xmm9, .xmm10, .xmm11, .xmm15], r ∉ rs := fun r hr' hr'' => by
+ have := hr r hr''; simp only [List.mem_cons, List.not_mem_nil, or_false] at hr' this
+ rcases hr' with rfl | rfl | rfl | rfl | rfl <;> rcases this with h | h | h | h | h | h <;> cases h
+ refine ⟨fun l hl' => ?_, fun l hl' => by rw [hl _ (n _ (by simp)) l hl']; exact h11 l hl'⟩
+ rw [hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl', hl _ (n _ (by simp)) l hl']
+ exact hyc l hl'
+
+/-- A hint shifted to bit 7: bit `8m + 7` is the hint's bit 0 if `m = 0`, and clear otherwise. -/
+theorem shl7_bit {v : BitVec 32} (hv : v.toNat ≤ 1) {m : Nat} (hm : m < 4) :
+ (v <<< 7).getLsbD (8 * m + 7) = (decide (m = 0) && v.getLsbD 0) := by
+ rw [BitVec.getLsbD_shiftLeft]
+ rcases (by omega : m = 0 ∨ 0 < m) with rfl | h
+ · simp
+ · have : v.getLsbD (8 * m) = false := by
+ rw [← BitVec.testBit_toNat, Nat.testBit_lt_two_pow (Nat.lt_of_le_of_lt hv
+ (Nat.one_lt_two_pow (by omega)))]
+ simp [show 8 * m + 7 - 7 = 8 * m by omega, this, show m ≠ 0 by omega]
+
+theorem bit0_toNat {v : BitVec 32} (hv : v.toNat ≤ 1) : (v.getLsbD 0).toNat = v.toNat := by
+ rw [← BitVec.testBit_toNat, Nat.testBit_zero]
+ rcases (by omega : v.toNat = 0 ∨ v.toNat = 1) with h | h <;> rw [h] <;> rfl
+
+theorem csum_le {f : Nat → Nat} : ∀ {m : Nat}, (∀ k < m, f k ≤ 1) → csum f m ≤ m
+ | 0, _ => Nat.le_refl _
+ | m + 1, h => by
+ have := csum_le (m := m) fun k hk => h k (by omega)
+ have := h m (by omega)
+ simp only [csum]; omega
+
+theorem csum8 (f : Nat → Nat) (b : Nat → Bool) (m : Nat) (h : ∀ d < 8, (b d).toNat = f (m + d)) :
+ csum f (m + 8) = csum f m + cnt8 b 8 := by
+ rw [show csum f (m + 8) = csum f m + f m + f (m + 1) + f (m + 2) + f (m + 3) + f (m + 4) + f (m + 5) +
+ f (m + 6) + f (m + 7) from rfl,
+ show cnt8 b 8 = (b 0).toNat + (b 1).toNat + (b 2).toNat + (b 3).toNat + (b 4).toNat + (b 5).toNat +
+ (b 6).toNat + (b 7).toNat by simp [cnt8],
+ h 0 (by decide), h 1 (by decide), h 2 (by decide), h 3 (by decide), h 4 (by decide), h 5 (by decide),
+ h 6 (by decide), h 7 (by decide)]
+ simp only [Nat.add_zero]
+ omega
+
+section
+variable {s₀ : State} (hrd : s₀.rd = [pR (s₀.gpr .rdi), pR (s₀.gpr .rsi)]) (hwr : s₀.wr = [pR (s₀.gpr .rcx)])
+ (hdz : (pR (s₀.gpr .rdi)).Disjoint (pR (s₀.gpr .rcx))) (hdr : (pR (s₀.gpr .rsi)).Disjoint (pR (s₀.gpr .rcx)))
+ (hz : Reduced s₀.mem (s₀.gpr .rdi)) (hr : Reduced s₀.mem (s₀.gpr .rsi)) {g : Nat} (hg : g = g32 ∨ g = g88)
+include hrd hwr hdz hdr hz hr hg
+
+theorem step {i : Nat} (hi : i < 32) {s : State} (hI : Inv s₀ g i s) :
+ WP isa (.block (mhBodyY g ++ ([.alu .sub .rcx (.imm 1)] : List Instr))) s fun s' =>
+ Inv s₀ g (i + 1) s' ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ s'.zf = some (s.gpr .rcx - 1 == 0) := by
+ have hg' : g ∈ gamma2s := by rcases hg with rfl | rfl <;> decide
+ have j0 : 8 * i + 8 ≤ 256 := by omega
+ have hw : pR (s₀.gpr .rcx) ∈ s.wr := by rw [hI.wr, hwr]; simp
+ have hrz : pR (s₀.gpr .rdi) ∈ s.rd ++ s.wr := by rw [hI.rd, hrd]; simp
+ have hrr : pR (s₀.gpr .rsi) ∈ s.rd ++ s.wr := by rw [hI.rd, hrd]; simp
+ have e1 : s.gpr .rsi + BitVec.ofNat 64 0 = coeffAddr (s₀.gpr .rsi) (8 * i) := by
+ rw [add_ofNat_zero, hI.rsi]; congr 2; omega
+ have e2 : s.gpr .rdi + BitVec.ofNat 64 0 = coeffAddr (s₀.gpr .rdi) (8 * i) := by
+ rw [add_ofNat_zero, hI.rdi]; congr 2; omega
+ have e3 : s.gpr .r10 = coeffAddr (s₀.gpr .rcx) (8 * i) := by rw [hI.r10]; congr 2; omega
+ rw [mhBodyY, List.append_assoc, List.append_assoc, List.append_assoc, List.append_assoc,
+ show ∀ a b : Instr, [a, b] = [a] ++ [b] from fun _ _ => rfl, List.append_assoc, WP.block_append_iff]
+ refine WP.mono (yld_ok (by rw [e1]; exact f_in32 hrr j0)) fun s1 ⟨L1, o1⟩ => ?_
+ rw [WP.block_append_iff]
+ refine WP.mono (yld_ok (by rw [o1.rd, o1.wr, o1.gpr, e2]; exact f_in32 hrz j0)) fun s2 ⟨L2, o2⟩ => ?_
+ have k2 := keep_consts hI.yc hI.c11 (rs := [.xmm0, .xmm5]) (by simp) fun r hr' l hl => by
+ simp only [List.mem_cons, List.not_mem_nil, or_false, not_or] at hr'
+ rw [o2.lane r (by simp [hr'.2]) l hl, o1.lane r (by simp [hr'.1]) l hl]
+ rw [WP.block_append_iff]
+ refine WP.mono (ylanes (lane_mhX g hg) (P := fun l t => ∀ e < 4,
+ dword (t.xmm .xmm0) e = mhL g (dword ((s2.proj l).xmm .xmm5) e) (dword ((s2.proj l).xmm .xmm0) e) ∧
+ dword (t.xmm .xmm1) e = mhL g (dword ((s2.proj l).xmm .xmm5) e) (dword ((s2.proj l).xmm .xmm0) e) <<< 7)
+ fun l hl => mhX_ok hg _ (k2.1.hbc hl) (k2.1.q hl)
+ (by rw [State.proj_xmm, k2.2 l hl]; exact bc_ofDwords _)) fun s3 ⟨B3, o3⟩ => ?_
+ have o13 := (o1.trans o2).trans o3
+ have k3 := keep_consts k2.1 k2.2 (rs := [.xmm0, .xmm1, .xmm2, .xmm3, .xmm4]) (by simp) o3.lane
+ have g3 : s3.gpr .r10 = coeffAddr (s₀.gpr .rcx) (8 * i) := by rw [o13.gpr, e3]
+ have w0 : InRegions s3.wr (s3.gpr .r10) 32 := by rw [o13.wr, g3]; exact f_in32 hw j0
+ -- The values of the lanes.
+ have hv : ∀ l < 2, ∀ e < 4, mhL g (dword ((s2.proj l).xmm .xmm5) e) (dword ((s2.proj l).xmm .xmm0) e) =
+ mhL g (coeffAt s₀.mem (s₀.gpr .rdi) (8 * i + 4 * l + e)) (coeffAt s₀.mem (s₀.gpr .rsi) (8 * i + 4 * l + e)) :=
+ fun l hl e he => by
+ rw [State.proj_xmm, State.proj_xmm, L2 l hl, o2.lane _ (by decide) l hl, L1 l hl, o1.mem, o1.gpr, e1, e2,
+ dword_readW _ _ he, dword_readW _ _ he, lane_load, lane_load, coeffAddr_add, coeffAddr_add, ← coeffAt_eq,
+ ← coeffAt_eq, coeffAt_frame hI.frame (by simpa using hdz) (by rw [n_eq]; omega),
+ coeffAt_frame hI.frame (by simpa using hdr) (by rw [n_eq]; omega)]
+ have hle : ∀ k < 256, (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat ≤ 1 :=
+ fun k hk => mhL_le hg' (hz k (by rw [n_eq]; exact hk)) (hr k (by rw [n_eq]; exact hk))
+ rw [WP.block_append_iff]
+ refine WP.mono (Q := fun (s4 : State) => s4.mem = s3.mem.writeW (s3.gpr .r10) (s3.ymm .xmm0) ∧
+ s4.gpr = (s3.setReg .rax (byteMask (s3.ymm .xmm1) 32)).gpr ∧ s4.rd = s3.rd ∧ s4.wr = s3.wr ∧
+ ∀ r l, s4.lane r l = s3.lane r l)
+ (by
+ vrund [State.store256_eq, State.setMem_gpr, State.setMem_wr, State.setMem_mem, State.setMem_rd,
+ State.setMem_ymm, w0]
+ exact ⟨rfl, fun r l => by simp only [lane_setReg, State.setMem_lane]⟩) fun s4 ⟨m4, g4, r4, w4, l4⟩ => ?_
+ rw [WP.block_append_iff]
+ refine WP.mono (cntH_ok s4) fun s5 ⟨⟨h9, m5, l5⟩, k5⟩ => ?_
+ have hax : s4.gpr .rax = byteMask (s3.ymm .xmm1) 32 := by rw [g4, RegUpd.gpr_setReg_self]
+ -- The count of the eight hints.
+ let b : Nat → Bool := fun d =>
+ (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) (8 * i + d)) (coeffAt s₀.mem (s₀.gpr .rsi) (8 * i + d))).getLsbD 0
+ have hbits : ∀ j < 4 * 8, (s3.ymm .xmm1).getLsbD (8 * j + 7) = (decide (j % 4 = 0) && b (j / 4)) := fun j hj => by
+ have hl : j / 16 < 2 := by omega
+ have he : j % 16 / 4 < 4 := by omega
+ rw [ymm_bit _ _ (by omega), ← State.proj_xmm, (B3 _ hl _ he).2, hv _ hl _ he,
+ shl7_bit (hle _ (by omega)) (by omega)]
+ simp only [b, show 8 * i + 4 * (j / 16) + j % 16 / 4 = 8 * i + j / 4 by omega]
+ have hcnt : (BitVec.setWidth 64 (nib ((s4.gpr .rax).setWidth 32))).toNat = cnt8 b 8 := by
+ rw [hax, VG.Proof.MlKem.X86_64.S4.byteMask_eq _ (by decide), BitVec.toNat_setWidth,
+ show BitVec.setWidth 32 (BitVec.ofNat 64 (VG.Proof.MlKem.X86_64.S4.bsum
+ (fun i => (s3.ymm .xmm1).getLsbD (8 * i + 7)) 32)) =
+ BitVec.ofNat 32 (VG.Proof.MlKem.X86_64.S4.bsum (fun i => (s3.ymm .xmm1).getLsbD (8 * i + 7)) 32) by
+ apply BitVec.eq_of_toNat_eq
+ have := VG.Proof.MlKem.X86_64.S4.bsum_lt (fun i => (s3.ymm .xmm1).getLsbD (8 * i + 7)) 32
+ simp only [BitVec.toNat_setWidth, BitVec.toNat_ofNat]; omega,
+ (bsum_sp 8 hbits : VG.Proof.MlKem.X86_64.S4.bsum _ 32 = _), nib_sp]
+ have : cnt8 b 8 ≤ 8 := by
+ simp only [cnt8]
+ have hb := fun d => Bool.toNat_le (b d)
+ have := hb 0; have := hb 1; have := hb 2; have := hb 3; have := hb 4; have := hb 5
+ have := hb 6; have := hb 7
+ omega
+ have := (nib_sp b).symm ▸ this
+ omega
+ simp only [List.cons_append, List.nil_append]
+ xrun
+ have G5 : ∀ r, r ∉ [Reg.rax, .rdx, .r9] → s5.gpr r = s.gpr r := fun r hr => by
+ rw [k5.gpr hr, g4, RegUpd.gpr_setReg_of_ne _ _ (fun h => hr (by simp [h])), o13.gpr]
+ refine ⟨⟨?_, ?_, ?_, ?_, ?_, ?_, ?_, ?_, fun k hk => ?_, ?_⟩, by rw [G5 .rcx (by simp)], by rw [G5 .rcx (by simp)]⟩
+ · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq]
+ rw [G5 .rdi (by simp), hI.rdi]; exact ptr_step _ i 32
+ · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq]
+ rw [G5 .rsi (by simp), hI.rsi]; exact ptr_step _ i 32
+ · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq]
+ rw [G5 .r10 (by simp), hI.r10]; exact ptr_step _ i 32
+ · simp only [RegUpd.rd_setReg, RegUpd.rd_setFlags]; rw [k5.2.1, r4, o13.rd, hI.rd]
+ · simp only [RegUpd.wr_setReg, RegUpd.wr_setFlags]; rw [k5.2.2, w4, o13.wr, hI.wr]
+ · exact k3.1.keep (rs := []) (by simp) fun r _ l _ => by simp only [lane_setReg, lane_setFlags]; rw [l5, l4]
+ · intro l hl; simp only [lane_setReg, lane_setFlags]; rw [l5, l4]; exact k3.2 l hl
+ · simp only [RegUpd.mem_setReg, RegUpd.mem_setFlags]
+ rw [m5, m4, g3, o13.mem]; exact hI.frame.writeW (List.mem_singleton_self _) _ (pR_contains32 _ j0)
+ · simp only [RegUpd.mem_setReg, RegUpd.mem_setFlags]
+ rw [m5, m4, g3, o13.mem, coeffAt_write256 _ _ j0 _ hk]
+ split
+ · rename_i h
+ rw [ifp (show k < 8 * (i + 1) by omega), State.ymm, extract_ymm _ _ (by omega)]
+ have hc : ∀ l < 2, ∀ e < 4, dword (s3.lane .xmm0 l) e =
+ mhL g (coeffAt s₀.mem (s₀.gpr .rdi) (8 * i + 4 * l + e)) (coeffAt s₀.mem (s₀.gpr .rsi) (8 * i + 4 * l + e)) :=
+ fun l hl e he => by rw [← State.proj_xmm, (B3 l hl e he).1, hv l hl e he]
+ split
+ · rename_i h4
+ have := hc 0 (by decide) (k - 8 * i) h4
+ rw [show 8 * i + 4 * 0 + (k - 8 * i) = k by omega] at this
+ exact this
+ · rename_i h4
+ have := hc 1 (by decide) (k - 8 * i - 4) (by omega)
+ rw [show 8 * i + 4 * 1 + (k - 8 * i - 4) = k by omega] at this
+ exact this
+ · rename_i h
+ rw [hI.coeff k hk]
+ by_cases h' : k < 8 * i
+ · rw [ifp h', ifp (by omega)]
+ · rw [ifn h', ifn (by omega)]
+ · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_false, reduceCtorEq]
+ have h94 : s4.gpr .r9 = s.gpr .r9 := by
+ rw [g4, RegUpd.gpr_setReg_of_ne _ _ (by decide), o13.gpr]
+ have e8 := csum8 (fun k => (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat) b
+ (8 * i) fun d hd => bit0_toNat (hle _ (by omega))
+ have hb := csum_le (f := fun k => (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat)
+ (m := 8 * i + 8) fun k hk => hle k (by omega)
+ rw [h9, BitVec.toNat_add, hcnt, h94, hI.r9, show 8 * (i + 1) = 8 * i + 8 by omega, e8]
+ rw [e8] at hb
+ exact Nat.mod_eq_of_lt (by omega)
+
+/-- The constants and the loop: the hints at `h`, and their count in `r9`. -/
+theorem loop_ok {s : State} (hs0 : s.gpr .rdi = s₀.gpr .rdi) (hs1 : s.gpr .rsi = s₀.gpr .rsi)
+ (hs10 : s.gpr .r10 = s₀.gpr .rcx) (hs9 : s.gpr .r9 = 0) (hsrd : s.rd = s₀.rd) (hswr : s.wr = s₀.wr)
+ (hsm : s.mem = s₀.mem) :
+ WP isa (mhY g) s fun s' => Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧
+ (∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k =
+ mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)) ∧
+ (s'.gpr .r9).toNat =
+ csum (fun k => (mhL g (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat) 256 := by
+ unfold mhY
+ refine WP.seq ?_
+ rw [WP.block_append_iff]
+ refine WP.mono (yC_ok g s) fun w ⟨yc, k1, m1, _, _⟩ =>
+ WP.mono (yconst_ok .xmm11 63 w) fun w2 ⟨l2, k2, m2, _, o2⟩ => ?_
+ have yc2 : YC g w2 := fun l hl => by
+ rw [o2 .xmm8 (by decide) l hl, o2 .xmm9 (by decide) l hl, o2 .xmm10 (by decide) l hl,
+ o2 .xmm15 (by decide) l hl]
+ exact yc l hl
+ refine WP.mono (wp_rcxLoopY (N := 32) (by decide) (by decide) _ (fun u o hy _ =>
+ ⟨by rw [o.keep.gpr (by decide), k2.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero, hs0],
+ by rw [o.keep.gpr (by decide), k2.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero, hs1],
+ by rw [o.keep.gpr (by decide), k2.gpr (by decide), k1.gpr (by decide), Nat.mul_zero, add_ofNat_zero, hs10],
+ by rw [o.keep.2.1, k2.2.1, k1.2.1, hsrd], by rw [o.keep.2.2, k2.2.2, k1.2.2, hswr],
+ fun l hl => by simp only [State.lane]; rw [o.xmm, hy]; exact yc2 l hl,
+ fun l hl => by simp only [State.lane]; rw [o.xmm, hy]; exact l2 l hl,
+ by rw [o.mem, m2, m1, hsm]; exact Frame.refl _ _, fun k _ => by rw [o.mem, m2, m1, hsm, ifn (by omega)],
+ by rw [o.keep.gpr (by decide), k2.gpr (by decide), k1.gpr (by decide), hs9]; rfl⟩)
+ fun i hi u hI => step hrd hwr hdz hdr hz hr hg hi hI) fun u hI => ⟨hI.frame, fun k hk => by
+ rw [hI.coeff k hk, ifp (by omega)], hI.r9⟩
+
+end
+
+end YHintL
+
+end VG.Proof.MlDsa.X86_64.Arith
+
+namespace VG.Proof.MlDsa.X86_64.Round
+
+open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round VG.Proof.MlDsa.Round
+open VG.Spec.MlDsa
+open VG.Proof.MlKem.X86_64 (Keep Keep.gpr WP.keep gprPreserved_of)
+open VG.Proof.MlDsa.X86_64.Arith (csum)
+
+/-- The count of the hints is `onesFrom` them. -/
+theorem csum_onesFrom (v : Vector Bool n) (f : Nat → Nat) (hf : ∀ k < 256, f k = v[k]!.toNat) :
+ ∀ m ≤ 256, csum f m + onesFrom v m = onesFrom v 0
+ | 0, _ => Nat.zero_add _
+ | m + 1, hm => by
+ have ih := csum_onesFrom v f hf m (by omega)
+ rw [onesFrom_step v (by rw [n_eq]; omega), ← hf m (by omega)] at ih
+ simp only [csum]; omega
+
+section
+variable {s₀ : State} (hp : makeHintK.pre s₀)
+include hp
+
+theorem makeHintY_correct : ∃ t s', Exec isa makeHintAvx2 s₀ t s' ∧ abiPreserved s₀ s' ∧ makeHintK.post s₀ s' := by
+ have hg : arg32 s₀ .rdx ∈ gamma2s := hp.2.2.2.2.2.2.2.1
+ have hz : Reduced s₀.mem (s₀.gpr .rdi) := hp.2.2.2.2.2.2.2.2.1
+ have hr : Reduced s₀.mem (s₀.gpr .rsi) := hp.2.2.2.2.2.2.2.2.2
+ let f : Nat → Nat := fun k =>
+ (mhL (arg32 s₀ .rdx) (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)).toNat
+ have wp : WP isa makeHintAvx2 s₀ fun s' => Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧
+ (∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k =
+ mhL (arg32 s₀ .rdx) (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)) ∧
+ (s'.gpr .rax).toNat = csum f 256 := by
+ unfold makeHintAvx2
+ refine WP.seq (WP.mono (mhPrologue_ok s₀) fun s₁ ⟨⟨h10, h9, hzf, hm⟩, hk⟩ => ?_)
+ have go : ∀ g, arg32 s₀ .rdx = g → (g = g32 ∨ g = g88) → WP isa (mhY g) s₁ fun s' =>
+ Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧
+ (∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k =
+ mhL (arg32 s₀ .rdx) (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)) ∧
+ (s'.gpr .r9).toNat = csum f 256 := fun g hge hg' => by
+ subst hge
+ exact Arith.YHintL.loop_ok hp.1 hp.2.1 hp.2.2.1 hp.2.2.2.1 hz hr hg' (hk.gpr (by decide))
+ (hk.gpr (by decide)) h10 h9 hk.2.1 hk.2.2 hm
+ have hite : WP isa (.ite .e (mhY g32) (mhY g88)) s₁ fun s' => Frame [pR (s₀.gpr .rcx)] s₀.mem s'.mem ∧
+ (∀ k < 256, coeffAt s'.mem (s₀.gpr .rcx) k =
+ mhL (arg32 s₀ .rdx) (coeffAt s₀.mem (s₀.gpr .rdi) k) (coeffAt s₀.mem (s₀.gpr .rsi) k)) ∧
+ (s'.gpr .r9).toNat = csum f 256 := by
+ refine WP.ite (M := isa) _ (show isa.eval .e s₁ = _ from hzf) (fun h => ?_) (fun h => ?_)
+ · rw [sub_beq_zero32, decide_eq_true_eq] at h
+ exact go _ ((gamma_cases hg).1 h) (.inl rfl)
+ · rw [sub_beq_zero32, decide_eq_false_iff_not] at h
+ exact go _ ((gamma_cases hg).2 h) (.inr rfl)
+ refine WP.seq (WP.mono hite fun u ⟨hf, hc, h9'⟩ => ?_)
+ refine WP.mono (Q := fun (u' : State) => u'.mem = u.mem ∧ u'.gpr .rax = u.gpr .r9) (by vrund; exact ⟨rfl, rfl⟩)
+ fun u' ⟨hm', hax⟩ => ?_
+ rw [hm', hax]
+ exact ⟨hf, hc, h9'⟩
+ obtain ⟨t, s', he, ⟨hf, hv, hax⟩, hk⟩ := WP.keep [.rax, .rcx, .rdx, .rsi, .rdi, .r9, .r10] wp (by decide +kernel)
+ refine ⟨t, s', he, abiPreserved_of_exec (by decide +kernel) he (gprPreserved_of hk (by decide) hf ?_), ?_, ?_⟩
+ · simpa using hp.2.2.2.2.2.2.1
+ · refine hintIs_of_toNat fun k hk' => ?_
+ apply BitVec.eq_of_toNat_eq
+ rw [hv k hk', mhL_toNat hg (hz k hk') (hr k hk'), zipWith_get _ _ _ hk', polyAt_get _ _ hk', polyAt_get _ _ hk']
+ simp only [BitVec.natCast_eq_ofNat, BitVec.toNat_ofNat]
+ exact (Nat.mod_eq_of_lt (Nat.lt_of_le_of_lt (Bool.toNat_le _) (by decide))).symm
+ · have e := csum_onesFrom (Vector.zipWith (makeHint (arg32 s₀ .rdx)) (polyAt s₀.mem (s₀.gpr .rdi))
+ (polyAt s₀.mem (s₀.gpr .rsi))) f (fun k hk' => by
+ rw [zipWith_get _ _ _ (by rw [n_eq]; exact hk'), polyAt_get _ _ (by rw [n_eq]; exact hk'),
+ polyAt_get _ _ (by rw [n_eq]; exact hk')]
+ exact mhL_toNat hg (hz k (by rw [n_eq]; exact hk')) (hr k (by rw [n_eq]; exact hk'))) 256 (Nat.le_refl _)
+ have e0 : onesFrom (Vector.zipWith (makeHint (arg32 s₀ .rdx)) (polyAt s₀.mem (s₀.gpr .rdi))
+ (polyAt s₀.mem (s₀.gpr .rsi))) 256 = 0 := onesFrom_n _
+ have : (s'.gpr .rax).toNat ≤ 256 := by
+ rw [hax]
+ exact Arith.YHintL.csum_le fun k hk' => mhL_le hg (hz k (by rw [n_eq]; exact hk')) (hr k (by rw [n_eq]; exact hk'))
+ rw [BitVec.toNat_setWidth, hintOnes_single]
+ omega
+
+end
+
+theorem makeHintY_ct : ConstantTime isa makeHintK.pre makeHintK.pub makeHintAvx2 :=
+ VG.Taint.constantTime (A := X86_64.taint) (regsLo [.rdi, .rsi, .rcx, .rsp] [.rdx])
+ (fun _ _ _ _ hp => agree_regsLo (fun r hr => by
+ simp only [List.mem_cons, List.not_mem_nil, or_false] at hr
+ rcases hr with rfl | rfl | rfl | rfl
+ exacts [hp.1, hp.2.1, hp.2.2.1, hp.2.2.2.1]) fun r hr => by
+ simp only [List.mem_singleton] at hr; subst hr; exact hp.2.2.2.2)
+ (by taint_decide)
+
+theorem makeHintY_verified : Verified X86_64.target makeHintAvx2 (makeHintContract X86_64.abi) :=
+ Verified.of_correct (fun _ hp => makeHintY_correct hp) makeHintY_ct (by
+ round_implies [makeHintContract, makeHintSig, makeHintK, hintK, X86_64.abi, X86_64.argRegs] [hintSat]
+ using hintSat)
+
+end VG.Proof.MlDsa.X86_64.Round
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YLane.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YLane.lean
new file mode 100644
index 000000000..6c354f4b1
--- /dev/null
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YLane.lean
@@ -0,0 +1,130 @@
+import VerifiedGarbage.Impl.MlDsa.X86_64.Round.Avx2
+import VerifiedGarbage.Proof.MlDsa.Round.Decompose
+import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.VArith
+
+/-!
+# ML-DSA on x86-64: what the AVX2 rounding code computes in a doubleword
+
+Untrusted: everything here is checked by Lean. What `hbX` and `lbX`
+(`Impl/MlDsa/X86_64/Round/Avx2.lean`) compute in each doubleword (`hbL`,
+`lbL`), and that it is `r₁` and `r₀` of `Decompose` (`hbL_toNat`,
+`lbL_toNat`): every intermediate value fits in 32 bits.
+-/
+
+namespace VG.Proof.MlDsa.X86_64.Round
+
+open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round
+open VG.Spec.MlDsa (q gamma2s ofInt lowBits)
+open VG.Proof.MlDsa.Round (hbF hbM q_eq mem_gamma2s hbF_le hbF_eq hbM_mul lowBits_val)
+open VG.Proof.MlDsa.X86_64.Arith (caddL caddL_toNat sshiftRight31)
+
+/-- `t · (1 + Σ 2^k)` by shifts and additions, as `mulX` computes it. -/
+def mulL (sh : List Nat) (t : BitVec 32) : BitVec 32 := sh.foldl (fun acc k => acc + (t <<< k)) t
+
+theorem foldl_toNat (t : BitVec 32) : ∀ (sh : List Nat) (acc : BitVec 32),
+ acc.toNat + t.toNat * (sh.map (2 ^ ·)).sum < 2 ^ 32 →
+ (sh.foldl (fun acc k => acc + (t <<< k)) acc).toNat = acc.toNat + t.toNat * (sh.map (2 ^ ·)).sum
+ | [], acc, _ => by simp
+ | k :: sh, acc, h => by
+ simp only [List.map_cons, List.sum_cons, Nat.mul_add] at h ⊢
+ have hk : (t <<< k).toNat = t.toNat * 2 ^ k := by
+ rw [BitVec.toNat_shiftLeft, Nat.shiftLeft_eq, Nat.mod_eq_of_lt (by omega)]
+ have ha : (acc + t <<< k).toNat = acc.toNat + t.toNat * 2 ^ k := by
+ rw [BitVec.toNat_add, hk, Nat.mod_eq_of_lt (by omega)]
+ rw [List.foldl_cons, foldl_toNat t sh _ (by rw [ha]; omega), ha]
+ omega
+
+theorem mulL_toNat (sh : List Nat) {t : BitVec 32} (h : t.toNat * (1 + (sh.map (2 ^ ·)).sum) < 2 ^ 32) :
+ (mulL sh t).toNat = t.toNat * (1 + (sh.map (2 ^ ·)).sum) := by
+ rw [mulL, foldl_toNat t sh t (by rw [Nat.mul_add] at h; omega), Nat.mul_add]; omega
+
+theorem dSh_sum {g : Nat} (h : g ∈ gamma2s) : 1 + ((dSh g).map (2 ^ ·)).sum = dMul g := by
+ rcases mem_gamma2s h with rfl | rfl <;> rfl
+
+/-- `f`, as `hbX` computes it in a doubleword. -/
+def hbFL (g : Nat) (x : BitVec 32) : BitVec 32 :=
+ (mulL (dSh g) ((x + BitVec.ofNat 32 127) >>> 7) + BitVec.ofNat 32 (dAdd g)) >>> dShift g
+
+/-- What `hbX` computes in a doubleword. -/
+def hbL (g : Nat) (x : BitVec 32) : BitVec 32 :=
+ hbFL g x &&& (hbFL g x - BitVec.ofNat 32 (dMod g)).sshiftRight (min (31 : BitVec 8).toNat 32)
+
+theorem dMod_eq' {g : Nat} (h : g ∈ gamma2s) : dMod g = hbM g := by
+ rcases mem_gamma2s h with rfl | rfl <;> rfl
+
+/-- `f` in 32 bits. -/
+theorem hbF32 {g : Nat} (h : g ∈ gamma2s) {x : BitVec 32} (hx : x.toNat < q) :
+ (hbFL g x).toNat = hbF g x.toNat := by
+ rw [q_eq] at hx
+ have e1 : (x + BitVec.ofNat 32 127).toNat = x.toNat + 127 := by
+ rw [BitVec.toNat_add, BitVec.toNat_ofNat]; omega
+ have e2 : ((x + BitVec.ofNat 32 127) >>> 7).toNat = (x.toNat + 127) / 128 := by
+ rw [BitVec.toNat_ushiftRight, e1, Nat.shiftRight_eq_div_pow]
+ have hM : dMul g ≤ 11275 := by unfold dMul; split <;> decide
+ have hA : dAdd g ≤ 2 ^ 23 := by unfold dAdd; split <;> decide
+ have hp : (x.toNat + 127) / 128 * dMul g ≤ 65473 * 11275 := Nat.mul_le_mul (by omega) hM
+ have e3 : (mulL (dSh g) ((x + BitVec.ofNat 32 127) >>> 7)).toNat = (x.toNat + 127) / 128 * dMul g := by
+ rw [mulL_toNat _ (by rw [e2, dSh_sum h]; omega), e2, dSh_sum h]
+ have e4 : (mulL (dSh g) ((x + BitVec.ofNat 32 127) >>> 7) + BitVec.ofNat 32 (dAdd g)).toNat =
+ (x.toNat + 127) / 128 * dMul g + dAdd g := by
+ rw [BitVec.toNat_add, e3, BitVec.toNat_ofNat, Nat.mod_eq_of_lt (show dAdd g < 2 ^ 32 by omega)]; omega
+ rw [hbFL, BitVec.toNat_ushiftRight, e4, Nat.shiftRight_eq_div_pow, hbF_eq h (by rw [q_eq]; exact hx)]
+ rfl
+
+theorem hbL_toNat {g : Nat} (h : g ∈ gamma2s) {x : BitVec 32} (hx : x.toNat < q) :
+ (hbL g x).toNat = hbF g x.toNat % hbM g := by
+ have hf := hbF32 h hx
+ have hle := hbF_le h hx
+ have hm : dMod g ≤ 44 ∧ 16 ≤ dMod g := by unfold dMod; split <;> decide
+ rw [← dMod_eq' h] at hle ⊢
+ unfold hbL
+ generalize hbFL g x = f at hf ⊢
+ rw [sshiftRight31]
+ have hsub : (f - BitVec.ofNat 32 (dMod g)).toNat = (f.toNat + 2 ^ 32 - dMod g) % 2 ^ 32 := by
+ rw [BitVec.toNat_sub, BitVec.toNat_ofNat, Nat.mod_eq_of_lt (show dMod g < 2 ^ 32 by omega)]; omega
+ by_cases e : hbF g x.toNat = dMod g
+ · rw [ite_eq_left (by rw [hsub, hf, e]; omega), show f &&& (0 : BitVec 32) = 0 from BitVec.and_zero, e, Nat.mod_self]; rfl
+ · rw [ite_eq_right (by rw [hsub, hf]; omega), show (-1 : BitVec 32) = BitVec.allOnes 32 by decide, BitVec.and_allOnes,
+ hf, Nat.mod_eq_of_lt (by omega)]
+
+/-- `r₁ · 2γ₂` by shifts and additions, as `mul2X` computes it. -/
+def mul2L (g : Nat) (r : BitVec 32) : BitVec 32 :=
+ if g = 261888 then r <<< 19 - r <<< 9
+ else [13, 14, 15, 17].foldl (fun acc k => acc + (r <<< k)) (r <<< 11)
+
+theorem mul2L_toNat {g : Nat} (h : g ∈ gamma2s) {r : BitVec 32} (hr : r.toNat ≤ 44) :
+ (mul2L g r).toNat = r.toNat * (2 * g) := by
+ have s : ∀ k ≤ 19, (r <<< k).toNat = r.toNat * 2 ^ k := fun k hk => by
+ rw [BitVec.toNat_shiftLeft, Nat.shiftLeft_eq, Nat.mod_eq_of_lt]
+ calc r.toNat * 2 ^ k ≤ 44 * 2 ^ 19 := Nat.mul_le_mul hr (Nat.pow_le_pow_right (by decide) hk)
+ _ < 2 ^ 32 := by decide
+ rcases mem_gamma2s h with rfl | rfl
+ · unfold mul2L
+ rw [ite_eq_right (by decide)]
+ simp only [List.foldl_cons, List.foldl_nil]
+ rw [BitVec.toNat_add, BitVec.toNat_add, BitVec.toNat_add, BitVec.toNat_add, s 11 (by decide), s 13 (by decide),
+ s 14 (by decide), s 15 (by decide), s 17 (by decide)]
+ omega
+ · unfold mul2L
+ rw [ite_eq_left rfl, BitVec.toNat_sub, s 19 (by decide), s 9 (by decide)]
+ omega
+
+/-- What `lbX` computes in a doubleword. -/
+def lbL (g : Nat) (x : BitVec 32) : BitVec 32 := caddL (x - mul2L g (hbL g x))
+
+theorem lbL_toNat {g : Nat} (h : g ∈ gamma2s) {x : BitVec 32} (hx : x.toNat < q) :
+ (lbL g x).toNat = (ofInt (lowBits g (Fin.ofNat q x.toNat))).val := by
+ have hv : (Fin.ofNat q x.toNat).val = x.toNat := Nat.mod_eq_of_lt hx
+ rw [lowBits_val h, hv]
+ have hlt : hbF g x.toNat % hbM g < hbM g := Nat.mod_lt _ (by rcases mem_gamma2s h with rfl | rfl <;> decide)
+ have hM : hbM g ≤ 44 := by rcases mem_gamma2s h with rfl | rfl <;> decide
+ have hle : hbF g x.toNat % hbM g * (2 * g) ≤ q - 1 := by
+ rw [← hbM_mul h]; exact Nat.mul_le_mul_right _ (Nat.le_of_lt hlt)
+ have e1 : (mul2L g (hbL g x)).toNat = hbF g x.toNat % hbM g * (2 * g) := by
+ rw [mul2L_toNat h (by rw [hbL_toNat h hx]; omega), hbL_toNat h hx]
+ unfold lbL
+ rw [caddL_toNat, BitVec.toNat_sub, e1]
+ rw [q_eq] at hx hle ⊢
+ split <;> split <;> omega
+
+end VG.Proof.MlDsa.X86_64.Round
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YNorm.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YNorm.lean
new file mode 100644
index 000000000..a94ca1e18
--- /dev/null
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Round/YNorm.lean
@@ -0,0 +1,403 @@
+import VerifiedGarbage.Proof.MlDsa.X86_64.Round.YBits
+import VerifiedGarbage.Proof.MlDsa.X86_64.Round.NormLt
+import VerifiedGarbage.Proof.MlKem.X86_64.S4Vec
+
+/-!
+# ML-DSA on x86-64: `vg_mldsa_norm_lt_avx2`
+
+Untrusted: everything here is checked by Lean. With the bound clamped to
+`b ≤ q` (which changes no result, `good_clamp`), each doubleword of `ymm10`
+keeps its top bit while every coefficient it has seen is good
+(`Good b a`: `a < b` or `q - a < b`, `nlL_msb`); at the end the top bits
+of the 32 bytes are all set exactly when every coefficient is good
+(`allOnes_bsum`).
+-/
+
+namespace VG.Proof.MlDsa.X86_64.Round
+
+open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round VG.Proof.MlDsa.Round
+open VG.Spec.MlDsa
+open VG.Proof.MlKem.X86_64 (Keep XOnly YOnly ylanes ifp ifn)
+open VG.Impl.MlKem.X86_64 (xb xmov toY)
+open VG.Proof.MlDsa.X86_64.Arith (dword_psubd dword_pand dword_psrad)
+open VG.Proof.MlKem.X86_64.S4 (bsum bsum_lt byteMask_eq)
+open VG.Proof.MlDsa.X86_64.Arith (bc)
+
+/-! ## A doubleword -/
+
+/-- The top bit of a difference of doublewords below `2³¹`: whether it borrows. -/
+theorem msb_sub32 {u v : BitVec 32} (hu : u.toNat < 2 ^ 31) (hv : v.toNat < 2 ^ 31) :
+ (u - v).msb = decide (u.toNat < v.toNat) := by
+ rw [BitVec.msb_eq_decide, BitVec.toNat_sub]
+ by_cases h : u.toNat < v.toNat
+ · rw [decide_eq_true h, decide_eq_true_iff]; omega
+ · rw [decide_eq_false h, decide_eq_false_iff_not]; omega
+
+/-- What `nlX` ORs into a doubleword: `(a - b) | ((q - b) - a)`. -/
+def nlL (a b c : BitVec 32) : BitVec 32 := (a - b) ||| (c - a)
+
+theorem nlL_msb {a b : BitVec 32} (ha : a.toNat < q) (hb : b.toNat ≤ q) :
+ (nlL a b (BitVec.ofNat 32 (q - b.toNat))).msb = decide (Good b.toNat a.toNat) := by
+ rw [q_eq] at ha hb
+ have hc : (BitVec.ofNat 32 (q - b.toNat)).toNat = q - b.toNat := by rw [BitVec.toNat_ofNat, q_eq]; omega
+ rw [nlL, BitVec.msb_or, msb_sub32 (by omega) (by omega), msb_sub32 (by rw [hc, q_eq]; omega) (by omega), hc]
+ simp only [Good, Bool.decide_or, q_eq]
+ congr 1
+ apply Bool.eq_iff_iff.mpr
+ simp only [decide_eq_true_iff]
+ omega
+
+/-- Clamping the bound to `q` changes no coefficient's goodness. -/
+theorem good_clamp {B a : Nat} (ha : a < q) : Good (min B q) a ↔ Good B a := by
+ simp only [Good]
+ constructor
+ · rintro (h | h) <;> [left; right] <;> omega
+ · intro h
+ by_cases hB : B ≤ q
+ · rw [Nat.min_eq_left hB]; exact h
+ · rw [Nat.min_eq_right (by omega)]; left; exact ha
+
+/-! ## The mask -/
+
+theorem bsum_allOnes (f : Nat → Bool) : ∀ n, bsum f n = 2 ^ n - 1 ↔ ∀ i < n, f i = true
+ | 0 => by simp [bsum]
+ | n + 1 => by
+ have hl := bsum_lt f n
+ have ih := bsum_allOnes f n
+ have h1 : 1 ≤ 2 ^ n := Nat.one_le_two_pow
+ simp only [bsum, Nat.pow_succ]
+ generalize 2 ^ n = N at hl ih h1 ⊢
+ constructor
+ · intro h i hi
+ have hfn : f n = true := by
+ cases e : f n
+ · simp only [e, Bool.toNat_false, Nat.zero_mul, Nat.add_zero] at h; omega
+ · rfl
+ rw [hfn] at h
+ simp only [Bool.toNat_true, Nat.one_mul] at h
+ rcases (by omega : i < n ∨ i = n) with hi | rfl
+ · exact (ih.mp (by omega)) i hi
+ · exact hfn
+ · intro h
+ rw [ih.mpr fun i hi => h i (by omega), h n (by omega)]
+ simp only [Bool.toNat_true, Nat.one_mul]
+ omega
+
+/-- A bit of a 256-bit register, in its lanes and doublewords. -/
+theorem ymm_bit (s : State) (r : XReg) {i : Nat} (hi : i < 32) :
+ (s.ymm r).getLsbD (8 * i + 7) = (dword (s.lane r (i / 16)) (i % 16 / 4)).getLsbD (8 * (i % 4) + 7) := by
+ simp only [State.ymm, State.lane, dword, BitVec.getLsbD_append, BitVec.getLsbD_extractLsb',
+ show 8 * (i % 4) + 7 < 32 by omega, decide_true, Bool.true_and]
+ by_cases h : i < 16
+ · rw [ite_eq_left (show 8 * i + 7 < 128 by omega), ite_eq_left (show i / 16 = 0 by omega),
+ show 32 * (i % 16 / 4) + (8 * (i % 4) + 7) = 8 * i + 7 by omega]
+ · rw [ite_eq_right (show ¬ 8 * i + 7 < 128 by omega), ite_eq_right (show ¬ i / 16 = 0 by omega),
+ show 32 * (i % 16 / 4) + (8 * (i % 4) + 7) = 8 * i + 7 - 128 by omega]
+
+/-- A doubleword spread from its top bit. -/
+theorem bit_of_spread {d : BitVec 32} (h : d = 0 ∨ d = BitVec.allOnes 32) {j : Nat} (hj : j < 32) :
+ d.getLsbD j = d.msb := by
+ rcases h with rfl | rfl
+ · simp
+ · rw [BitVec.msb_allOnes (by decide), BitVec.getLsbD_allOnes]; simp [hj]
+
+/-! ## The code on a register -/
+
+theorem nlX_ok (s : State) :
+ WP isa (.block nlX) s fun s' => (∀ e < 4, dword (s'.xmm .xmm10) e =
+ dword (s.xmm .xmm10) e &&& nlL (dword (s.xmm .xmm0) e) (dword (s.xmm .xmm8) e) (dword (s.xmm .xmm9) e)) ∧
+ XOnly [.xmm1, .xmm2, .xmm10] s s' := by
+ simp only [nlX, xmov, xb]
+ vrun [VG.X86_64.eval_movdqa]
+ refine ⟨fun e he => ?_, by xonly⟩
+ simp only [dword_pand, dword_por, dword_psubd _ _ he]
+ rfl
+
+theorem lane_nlX : laneSseBlock (toY nlX) = some nlX := by decide +kernel
+
+/-- `rax`'s low doubleword in each doubleword of `ymm r`. -/
+theorem ybc_ok (r : XReg) (s : State) :
+ WP isa (.block (ybcast r)) s fun s' =>
+ (∀ l < 2, s'.lane r l = bc ((s.gpr .rax).setWidth 32)) ∧ s'.gpr = s.gpr ∧ s'.mem = s.mem ∧ s'.rd = s.rd ∧
+ s'.wr = s.wr ∧ s'.mxcsr = s.mxcsr ∧ ∀ r' ≠ r, ∀ l < 2, s'.lane r' l = s.lane r' l := by
+ apply WP.of_runBlock
+ simp only [ybcast, runBlock_cons, runStep_some, runBlock_nil, exec, VOp.exec, Option.some.injEq,
+ exists_eq_left']
+ refine ⟨fun l hl => ?_, rfl, rfl, rfl, rfl, rfl, fun r' hr' l hl => ?_⟩
+ · rw [State.lane_setV256, ifp rfl]
+ have : dword ((s.setV .l128 r ((0 : BitVec 64) ++ s.gpr .rax) 0).xmm r) 0 = (s.gpr .rax).setWidth 32 := by
+ rw [RegUpd.xmm_setV, ifp rfl]
+ apply BitVec.eq_of_getLsbD_eq; intro i hi
+ simp only [dword, BitVec.getLsbD_extractLsb', hi, decide_true, Bool.true_and]
+ rw [BitVec.getLsbD_append, ifp (by omega), BitVec.getLsbD_setWidth]
+ simp [show i < 64 by omega, hi]
+ rcases lane01 hl with rfl | rfl <;> simp only [ite_true, ite_false, this, Nat.one_ne_zero] <;> rfl
+ · rw [State.lane_setV256, ifn hr', State.lane_setV128, ifn hr']
+
+end VG.Proof.MlDsa.X86_64.Round
+
+namespace VG.Proof.MlDsa.X86_64.Arith
+
+open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round
+open VG.Proof.MlDsa.Arith
+open VG.Proof.MlDsa.X86_64.Round (Good nlL nlL_msb nlX_ok lane_nlX bc_ofDwords)
+open VG.Proof.MlKem.X86_64 (Keep XOnly YOnly ylanes yld_ok ifp ifn ptr_step add_ofNat_zero lane_setReg lane_setFlags)
+open VG.Spec.MlDsa (q n coeffAt Reduced)
+
+namespace YNormL
+
+/-- After `i` vectors of eight: the top bit of doubleword `e` of lane `l` of
+`ymm10` is whether coefficients `8j + 4l + e`, `j < i`, are good. -/
+structure Inv (s₀ : State) (b : BitVec 32) (i : Nat) (s : State) : Prop where
+ rdi : s.gpr .rdi = s₀.gpr .rdi + BitVec.ofNat 64 (32 * i)
+ rd : s.rd = s₀.rd
+ wr : s.wr = s₀.wr
+ mem : s.mem = s₀.mem
+ c8 : ∀ l < 2, s.lane .xmm8 l = bc b
+ c9 : ∀ l < 2, s.lane .xmm9 l = bc (BitVec.ofNat 32 (q - b.toNat))
+ acc : ∀ l < 2, ∀ e < 4, (dword (s.lane .xmm10 l) e).msb = true ↔
+ ∀ j < i, Good b.toNat (coeffAt s₀.mem (s₀.gpr .rdi) (8 * j + 4 * l + e)).toNat
+
+theorem step {s₀ : State} (hrd : pR (s₀.gpr .rdi) ∈ s₀.rd) (hr : Reduced s₀.mem (s₀.gpr .rdi)) {b : BitVec 32}
+ (hb : b.toNat ≤ q) {i : Nat} (hi : i < 32) {s : State} (hI : Inv s₀ b i s) :
+ WP isa (.block (nlBodyY ++ ([.alu .sub .rcx (.imm 1)] : List Instr))) s fun s' =>
+ Inv s₀ b (i + 1) s' ∧ s'.gpr .rcx = s.gpr .rcx - 1 ∧ s'.zf = some (s.gpr .rcx - 1 == 0) := by
+ have j0 : 8 * i + 8 ≤ 256 := by omega
+ have hr' : pR (s₀.gpr .rdi) ∈ s.rd ++ s.wr := by rw [hI.rd]; exact List.mem_append_left _ hrd
+ have e1 : s.gpr .rdi + BitVec.ofNat 64 0 = coeffAddr (s₀.gpr .rdi) (8 * i) := by
+ rw [add_ofNat_zero, hI.rdi]; congr 2; omega
+ rw [nlBodyY, List.append_assoc, List.append_assoc, WP.block_append_iff]
+ refine WP.mono (yld_ok (by rw [e1]; exact f_in32 hr' j0)) fun s1 ⟨L1, o1⟩ => ?_
+ rw [WP.block_append_iff]
+ refine WP.mono (ylanes lane_nlX (P := fun l t => ∀ e < 4, dword (t.xmm .xmm10) e =
+ dword ((s1.proj l).xmm .xmm10) e &&& nlL (dword ((s1.proj l).xmm .xmm0) e) (dword ((s1.proj l).xmm .xmm8) e)
+ (dword ((s1.proj l).xmm .xmm9) e)) fun l hl => nlX_ok _) fun s3 ⟨B3, o3⟩ => ?_
+ have o13 := o1.trans o3
+ vrund
+ refine ⟨⟨?_, ?_, ?_, ?_, fun l hl => ?_, fun l hl => ?_, fun l hl e he => ?_⟩, ?_⟩
+ · simp only [RegUpd.gpr_setReg, RegUpd.gpr_setFlags, ite_true, ite_false, reduceCtorEq]
+ rw [o13.gpr, hI.rdi]; exact ptr_step _ i 32
+ · simp only [RegUpd.rd_setReg, RegUpd.rd_setFlags]; rw [o13.rd, hI.rd]
+ · simp only [RegUpd.wr_setReg, RegUpd.wr_setFlags]; rw [o13.wr, hI.wr]
+ · simp only [RegUpd.mem_setReg, RegUpd.mem_setFlags]; rw [o13.mem, hI.mem]
+ · simp only [lane_setReg, lane_setFlags]; rw [o13.lane _ (by decide) l hl, hI.c8 l hl]
+ · simp only [lane_setReg, lane_setFlags]; rw [o13.lane _ (by decide) l hl, hI.c9 l hl]
+ · simp only [lane_setReg, lane_setFlags]
+ have hx : dword ((s1.proj l).xmm .xmm0) e = coeffAt s₀.mem (s₀.gpr .rdi) (8 * i + 4 * l + e) := by
+ rw [State.proj_xmm, L1 l hl, e1, dword_readW _ _ he, lane_load, coeffAddr_add, ← coeffAt_eq, hI.mem]
+ have h8 : dword ((s1.proj l).xmm .xmm8) e = b := by
+ rw [State.proj_xmm, o1.lane _ (by decide) l hl, hI.c8 l hl]; exact bc_ofDwords _ e he
+ have h9 : dword ((s1.proj l).xmm .xmm9) e = BitVec.ofNat 32 (q - b.toNat) := by
+ rw [State.proj_xmm, o1.lane _ (by decide) l hl, hI.c9 l hl]; exact bc_ofDwords _ e he
+ have h10 : (s1.proj l).xmm .xmm10 = s.lane .xmm10 l := by rw [State.proj_xmm, o1.lane _ (by decide) l hl]
+ rw [← State.proj_xmm, B3 l hl e he, BitVec.msb_and, Bool.and_eq_true, hx, h8, h9, h10,
+ nlL_msb (by rw [← VG.Proof.MlDsa.Round.n_eq] at *; exact hr _ (by rw [VG.Proof.MlDsa.Round.n_eq]; omega)) hb,
+ decide_eq_true_iff, hI.acc l hl e he]
+ refine ⟨fun ⟨h₁, h₂⟩ j hj => ?_, fun h => ⟨fun j hj => h j (by omega), h i (by omega)⟩⟩
+ rcases (by omega : j < i ∨ j = i) with hj | rfl
+ exacts [h₁ j hj, h₂]
+ · exact ⟨by rw [o13.gpr], by rw [o13.gpr]⟩
+
+end YNormL
+
+end VG.Proof.MlDsa.X86_64.Arith
+
+namespace VG.Proof.MlDsa.X86_64.Round
+
+open VG VG.X86_64 VG.Impl.MlDsa.X86_64.Round VG.Proof.MlDsa.Round
+open VG.Spec.MlDsa
+open VG.Proof.MlKem.X86_64 (Keep Keep.gpr WP.keep gprPreserved_of XOnly YOnly ylanes yconst_ok ifp ifn)
+open VG.Proof.MlDsa.X86_64.Arith (bc dword_psrad sshiftRight31)
+open VG.Impl.MlKem.X86_64 (toY)
+open VG.Proof.MlKem.X86_64.S4 (bsum byteMask_eq)
+
+theorem sw3264 (y : BitVec 32) : BitVec.setWidth 32 (BitVec.setWidth 64 y) = y := by
+ apply BitVec.eq_of_toNat_eq; rw [BitVec.toNat_setWidth, BitVec.toNat_setWidth]; have := y.isLt; omega
+
+/-- The bound, clamped to `q`. -/
+theorem nlPro_ok (s₀ : State) :
+ WP isa nlPro s₀ fun s₂ => (s₂.mem = s₀.mem ∧ ((s₂.gpr .rsi).setWidth 32).toNat = min (arg32 s₀ .rsi) q) ∧
+ Keep [.rsi] s₀ s₂ := by
+ unfold nlPro
+ refine WP.seq (WP.mono (Q := fun (s₁ : State) => (s₁.cf = some (decide (((s₀.gpr .rsi).setWidth 32).toNat < q)) ∧
+ s₁.mem = s₀.mem ∧ s₁.gpr .rsi = BitVec.setWidth 64 ((s₀.gpr .rsi).setWidth 32)) ∧ Keep [.rsi] s₀ s₁)
+ (by refine WP.keep _ ?_ (by decide); xrun; rfl)
+ fun s₁ ⟨⟨hc, hm, hsi⟩, hk⟩ => ?_)
+ refine WP.ite (M := isa) _ (show isa.eval .b s₁ = _ from hc) (fun h => ?_) (fun h => ?_)
+ · rw [decide_eq_true_iff] at h
+ refine WP.mono (Q := fun s₂ => s₂ = s₁) (by vrund) fun s₂ e => ?_
+ subst e
+ refine ⟨⟨hm, ?_⟩, hk⟩
+ rw [hsi, sw3264, Nat.min_eq_left (by unfold arg32; omega)]
+ · rw [decide_eq_false_iff_not] at h
+ refine WP.mono (Q := fun (s₂ : State) => (s₂.gpr .rsi = BitVec.setWidth 64 qImm ∧ s₂.mem = s₁.mem) ∧
+ Keep [.rsi] s₁ s₂) (by refine WP.keep _ ?_ (by decide); xrun) fun s₂ ⟨⟨h1, h2⟩, k2⟩ => ?_
+ refine ⟨⟨h2.trans hm, ?_⟩, (hk.trans k2).mono (by simp)⟩
+ rw [h1, Nat.min_eq_right (by unfold arg32; omega)]; rfl
+
+theorem qsub_eq {x : BitVec 32} (hx : x.toNat ≤ q) : qImm - x = BitVec.ofNat 32 (q - x.toNat) := by
+ apply BitVec.eq_of_toNat_eq
+ rw [BitVec.toNat_sub, BitVec.toNat_ofNat, show qImm.toNat = q from rfl]
+ rw [q_eq] at hx ⊢; omega
+
+/-- The constants of the loop: `b`, `q - b` and all ones. -/
+theorem nlConsts_ok (s : State) (hb : ((s.gpr .rsi).setWidth 32).toNat ≤ q) :
+ WP isa (.block nlConsts) s fun s' =>
+ (∀ l < 2, s'.lane .xmm8 l = bc ((s.gpr .rsi).setWidth 32) ∧
+ s'.lane .xmm9 l = bc (BitVec.ofNat 32 (q - ((s.gpr .rsi).setWidth 32).toNat)) ∧
+ s'.lane .xmm10 l = bc (BitVec.allOnes 32)) ∧ Keep [.rax] s s' ∧ s'.mem = s.mem ∧ s'.mxcsr = s.mxcsr := by
+ simp only [nlConsts, List.append_assoc]
+ rw [WP.block_append_iff]
+ refine WP.mono (Q := fun (s1 : State) => (s1.gpr .rax = BitVec.setWidth 64 ((s.gpr .rsi).setWidth 32) ∧
+ s1.mem = s.mem ∧ s1.mxcsr = s.mxcsr ∧ ∀ r l, s1.lane r l = s.lane r l) ∧ Keep [.rax] s s1)
+ (by refine WP.keep _ ?_ (by decide); xrun; exact ⟨rfl, fun _ _ => rfl⟩) fun s1 ⟨⟨a1, m1, x1, l1⟩, k1⟩ => ?_
+ rw [WP.block_append_iff]
+ refine WP.mono (ybc_ok .xmm8 s1) fun s2 ⟨c2, g2, m2, r2, w2, x2, o2⟩ => ?_
+ rw [WP.block_append_iff]
+ refine WP.mono (Q := fun (s3 : State) => (s3.gpr .rax = BitVec.setWidth 64 (qImm - (s2.gpr .rsi).setWidth 32) ∧
+ s3.mem = s2.mem ∧ s3.mxcsr = s2.mxcsr ∧ ∀ r l, s3.lane r l = s2.lane r l) ∧ Keep [.rax] s2 s3)
+ (by refine WP.keep _ ?_ (by decide); xrun; exact ⟨rfl, fun _ _ => rfl⟩) fun s3 ⟨⟨a3, m3, x3, l3⟩, k3⟩ => ?_
+ rw [WP.block_append_iff]
+ refine WP.mono (ybc_ok .xmm9 s3) fun s4 ⟨c4, g4, m4, r4, w4, x4, o4⟩ => ?_
+ refine WP.mono (yconst_ok .xmm10 _ s4) fun s5 ⟨c5, k5, m5, x5, o5⟩ => ?_
+ have e2 : s2.gpr .rsi = s.gpr .rsi := by rw [g2, k1.gpr (by decide)]
+ refine ⟨fun l hl => ⟨?_, ?_, ?_⟩, ?_, by rw [m5, m4, m3, m2, m1], by rw [x5, x4, x3, x2, x1]⟩
+ · rw [o5 _ (by decide) l hl, o4 _ (by decide) l hl, l3, c2 l hl, a1, sw3264]
+ · rw [o5 _ (by decide) l hl, c4 l hl, a3, e2, sw3264, qsub_eq hb]
+ · rw [c5 l hl]; rfl
+ · refine ⟨fun r hr => ?_, ?_, ?_⟩
+ · rw [k5.gpr hr, g4, k3.gpr hr, g2, k1.gpr hr]
+ · rw [k5.2.1, r4, k3.2.1, r2, k1.2.1]
+ · rw [k5.2.2, w4, k3.2.2, w2, k1.2.2]
+
+theorem lane_psrad : laneSseBlock (toY [.xop (.shift .psrad .xmm10 31)]) = some [.xop (.shift .psrad .xmm10 31)] := by
+ decide +kernel
+
+/-- Whether every doubleword of `ymm10` has its top bit set. -/
+def allMsb (s : State) : Bool := (List.range 2).all fun l => (List.range 4).all fun e => (dword (s.lane .xmm10 l) e).msb
+
+/-- The end: 1 if every doubleword of `ymm10` has its top bit set, and 0 otherwise. -/
+theorem nlEnd_ok (s : State) :
+ WP isa (.block nlEnd) s fun s' =>
+ ((s'.gpr .rax).setWidth 32 = if allMsb s = true then 1 else 0) ∧
+ s'.mem = s.mem ∧ s'.rd = s.rd ∧ s'.wr = s.wr ∧ ∀ r, r ≠ .rax → s'.gpr r = s.gpr r := by
+ rw [nlEnd, List.append_assoc, WP.block_append_iff]
+ refine WP.mono (ylanes lane_psrad (rs := [.xmm10]) (P := fun l t => ∀ e < 4,
+ dword (t.xmm .xmm10) e = (dword ((s.proj l).xmm .xmm10) e).sshiftRight (min (31 : BitVec 8).toNat 32))
+ fun l hl => by vrun; exact ⟨fun e he => dword_psrad _ _ he, by xonly⟩) fun s1 ⟨B1, o1⟩ => ?_
+ rw [WP.block_append_iff]
+ refine WP.mono (Q := fun (s2 : State) => s2.gpr = (s1.setReg .rax (byteMask (s1.ymm .xmm10) 32)).gpr ∧
+ s2.mem = s1.mem ∧ s2.rd = s1.rd ∧ s2.wr = s1.wr) (by vrund; exact ⟨rfl, rfl, rfl, rfl⟩) fun s2 ⟨g2, m2, r2, w2⟩ => ?_
+ xrun
+ have hf : ∀ i < 32, (s1.ymm .xmm10).getLsbD (8 * i + 7) = (dword (s.lane .xmm10 (i / 16)) (i % 16 / 4)).msb :=
+ fun i hi => by
+ have hl : i / 16 < 2 := by omega
+ have he : i % 16 / 4 < 4 := by omega
+ have hd := B1 (i / 16) hl (i % 16 / 4) he
+ rw [State.proj_xmm, State.proj_xmm, sshiftRight31] at hd
+ rw [ymm_bit _ _ hi, hd]
+ split
+ · rename_i h
+ rw [BitVec.msb_eq_decide]; simp [h]
+ · rename_i h
+ rw [show (-1 : BitVec 32) = BitVec.allOnes 32 by decide, BitVec.getLsbD_allOnes, BitVec.msb_eq_decide]
+ simp only [decide_eq_true (show 8 * (i % 4) + 7 < 32 by omega)]
+ rw [decide_eq_true (by omega)]
+ have hall : allMsb s = true ↔ ∀ i < 32, (s1.ymm .xmm10).getLsbD (8 * i + 7) = true := by
+ simp only [allMsb, List.all_eq_true, List.mem_range]
+ constructor
+ · intro h i hi; rw [hf i hi]; exact h _ (by omega) _ (by omega)
+ · intro h l hl e he
+ have := h (16 * l + 4 * e) (by omega)
+ rwa [hf _ (by omega), show (16 * l + 4 * e) / 16 = l by omega, show (16 * l + 4 * e) % 16 / 4 = e by omega]
+ at this
+ have hb := VG.Proof.MlKem.X86_64.S4.bsum_lt (fun i => (s1.ymm .xmm10).getLsbD (8 * i + 7)) 32
+ have hax : s2.gpr .rax = BitVec.ofNat 64 (bsum (fun i => (s1.ymm .xmm10).getLsbD (8 * i + 7)) 32) := by
+ rw [g2, RegUpd.gpr_setReg_self, byteMask_eq _ (by decide)]
+ refine ⟨?_, by rw [m2, o1.mem], by rw [r2, o1.rd], by rw [w2, o1.wr], fun r hr => ?_⟩
+ · rw [hax]
+ apply BitVec.eq_of_toNat_eq
+ rw [BitVec.toNat_setWidth, BitVec.toNat_ushiftRight, BitVec.toNat_add, BitVec.toNat_ofNat,
+ Nat.shiftRight_eq_div_pow]
+ have e1 : (1 : BitVec 64).toNat = 1 := rfl
+ rw [e1]
+ by_cases h : allMsb s = true
+ · rw [ite_eq_left h, (bsum_allOnes _ 32).mpr (hall.mp h)]; rfl
+ · rw [ite_eq_right h]
+ have : bsum (fun i => (s1.ymm .xmm10).getLsbD (8 * i + 7)) 32 ≠ 2 ^ 32 - 1 := fun e =>
+ h (hall.mpr ((bsum_allOnes _ 32).mp e))
+ show _ = 0
+ omega
+ · rw [ite_eq_right hr, ite_eq_right hr, g2, RegUpd.gpr_setReg_of_ne _ _ hr, o1.gpr]
+
+section
+variable {s₀ : State} (hp : normLtK.pre s₀)
+include hp
+
+theorem normLtY_wp : WP isa normLtAvx2 s₀ fun s' =>
+ ((s'.gpr .rax).setWidth 32 = if normRq [polyAt s₀.mem (s₀.gpr .rdi)] < arg32 s₀ .rsi then 1 else 0) ∧
+ Frame [] s₀.mem s'.mem := by
+ have hr : Reduced s₀.mem (s₀.gpr .rdi) := hp.2.2.2
+ have hrd : pR (s₀.gpr .rdi) ∈ s₀.rd := by rw [hp.1]; simp
+ unfold normLtAvx2
+ refine WP.seq (WP.mono (nlPro_ok s₀) fun s₂ ⟨⟨hm2, hb2⟩, k2⟩ => ?_)
+ have hbq : ((s₂.gpr .rsi).setWidth 32).toNat ≤ q := by rw [hb2]; exact Nat.min_le_right _ _
+ refine WP.seq (WP.mono (nlConsts_ok s₂ hbq) fun s₃ ⟨hc, k3, m3, _⟩ => ?_)
+ refine WP.seq (WP.mono (VG.Proof.MlKem.X86_64.wp_rcxLoopY (N := 32) (by decide) (by decide)
+ (Arith.YNormL.Inv s₀ ((s₂.gpr .rsi).setWidth 32)) (fun u o hy _ => ⟨?_, ?_, ?_, ?_, fun l hl => ?_, fun l hl => ?_,
+ fun l hl e he => ?_⟩) fun i hi u hI => Arith.YNormL.step hrd hr hbq hi hI) fun s₄ hI => ?_)
+ · rw [o.keep.gpr (by decide), k3.gpr (by decide), k2.gpr (by decide), Nat.mul_zero,
+ VG.Proof.MlKem.X86_64.add_ofNat_zero]
+ · rw [o.keep.2.1, k3.2.1, k2.2.1]
+ · rw [o.keep.2.2, k3.2.2, k2.2.2]
+ · rw [o.mem, m3, hm2]
+ · simp only [State.lane]; rw [o.xmm, hy]; exact (hc l hl).1
+ · simp only [State.lane]; rw [o.xmm, hy]; exact (hc l hl).2.1
+ · have : u.lane .xmm10 l = bc (BitVec.allOnes 32) := by simp only [State.lane]; rw [o.xmm, hy]; exact (hc l hl).2.2
+ rw [this, show dword (bc (BitVec.allOnes 32)) e = BitVec.allOnes 32 from bc_ofDwords _ e he]
+ exact iff_of_true (by decide) fun j hj => absurd hj (Nat.not_lt_zero j)
+ refine WP.mono (nlEnd_ok s₄) fun s' ⟨hv, hm', _, _, _⟩ => ⟨?_, ?_⟩
+ · have hnorm : normRq [polyAt s₀.mem (s₀.gpr .rdi)] < arg32 s₀ .rsi ↔ allMsb s₄ = true := by
+ rw [normRq_lt]
+ simp only [allMsb, List.all_eq_true, List.mem_range]
+ constructor
+ · intro h l hl e he
+ rw [(hI.acc l hl e he)]
+ intro j hj
+ have hk : 8 * j + 4 * l + e < 256 := by omega
+ have := h _ hk
+ rw [normZq_lt, polyAt_val hr hk] at this
+ rw [hb2]
+ exact (good_clamp (by rw [← VG.Proof.MlDsa.Round.n_eq] at *; exact hr _ hk)).mpr this
+ · intro h k hk
+ have hk' : k < 256 := by rw [VG.Proof.MlDsa.Round.n_eq] at hk; exact hk
+ have := (hI.acc (k % 8 / 4) (by omega) (k % 4) (by omega)).mp (h _ (by omega) _ (by omega)) (k / 8) (by omega)
+ rw [show 8 * (k / 8) + 4 * (k % 8 / 4) + k % 4 = k by omega, hb2] at this
+ rw [normZq_lt, polyAt_val hr hk]
+ exact (good_clamp (by rw [← VG.Proof.MlDsa.Round.n_eq] at *; exact hr _ hk)).mp this
+ rw [hv]
+ by_cases h : allMsb s₄ = true
+ · rw [ite_eq_left h, ite_eq_left (hnorm.mpr h)]
+ · rw [ite_eq_right h, ite_eq_right (fun h' => h (hnorm.mp h'))]
+ · rw [hm', hI.mem]; exact Frame.refl _ _
+
+theorem normLtY_correct : ∃ t s', Exec isa normLtAvx2 s₀ t s' ∧ abiPreserved s₀ s' ∧ normLtK.post s₀ s' := by
+ obtain ⟨t, s', he, ⟨hv, hf⟩, hk⟩ := WP.keep [.rax, .rcx, .rsi, .rdi] (normLtY_wp hp) (by decide +kernel)
+ exact ⟨t, s', he, abiPreserved_of_exec (by decide +kernel) he (gprPreserved_of hk (by decide) hf (by simp)), hv⟩
+
+end
+
+theorem normLtY_ct : ConstantTime isa normLtK.pre normLtK.pub normLtAvx2 :=
+ VG.Taint.constantTime (A := X86_64.taint) (regsLo [.rdi, .rsp] [.rsi])
+ (fun _ _ _ _ hp => agree_regsLo (fun r hr => by
+ simp only [List.mem_cons, List.not_mem_nil, or_false] at hr
+ rcases hr with rfl | rfl
+ exacts [hp.1, hp.2.1]) fun r hr => by
+ simp only [List.mem_singleton] at hr; subst hr; exact hp.2.2)
+ (by taint_decide)
+
+theorem normLtY_verified : Verified X86_64.target normLtAvx2 (normLtContract X86_64.abi) :=
+ Verified.of_correct (fun _ hp => normLtY_correct hp) normLtY_ct (by
+ round_implies [normLtContract, normLtSig, normLtK, X86_64.abi, X86_64.argRegs] [normSat] using normSat)
+
+end VG.Proof.MlDsa.X86_64.Round
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean
index 0d9b4dcd2..938cf8fa4 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Inst.lean
@@ -63,6 +63,10 @@ def primsWith (B : Impl.MlDsa.X86_64.Arith.Backend) : Prims :=
mulAdd := B.mulAdd
add := B.add
sub := B.sub
+ highBits := B.highBits
+ lowBits := B.lowBits
+ normLt := B.normLt
+ makeHint := B.makeHint
sfx := B.sfx }
theorem nosp_of {c : Prog isa} (h : c.allInstrs (fun i => !Taint.clobbers i .rsp) = true) : NoSp c := by
@@ -141,14 +145,10 @@ def prims_okWith (v : ArithImpl) : PrimsOk (primsWith v.code) signStack where
ball := (⟨16, by decide, Proof.MlDsa.X86_64.Sample.sampleInBall_verified, nosp_of (by decide +kernel),
by decide +kernel⟩ :
Callee _ signStack prims.ball)
- highBits := (⟨0, by decide, Proof.MlDsa.X86_64.Round.highBits_verified, nosp_of (by decide +kernel), by decide +kernel⟩ :
- Callee _ signStack prims.highBits)
- lowBits := (⟨0, by decide, Proof.MlDsa.X86_64.Round.lowBits_verified, nosp_of (by decide +kernel), by decide +kernel⟩ :
- Callee _ signStack prims.lowBits)
- normLt := (⟨0, by decide, Proof.MlDsa.X86_64.Round.normLt_verified, nosp_of (by decide +kernel), by decide +kernel⟩ :
- Callee _ signStack prims.normLt)
- makeHint := (⟨0, by decide, Proof.MlDsa.X86_64.Round.makeHint_verified, nosp_of (by decide +kernel), by decide +kernel⟩ :
- Callee _ signStack prims.makeHint)
+ highBits := calleeOf v.ok.highBits
+ lowBits := calleeOf v.ok.lowBits
+ normLt := calleeOf v.ok.normLt
+ makeHint := calleeOf v.ok.makeHint
simpleBitPack := (⟨0, by decide, Proof.MlDsa.X86_64.Pack.simpleBitPack_verified, nosp_of (by decide +kernel),
by decide +kernel⟩ :
Callee _ signStack prims.simpleBitPack)
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/PrimsC.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/PrimsC.lean
index 8560f2aa7..450aeefea 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/PrimsC.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/PrimsC.lean
@@ -160,9 +160,9 @@ theorem normArgs_ok {bs : List (Reg × Nat)} {f : Ptr} {B : Nat} (hB : B < 2 ^ 3
simp only [normChk, Bool.and_eq_true, decide_eq_true_eq] at hc
simp only [List.all_cons, List.all_nil, Arg.ok, hc.1.2, hc.2, decide_true, Bool.and_true, decide_eq_true hB]
-theorem normCall_ok {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wbs s) {f : Ptr} {B : Nat}
+theorem normCall_ok {nm : String} {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wbs s) {f : Ptr} {B : Nat}
(hB : B < 2 ^ 32) (hc : normChk (rbs ++ wbs) f = true) (hr : Reduced s.mem (pa s f)) :
- WP isa (callP "vg_mldsa_norm_lt" P.normLt [.ptr f, .imm B]) s fun s' => PPostB D s s' [] ∧
+ WP isa (callP nm P.normLt [.ptr f, .imm B]) s fun s' => PPostB D s s' [] ∧
(∀ r ∈ calleeSaved, s'.gpr r = s.gpr r) ∧
(s'.gpr .rax).setWidth 32 = if normRq [polyAt s.mem (pa s f)] < B then 1 else 0 := by
have i1 : inB (rbs ++ wbs) f 1024 = true := by
@@ -179,10 +179,10 @@ theorem normCall_ok {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wb
simp only [e1, e2, Arg.val, er, A.poly' i1 hD, sw32_ofNat hB] at hq
exact hq
-theorem normCall_tr {P : Prims} (hP : PrimsOk P D) {f : Ptr} {B : Nat} (hB : B < 2 ^ 32)
+theorem normCall_tr {nm : String} {P : Prims} (hP : PrimsOk P D) {f : Ptr} {B : Nat} (hB : B < 2 ^ 32)
(hc : normChk (rbs ++ wbs) f = true) :
RelCT isa (fun x y => LRel D rbs wbs x y ∧ Reduced x.mem (pa x f) ∧ Reduced y.mem (pa y f))
- (callP "vg_mldsa_norm_lt" P.normLt [.ptr f, .imm B]) fun _ _ => True := by
+ (callP nm P.normLt [.ptr f, .imm B]) fun _ _ => True := by
have i1 : inB (rbs ++ wbs) f 1024 = true := by
simp only [normChk, Bool.and_eq_true] at hc; exact hc.1.1
refine callP_tr hP.normLt.ver.1 hP.normLt.ver.2.1 (normArgs_ok hB hc)
@@ -238,10 +238,10 @@ theorem hintPre {S : Nat} (hS : S + 8 ≤ D) {s s1 : State} (A : At D rbs wbs s
simp only [List.mem_cons, List.mem_nil_iff, or_false, forall_eq_or_imp, forall_eq]
exact ⟨A.stk hS i1, A.stk hS i2, A.stk hS i3⟩
-theorem hintCall_ok {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wbs s) {z r h : Ptr} {γ : Nat}
+theorem hintCall_ok {nm : String} {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wbs s) {z r h : Ptr} {γ : Nat}
(hγ : γ ∈ gamma2s) (hc : hintChk (rbs ++ wbs) wbs z r h = true) (rz : Reduced s.mem (pa s z))
(rr : Reduced s.mem (pa s r)) :
- WP isa (callP "vg_mldsa_make_hint" P.makeHint [.ptr z, .ptr r, .imm γ, .ptr h]) s fun s' =>
+ WP isa (callP nm P.makeHint [.ptr z, .ptr r, .imm γ, .ptr h]) s fun s' =>
PPostB D s s' [(h, 1024)] ∧ (∀ r ∈ calleeSaved, s'.gpr r = s.gpr r) ∧
HintIs s'.mem (pa s h) 1 [Vector.zipWith (makeHint γ) (polyAt s.mem (pa s z)) (polyAt s.mem (pa s r))] ∧
((s'.gpr .rax).setWidth 32).toNat =
@@ -259,11 +259,11 @@ theorem hintCall_ok {P : Prims} (hP : PrimsOk P D) {s : State} (L : Lay D rbs wb
simp only [e1, e2, e3, e4, Arg.val, hm₂, er, A.poly' i1 hD, A.poly' i2 hD, sw32_ofNat (gamma2_lt hγ)] at hq
exact hq
-theorem hintCall_tr {P : Prims} (hP : PrimsOk P D) {z r h : Ptr} {γ : Nat} (hγ : γ ∈ gamma2s)
+theorem hintCall_tr {nm : String} {P : Prims} (hP : PrimsOk P D) {z r h : Ptr} {γ : Nat} (hγ : γ ∈ gamma2s)
(hc : hintChk (rbs ++ wbs) wbs z r h = true) :
RelCT isa (fun x y => LRel D rbs wbs x y ∧ (Reduced x.mem (pa x z) ∧ Reduced x.mem (pa x r)) ∧
(Reduced y.mem (pa y z) ∧ Reduced y.mem (pa y r)))
- (callP "vg_mldsa_make_hint" P.makeHint [.ptr z, .ptr r, .imm γ, .ptr h]) fun _ _ => True := by
+ (callP nm P.makeHint [.ptr z, .ptr r, .imm γ, .ptr h]) fun _ _ => True := by
obtain ⟨w1, i1, i2, i3, _⟩ := hintChk_spec hc
refine callP_tr hP.makeHint.ver.1 hP.makeHint.ver.2.1 (hintArgs_ok hγ hc)
fun x y x1 y1 ⟨R, ⟨rzx, rrx⟩, ⟨rzy, rry⟩⟩ ⟨⟨hAx, hmx⟩, kx⟩ ⟨⟨hAy, hmy⟩, ky⟩ =>
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean
index 3b2dc921e..8d11355f3 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Sign/Verified.lean
@@ -76,7 +76,8 @@ with no implementation of it. -/
theorem sign_same {m mc : Prog isa → Bool} (hm : Comp m mc) {B : Backend} (h1 : mc B.ntt = true)
(h2 : mc B.invNtt = true) (h3 : mc B.mul = true) (h4 : mc B.mulAdd = true) (h5 : mc B.add = true)
- (h6 : mc B.sub = true) (p : Params) :
+ (h6 : mc B.sub = true) (h8 : mc B.highBits = true) (h9 : mc B.lowBits = true)
+ (h10 : mc B.normLt = true) (h11 : mc B.makeHint = true) (p : Params) :
Same m (Impl.MlDsa.X86_64.Sign.sign (primsWith B) p) (Impl.MlDsa.X86_64.Sign.sign (primsWith .empty) p) := by
unfold Impl.MlDsa.X86_64.Sign.sign
same_tac hm
@@ -93,13 +94,15 @@ include h3
theorem sign_ctl : ctlOk (Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) p) = true :=
ctlOk_of_ctlC (Same.ok (sign_same Comp.ctlC v.ok.ntt.ctl v.ok.invNtt.ctl v.ok.mul.ctl v.ok.mulAdd.ctl
- v.ok.add.ctl v.ok.sub.ctl p) (sign0_ctlC h3))
+ v.ok.add.ctl v.ok.sub.ctl v.ok.highBits.ctl v.ok.lowBits.ctl v.ok.normLt.ctl v.ok.makeHint.ctl p) (sign0_ctlC h3))
theorem sign_spSafe : (Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) p).all (fun i => !isa.writesSp i) = true :=
Code.all_of_allInstrs (Same.ok (sign_same (Comp.all _) (Code.allInstrs_of_all v.ok.ntt.sp)
(Code.allInstrs_of_all v.ok.invNtt.sp) (Code.allInstrs_of_all v.ok.mul.sp)
(Code.allInstrs_of_all v.ok.mulAdd.sp) (Code.allInstrs_of_all v.ok.add.sp)
- (Code.allInstrs_of_all v.ok.sub.sp) p) (sign0_sp h3))
+ (Code.allInstrs_of_all v.ok.sub.sp)
+ (Code.allInstrs_of_all v.ok.highBits.sp) (Code.allInstrs_of_all v.ok.lowBits.sp)
+ (Code.allInstrs_of_all v.ok.normLt.sp) (Code.allInstrs_of_all v.ok.makeHint.sp) p) (sign0_sp h3))
theorem sign_verified :
Verified X86_64.target (Impl.MlDsa.X86_64.Sign.sign (primsWith v.code) p) (signContractT p X86_64.abi signStack) :=
diff --git a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean
index ef0be7ae9..744298288 100644
--- a/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean
+++ b/lean/VerifiedGarbage/Proof/MlDsa/X86_64/Verify/Prims.lean
@@ -54,6 +54,7 @@ def primsWith (B : Arith.Backend) : Prims :=
mul := B.mul
mulAdd := B.mulAdd
sub := B.sub
+ normLt := B.normLt
rej4 := B.rej4
sfx := B.sfx }
@@ -104,10 +105,7 @@ theorem prims_okWith (v : ArithImpl) : PrimsOk (primsWith v.code) where
(Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
(Code.all_of_allInstrs (by lit_decide)) :
CalleeOk prims.hintUnpack _)
- normLt := (CalleeOk.of_verified Proof.MlDsa.X86_64.Round.normLt_verified (by decide)
- (Proof.MlKem.X86_64.nosp_of (by lit_decide)) (by lit_decide) (by lit_decide)
- (Code.all_of_allInstrs (by lit_decide)) :
- CalleeOk prims.normLt _)
+ normLt := calleeOf v.ok.normLt
rej4 := ⟨v.ok.rej4.ver.1, v.ok.rej4.ver.2.1, v.ok.rej4.nosp, v.ok.rej4.depth, v.ok.rej4.ctl, v.ok.rej4.sp⟩
/-- `vg_mldsa*_verify` for the parameter set `p`, calling the x86-64
diff --git a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean
index 5b74ee5fe..5f57b2be7 100644
--- a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean
+++ b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Avx2.lean
@@ -6,8 +6,10 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.BackendAvx2
A variant of `MlDsaArith` on x86-64 (see `TCB/Emit.lean`):
`vg_mldsa_ntt_avx2`, `vg_mldsa_inv_ntt_avx2`,
`vg_mldsa_multiply_ntt_avx2`, `vg_mldsa_multiply_add_ntt_avx2`,
-`vg_mldsa_add_avx2` and `vg_mldsa_sub_avx2`, on eight coefficients at a time
-in AVX2 registers, which need AVX and AVX2; key generation, signing and
+`vg_mldsa_add_avx2`, `vg_mldsa_sub_avx2`, `vg_mldsa_high_bits_avx2`,
+`vg_mldsa_low_bits_avx2`, `vg_mldsa_norm_lt_avx2` and
+`vg_mldsa_make_hint_avx2`, on eight coefficients at a time in AVX2
+registers, which need AVX and AVX2; key generation, signing and
verification calling them need them too.
-/
diff --git a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean
index 51dccc5dd..fd2c251ab 100644
--- a/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean
+++ b/lean/VerifiedGarbage/Variants/MlDsaArith/X86_64/Sse2.lean
@@ -5,8 +5,9 @@ import VerifiedGarbage.Proof.MlDsa.X86_64.Arith.Backend
A variant of `MlDsaArith` on x86-64 (see `TCB/Emit.lean`): `vg_mldsa_ntt`,
`vg_mldsa_inv_ntt`, `vg_mldsa_multiply_ntt`, `vg_mldsa_multiply_add_ntt`,
-`vg_mldsa_add` and `vg_mldsa_sub`, in the baseline ISA (SSE2), which key
-generation, signing and verification call.
+`vg_mldsa_add`, `vg_mldsa_sub`, `vg_mldsa_high_bits`,
+`vg_mldsa_low_bits`, `vg_mldsa_norm_lt` and `vg_mldsa_make_hint`, in the
+baseline ISA (SSE2), which key generation, signing and verification call.
-/
namespace VG.Variants.MlDsaArith.X86_64.Sse2
diff --git a/src/asm/x86_64/mldsa.rs b/src/asm/x86_64/mldsa.rs
index 2997f10a4..41484c2ca 100644
--- a/src/asm/x86_64/mldsa.rs
+++ b/src/asm/x86_64/mldsa.rs
@@ -4588,6 +4588,230 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_low_bits(r: *const [u32; 256], gam
)
}
+/// The CPU features `vg_mldsa_high_bits_avx2` requires (`Artifact.features`).
+pub(crate) const VG_MLDSA_HIGH_BITS_AVX2_FEATURES: &[&str] = &["avx", "avx2"];
+
+/// `HighBits` (FIPS 204 Algorithm 37) of each coefficient of `*r`, with `gamma2` = `γ₂`: writes the `r1`s to `*out`.
+///
+/// Contract: `VG.Spec.MlDsa.highBitsContract`. Constant time: only the pointers and `gamma2` may affect timing, not the data.
+///
+/// The function computes on eight coefficients at a time in AVX2 registers, multiplying by shifts and additions; it needs AVX and AVX2.
+///
+/// # Safety
+///
+/// * `r` must be valid for reads of 1024 bytes.
+/// * `out` must be valid for reads and writes of 1024 bytes.
+/// * `gamma2` must be (q - 1)/88 = 95232 or (q - 1)/32 = 261888.
+/// * Each of the 256 `u32`s of `r` must be less than `q` = 8380417.
+/// * `out` must not overlap `r` (distinct Rust objects never do).
+/// * Neither `r` nor `out` may overlap the return address on the stack, or wrap around the end of the address space (no Rust object does).
+/// * The CPU must support the `avx` and `avx2` target features.
+#[unsafe(naked)]
+pub(crate) unsafe extern "sysv64" fn vg_mldsa_high_bits_avx2(r: *const [u32; 256], gamma2: u32, out: *mut [u32; 256]) {
+ core::arch::naked_asm!(
+ "mov esi, esi",
+ "cmp esi, 261888",
+ "mov r10, rdx",
+ "je 20f",
+ "mov eax, 127",
+ "vmovq xmm8, rax",
+ "vpbroadcastd ymm8, xmm8",
+ "mov eax, 8388608",
+ "vmovq xmm9, rax",
+ "vpbroadcastd ymm9, xmm9",
+ "mov eax, 44",
+ "vmovq xmm10, rax",
+ "vpbroadcastd ymm10, xmm10",
+ "mov eax, 8380417",
+ "vmovq xmm15, rax",
+ "vpbroadcastd ymm15, xmm15",
+ "mov ecx, 32",
+ "22:",
+ "vmovdqu ymm0, YMMWORD PTR [rdi]",
+ "vpaddd ymm0, ymm0, ymm8",
+ "vpsrld ymm0, ymm0, 7",
+ "vmovdqa ymm1, ymm0",
+ "vpslld ymm2, ymm1, 1",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 3",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 10",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 11",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 13",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpaddd ymm0, ymm0, ymm9",
+ "vpsrld ymm0, ymm0, 24",
+ "vpsubd ymm1, ymm0, ymm10",
+ "vpsrad ymm1, ymm1, 31",
+ "vpand ymm0, ymm0, ymm1",
+ "vmovdqu YMMWORD PTR [r10], ymm0",
+ "add rdi, 32",
+ "add r10, 32",
+ "sub rcx, 1",
+ "jne 22b",
+ "jmp 21f",
+ "20:",
+ "mov eax, 127",
+ "vmovq xmm8, rax",
+ "vpbroadcastd ymm8, xmm8",
+ "mov eax, 2097152",
+ "vmovq xmm9, rax",
+ "vpbroadcastd ymm9, xmm9",
+ "mov eax, 16",
+ "vmovq xmm10, rax",
+ "vpbroadcastd ymm10, xmm10",
+ "mov eax, 8380417",
+ "vmovq xmm15, rax",
+ "vpbroadcastd ymm15, xmm15",
+ "mov ecx, 32",
+ "23:",
+ "vmovdqu ymm0, YMMWORD PTR [rdi]",
+ "vpaddd ymm0, ymm0, ymm8",
+ "vpsrld ymm0, ymm0, 7",
+ "vmovdqa ymm1, ymm0",
+ "vpslld ymm2, ymm1, 10",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpaddd ymm0, ymm0, ymm9",
+ "vpsrld ymm0, ymm0, 22",
+ "vpsubd ymm1, ymm0, ymm10",
+ "vpsrad ymm1, ymm1, 31",
+ "vpand ymm0, ymm0, ymm1",
+ "vmovdqu YMMWORD PTR [r10], ymm0",
+ "add rdi, 32",
+ "add r10, 32",
+ "sub rcx, 1",
+ "jne 23b",
+ "21:",
+ "vzeroupper",
+ "ret",
+ )
+}
+
+/// The CPU features `vg_mldsa_low_bits_avx2` requires (`Artifact.features`).
+pub(crate) const VG_MLDSA_LOW_BITS_AVX2_FEATURES: &[&str] = &["avx", "avx2"];
+
+/// `LowBits` (FIPS 204 Algorithm 38) of each coefficient of `*r`, with `gamma2` = `γ₂`: writes the `r0`s, modulo `q` = 8380417, to `*out`.
+///
+/// Contract: `VG.Spec.MlDsa.lowBitsContract`. Constant time: only the pointers and `gamma2` may affect timing, not the data.
+///
+/// The function computes on eight coefficients at a time in AVX2 registers, multiplying by shifts and additions; it needs AVX and AVX2.
+///
+/// # Safety
+///
+/// * `r` must be valid for reads of 1024 bytes.
+/// * `out` must be valid for reads and writes of 1024 bytes.
+/// * `gamma2` must be (q - 1)/88 = 95232 or (q - 1)/32 = 261888.
+/// * Each of the 256 `u32`s of `r` must be less than `q` = 8380417.
+/// * `out` must not overlap `r` (distinct Rust objects never do).
+/// * Neither `r` nor `out` may overlap the return address on the stack, or wrap around the end of the address space (no Rust object does).
+/// * The CPU must support the `avx` and `avx2` target features.
+#[unsafe(naked)]
+pub(crate) unsafe extern "sysv64" fn vg_mldsa_low_bits_avx2(r: *const [u32; 256], gamma2: u32, out: *mut [u32; 256]) {
+ core::arch::naked_asm!(
+ "mov esi, esi",
+ "cmp esi, 261888",
+ "mov r10, rdx",
+ "je 20f",
+ "mov eax, 127",
+ "vmovq xmm8, rax",
+ "vpbroadcastd ymm8, xmm8",
+ "mov eax, 8388608",
+ "vmovq xmm9, rax",
+ "vpbroadcastd ymm9, xmm9",
+ "mov eax, 44",
+ "vmovq xmm10, rax",
+ "vpbroadcastd ymm10, xmm10",
+ "mov eax, 8380417",
+ "vmovq xmm15, rax",
+ "vpbroadcastd ymm15, xmm15",
+ "mov ecx, 32",
+ "22:",
+ "vmovdqu ymm0, YMMWORD PTR [rdi]",
+ "vmovdqa ymm3, ymm0",
+ "vpaddd ymm0, ymm0, ymm8",
+ "vpsrld ymm0, ymm0, 7",
+ "vmovdqa ymm1, ymm0",
+ "vpslld ymm2, ymm1, 1",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 3",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 10",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 11",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 13",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpaddd ymm0, ymm0, ymm9",
+ "vpsrld ymm0, ymm0, 24",
+ "vpsubd ymm1, ymm0, ymm10",
+ "vpsrad ymm1, ymm1, 31",
+ "vpand ymm0, ymm0, ymm1",
+ "vpslld ymm1, ymm0, 11",
+ "vpslld ymm2, ymm0, 13",
+ "vpaddd ymm1, ymm1, ymm2",
+ "vpslld ymm2, ymm0, 14",
+ "vpaddd ymm1, ymm1, ymm2",
+ "vpslld ymm2, ymm0, 15",
+ "vpaddd ymm1, ymm1, ymm2",
+ "vpslld ymm2, ymm0, 17",
+ "vpaddd ymm1, ymm1, ymm2",
+ "vpsubd ymm3, ymm3, ymm1",
+ "vpsrad ymm1, ymm3, 31",
+ "vpand ymm1, ymm1, ymm15",
+ "vpaddd ymm3, ymm3, ymm1",
+ "vmovdqu YMMWORD PTR [r10], ymm3",
+ "add rdi, 32",
+ "add r10, 32",
+ "sub rcx, 1",
+ "jne 22b",
+ "jmp 21f",
+ "20:",
+ "mov eax, 127",
+ "vmovq xmm8, rax",
+ "vpbroadcastd ymm8, xmm8",
+ "mov eax, 2097152",
+ "vmovq xmm9, rax",
+ "vpbroadcastd ymm9, xmm9",
+ "mov eax, 16",
+ "vmovq xmm10, rax",
+ "vpbroadcastd ymm10, xmm10",
+ "mov eax, 8380417",
+ "vmovq xmm15, rax",
+ "vpbroadcastd ymm15, xmm15",
+ "mov ecx, 32",
+ "23:",
+ "vmovdqu ymm0, YMMWORD PTR [rdi]",
+ "vmovdqa ymm3, ymm0",
+ "vpaddd ymm0, ymm0, ymm8",
+ "vpsrld ymm0, ymm0, 7",
+ "vmovdqa ymm1, ymm0",
+ "vpslld ymm2, ymm1, 10",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpaddd ymm0, ymm0, ymm9",
+ "vpsrld ymm0, ymm0, 22",
+ "vpsubd ymm1, ymm0, ymm10",
+ "vpsrad ymm1, ymm1, 31",
+ "vpand ymm0, ymm0, ymm1",
+ "vpslld ymm1, ymm0, 19",
+ "vpslld ymm2, ymm0, 9",
+ "vpsubd ymm1, ymm1, ymm2",
+ "vpsubd ymm3, ymm3, ymm1",
+ "vpsrad ymm1, ymm3, 31",
+ "vpand ymm1, ymm1, ymm15",
+ "vpaddd ymm3, ymm3, ymm1",
+ "vmovdqu YMMWORD PTR [r10], ymm3",
+ "add rdi, 32",
+ "add r10, 32",
+ "sub rcx, 1",
+ "jne 23b",
+ "21:",
+ "vzeroupper",
+ "ret",
+ )
+}
+
/// Returns 1 if the infinity norm of the polynomial `*f` (FIPS 204 §2.3: the largest `|fᵢ mod± q|`) is less than `bound`, and 0 otherwise.
///
/// Contract: `VG.Spec.MlDsa.normLtContract`. Constant time: only the pointer and `bound` may affect timing, not the data.
@@ -4723,6 +4947,249 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa_make_hint(z: *const [u32; 256], r:
)
}
+/// The CPU features `vg_mldsa_norm_lt_avx2` requires (`Artifact.features`).
+pub(crate) const VG_MLDSA_NORM_LT_AVX2_FEATURES: &[&str] = &["avx", "avx2"];
+
+/// Returns 1 if the infinity norm of the polynomial `*f` (FIPS 204 §2.3: the largest `|fᵢ mod± q|`) is less than `bound`, and 0 otherwise.
+///
+/// Contract: `VG.Spec.MlDsa.normLtContract`. Constant time: only the pointer and `bound` may affect timing, not the data.
+///
+/// The function compares eight coefficients at a time in AVX2 registers; it needs AVX and AVX2.
+///
+/// # Safety
+///
+/// * `f` must be valid for reads of 1024 bytes.
+/// * Each of the 256 `u32`s of `f` must be less than `q` = 8380417.
+/// * `f` must not overlap the return address on the stack, or wrap around the end of the address space (no Rust object does).
+/// * The CPU must support the `avx` and `avx2` target features.
+#[unsafe(naked)]
+pub(crate) unsafe extern "sysv64" fn vg_mldsa_norm_lt_avx2(f: *const [u32; 256], bound: u32) -> u32 {
+ core::arch::naked_asm!(
+ "mov esi, esi",
+ "cmp esi, 8380417",
+ "jb 20f",
+ "mov esi, 8380417",
+ "jmp 21f",
+ "20:",
+ "21:",
+ "mov eax, esi",
+ "vmovq xmm8, rax",
+ "vpbroadcastd ymm8, xmm8",
+ "mov eax, 8380417",
+ "sub eax, esi",
+ "vmovq xmm9, rax",
+ "vpbroadcastd ymm9, xmm9",
+ "mov eax, -1",
+ "vmovq xmm10, rax",
+ "vpbroadcastd ymm10, xmm10",
+ "mov ecx, 32",
+ "22:",
+ "vmovdqu ymm0, YMMWORD PTR [rdi]",
+ "vpsubd ymm1, ymm0, ymm8",
+ "vpsubd ymm2, ymm9, ymm0",
+ "vpor ymm1, ymm1, ymm2",
+ "vpand ymm10, ymm10, ymm1",
+ "add rdi, 32",
+ "sub rcx, 1",
+ "jne 22b",
+ "vpsrad ymm10, ymm10, 31",
+ "vpmovmskb eax, ymm10",
+ "vzeroupper",
+ "add rax, 1",
+ "shr rax, 32",
+ "ret",
+ )
+}
+
+/// The CPU features `vg_mldsa_make_hint_avx2` requires (`Artifact.features`).
+pub(crate) const VG_MLDSA_MAKE_HINT_AVX2_FEATURES: &[&str] = &["avx", "avx2"];
+
+/// `MakeHint` (FIPS 204 Algorithm 39) of each pair of coefficients of `*z` and `*r`, with `gamma2` = `γ₂`: writes 1 for true and 0 for false to `*h`, and returns the number of 1s.
+///
+/// Contract: `VG.Spec.MlDsa.makeHintContract`. Constant time: only the pointers and `gamma2` may affect timing, not the data.
+///
+/// The function computes eight hints at a time in AVX2 registers, multiplying by shifts and additions; it needs AVX and AVX2.
+///
+/// # Safety
+///
+/// * `z` must be valid for reads of 1024 bytes.
+/// * `r` must be valid for reads of 1024 bytes.
+/// * `h` must be valid for reads and writes of 1024 bytes.
+/// * `gamma2` must be (q - 1)/88 = 95232 or (q - 1)/32 = 261888.
+/// * Each of the 256 `u32`s of `z` must be less than `q` = 8380417.
+/// * Each of the 256 `u32`s of `r` must be less than `q` = 8380417.
+/// * `h` must not overlap `z` or `r` (distinct Rust objects never do).
+/// * None of `z`, `r` and `h` may overlap the return address on the stack, or wrap around the end of the address space (no Rust object does).
+/// * The CPU must support the `avx` and `avx2` target features.
+#[unsafe(naked)]
+pub(crate) unsafe extern "sysv64" fn vg_mldsa_make_hint_avx2(z: *const [u32; 256], r: *const [u32; 256], gamma2: u32, h: *mut [u32; 256]) -> u32 {
+ core::arch::naked_asm!(
+ "mov edx, edx",
+ "cmp edx, 261888",
+ "mov r10, rcx",
+ "mov r9d, 0",
+ "je 20f",
+ "mov eax, 127",
+ "vmovq xmm8, rax",
+ "vpbroadcastd ymm8, xmm8",
+ "mov eax, 8388608",
+ "vmovq xmm9, rax",
+ "vpbroadcastd ymm9, xmm9",
+ "mov eax, 44",
+ "vmovq xmm10, rax",
+ "vpbroadcastd ymm10, xmm10",
+ "mov eax, 8380417",
+ "vmovq xmm15, rax",
+ "vpbroadcastd ymm15, xmm15",
+ "mov eax, 63",
+ "vmovq xmm11, rax",
+ "vpbroadcastd ymm11, xmm11",
+ "mov ecx, 32",
+ "22:",
+ "vmovdqu ymm0, YMMWORD PTR [rsi]",
+ "vmovdqu ymm5, YMMWORD PTR [rdi]",
+ "vmovdqa ymm4, ymm0",
+ "vpaddd ymm0, ymm0, ymm8",
+ "vpsrld ymm0, ymm0, 7",
+ "vmovdqa ymm1, ymm0",
+ "vpslld ymm2, ymm1, 1",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 3",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 10",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 11",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 13",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpaddd ymm0, ymm0, ymm9",
+ "vpsrld ymm0, ymm0, 24",
+ "vpsubd ymm1, ymm0, ymm10",
+ "vpsrad ymm1, ymm1, 31",
+ "vpand ymm0, ymm0, ymm1",
+ "vmovdqa ymm3, ymm0",
+ "vpaddd ymm0, ymm4, ymm5",
+ "vpsubd ymm0, ymm0, ymm15",
+ "vpsrad ymm1, ymm0, 31",
+ "vpand ymm1, ymm1, ymm15",
+ "vpaddd ymm0, ymm0, ymm1",
+ "vpaddd ymm0, ymm0, ymm8",
+ "vpsrld ymm0, ymm0, 7",
+ "vmovdqa ymm1, ymm0",
+ "vpslld ymm2, ymm1, 1",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 3",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 10",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 11",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpslld ymm2, ymm1, 13",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpaddd ymm0, ymm0, ymm9",
+ "vpsrld ymm0, ymm0, 24",
+ "vpsubd ymm1, ymm0, ymm10",
+ "vpsrad ymm1, ymm1, 31",
+ "vpand ymm0, ymm0, ymm1",
+ "vpxor ymm0, ymm0, ymm3",
+ "vpaddd ymm0, ymm0, ymm11",
+ "vpsrld ymm0, ymm0, 6",
+ "vpslld ymm1, ymm0, 7",
+ "vmovdqu YMMWORD PTR [r10], ymm0",
+ "vpmovmskb eax, ymm1",
+ "mov edx, eax",
+ "shr edx, 4",
+ "add eax, edx",
+ "mov edx, eax",
+ "shr edx, 8",
+ "add eax, edx",
+ "mov edx, eax",
+ "shr edx, 16",
+ "add eax, edx",
+ "and eax, 15",
+ "add r9, rax",
+ "add rdi, 32",
+ "add rsi, 32",
+ "add r10, 32",
+ "sub rcx, 1",
+ "jne 22b",
+ "jmp 21f",
+ "20:",
+ "mov eax, 127",
+ "vmovq xmm8, rax",
+ "vpbroadcastd ymm8, xmm8",
+ "mov eax, 2097152",
+ "vmovq xmm9, rax",
+ "vpbroadcastd ymm9, xmm9",
+ "mov eax, 16",
+ "vmovq xmm10, rax",
+ "vpbroadcastd ymm10, xmm10",
+ "mov eax, 8380417",
+ "vmovq xmm15, rax",
+ "vpbroadcastd ymm15, xmm15",
+ "mov eax, 63",
+ "vmovq xmm11, rax",
+ "vpbroadcastd ymm11, xmm11",
+ "mov ecx, 32",
+ "23:",
+ "vmovdqu ymm0, YMMWORD PTR [rsi]",
+ "vmovdqu ymm5, YMMWORD PTR [rdi]",
+ "vmovdqa ymm4, ymm0",
+ "vpaddd ymm0, ymm0, ymm8",
+ "vpsrld ymm0, ymm0, 7",
+ "vmovdqa ymm1, ymm0",
+ "vpslld ymm2, ymm1, 10",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpaddd ymm0, ymm0, ymm9",
+ "vpsrld ymm0, ymm0, 22",
+ "vpsubd ymm1, ymm0, ymm10",
+ "vpsrad ymm1, ymm1, 31",
+ "vpand ymm0, ymm0, ymm1",
+ "vmovdqa ymm3, ymm0",
+ "vpaddd ymm0, ymm4, ymm5",
+ "vpsubd ymm0, ymm0, ymm15",
+ "vpsrad ymm1, ymm0, 31",
+ "vpand ymm1, ymm1, ymm15",
+ "vpaddd ymm0, ymm0, ymm1",
+ "vpaddd ymm0, ymm0, ymm8",
+ "vpsrld ymm0, ymm0, 7",
+ "vmovdqa ymm1, ymm0",
+ "vpslld ymm2, ymm1, 10",
+ "vpaddd ymm0, ymm0, ymm2",
+ "vpaddd ymm0, ymm0, ymm9",
+ "vpsrld ymm0, ymm0, 22",
+ "vpsubd ymm1, ymm0, ymm10",
+ "vpsrad ymm1, ymm1, 31",
+ "vpand ymm0, ymm0, ymm1",
+ "vpxor ymm0, ymm0, ymm3",
+ "vpaddd ymm0, ymm0, ymm11",
+ "vpsrld ymm0, ymm0, 6",
+ "vpslld ymm1, ymm0, 7",
+ "vmovdqu YMMWORD PTR [r10], ymm0",
+ "vpmovmskb eax, ymm1",
+ "mov edx, eax",
+ "shr edx, 4",
+ "add eax, edx",
+ "mov edx, eax",
+ "shr edx, 8",
+ "add eax, edx",
+ "mov edx, eax",
+ "shr edx, 16",
+ "add eax, edx",
+ "and eax, 15",
+ "add r9, rax",
+ "add rdi, 32",
+ "add rsi, 32",
+ "add r10, 32",
+ "sub rcx, 1",
+ "jne 23b",
+ "21:",
+ "mov rax, r9",
+ "vzeroupper",
+ "ret",
+ )
+}
+
/// `UseHint` (FIPS 204 Algorithm 40) of each pair of coefficients of `*h` (a hint bit: true if it is not 0) and `*r`, with `gamma2` = `γ₂`: writes the results to `*out`.
///
/// Contract: `VG.Spec.MlDsa.useHintContract`. Constant time: only the pointers and `gamma2` may affect timing, not the data.
diff --git a/src/asm/x86_64/mldsa44.rs b/src/asm/x86_64/mldsa44.rs
index b51c4cdf8..b8c0ab667 100644
--- a/src/asm/x86_64/mldsa44.rs
+++ b/src/asm/x86_64/mldsa44.rs
@@ -1685,7 +1685,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov esi, 95232",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 43",
@@ -1698,7 +1698,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov esi, 95232",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 43",
@@ -1711,7 +1711,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov esi, 95232",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 43",
@@ -1724,7 +1724,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov esi, 95232",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 43",
@@ -1840,7 +1840,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov rdi, rbx",
"add rdi, 14336",
"mov esi, 130994",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -1862,7 +1862,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov rdi, rbx",
"add rdi, 15360",
"mov esi, 130994",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -1884,7 +1884,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov rdi, rbx",
"add rdi, 16384",
"mov esi, 130994",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -1906,7 +1906,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov rdi, rbx",
"add rdi, 17408",
"mov esi, 130994",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -1930,11 +1930,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov esi, 95232",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 95154",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -1958,11 +1958,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov esi, 95232",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 95154",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -1986,11 +1986,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov esi, 95232",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 95154",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -2014,11 +2014,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov esi, 95232",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 95154",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 8192",
@@ -2035,7 +2035,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 95232",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -2066,7 +2066,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov edx, 95232",
"mov rcx, rbx",
"add rcx, 10240",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -2085,7 +2085,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 95232",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -2116,7 +2116,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov edx, 95232",
"mov rcx, rbx",
"add rcx, 11264",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -2135,7 +2135,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 95232",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -2166,7 +2166,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov edx, 95232",
"mov rcx, rbx",
"add rcx, 12288",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -2185,7 +2185,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 95232",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -2216,7 +2216,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
"mov edx, 95232",
"mov rcx, rbx",
"add rcx, 13312",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -2315,14 +2315,14 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_sign_avx2(sk: *const [u8; 2560],
vg_mldsa_multiply_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_ntt_avx2,
vg_mldsa_multiply_add_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_add_ntt_avx2,
vg_mldsa_inv_ntt_avx2 = sym super::mldsa::vg_mldsa_inv_ntt_avx2,
- vg_mldsa_high_bits = sym super::mldsa::vg_mldsa_high_bits,
+ vg_mldsa_high_bits_avx2 = sym super::mldsa::vg_mldsa_high_bits_avx2,
vg_mldsa_simple_bit_pack = sym super::mldsa::vg_mldsa_simple_bit_pack,
vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball,
vg_mldsa_add_avx2 = sym super::mldsa::vg_mldsa_add_avx2,
- vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt,
+ vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2,
vg_mldsa_sub_avx2 = sym super::mldsa::vg_mldsa_sub_avx2,
- vg_mldsa_low_bits = sym super::mldsa::vg_mldsa_low_bits,
- vg_mldsa_make_hint = sym super::mldsa::vg_mldsa_make_hint,
+ vg_mldsa_low_bits_avx2 = sym super::mldsa::vg_mldsa_low_bits_avx2,
+ vg_mldsa_make_hint_avx2 = sym super::mldsa::vg_mldsa_make_hint_avx2,
vg_mldsa_bit_pack = sym super::mldsa::vg_mldsa_bit_pack,
vg_mldsa_hint_bit_pack = sym super::mldsa::vg_mldsa_hint_bit_pack,
)
@@ -2385,7 +2385,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312
"mov rdi, rbx",
"add rdi, 16384",
"mov esi, 130994",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 608",
@@ -2398,7 +2398,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312
"mov rdi, rbx",
"add rdi, 17408",
"mov esi, 130994",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 1184",
@@ -2411,7 +2411,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312
"mov rdi, rbx",
"add rdi, 18432",
"mov esi, 130994",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 1760",
@@ -2424,7 +2424,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312
"mov rdi, rbx",
"add rdi, 19456",
"mov esi, 130994",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"test r15d, r15d",
"jne 22f",
@@ -3055,7 +3055,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa44_verify_avx2(pk: *const [u8; 1312
"ret",
vg_mldsa_hint_bit_unpack = sym super::mldsa::vg_mldsa_hint_bit_unpack,
vg_mldsa_bit_unpack = sym super::mldsa::vg_mldsa_bit_unpack,
- vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt,
+ vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2,
vg_mldsa_rej_ntt_poly4_avx2 = sym super::mldsa::vg_mldsa_rej_ntt_poly4_avx2,
vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball,
vg_mldsa_ntt_avx2 = sym super::mldsa::vg_mldsa_ntt_avx2,
diff --git a/src/asm/x86_64/mldsa65.rs b/src/asm/x86_64/mldsa65.rs
index 7d80620cf..6feef85c9 100644
--- a/src/asm/x86_64/mldsa65.rs
+++ b/src/asm/x86_64/mldsa65.rs
@@ -2473,7 +2473,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -2486,7 +2486,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -2499,7 +2499,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -2512,7 +2512,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -2525,7 +2525,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -2538,7 +2538,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -2654,7 +2654,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov rdi, rbx",
"add rdi, 16384",
"mov esi, 524092",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -2676,7 +2676,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov rdi, rbx",
"add rdi, 17408",
"mov esi, 524092",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -2698,7 +2698,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov rdi, rbx",
"add rdi, 18432",
"mov esi, 524092",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -2720,7 +2720,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov rdi, rbx",
"add rdi, 19456",
"mov esi, 524092",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -2742,7 +2742,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov rdi, rbx",
"add rdi, 20480",
"mov esi, 524092",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -2766,11 +2766,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261692",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -2794,11 +2794,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261692",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -2822,11 +2822,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261692",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -2850,11 +2850,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261692",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -2878,11 +2878,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261692",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -2906,11 +2906,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261692",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 8192",
@@ -2927,7 +2927,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -2958,7 +2958,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 10240",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -2977,7 +2977,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -3008,7 +3008,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 11264",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -3027,7 +3027,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -3058,7 +3058,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 12288",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -3077,7 +3077,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -3108,7 +3108,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 13312",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -3127,7 +3127,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -3158,7 +3158,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 14336",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -3177,7 +3177,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -3208,7 +3208,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 15360",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -3315,14 +3315,14 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_sign_avx2(sk: *const [u8; 4032],
vg_mldsa_multiply_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_ntt_avx2,
vg_mldsa_multiply_add_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_add_ntt_avx2,
vg_mldsa_inv_ntt_avx2 = sym super::mldsa::vg_mldsa_inv_ntt_avx2,
- vg_mldsa_high_bits = sym super::mldsa::vg_mldsa_high_bits,
+ vg_mldsa_high_bits_avx2 = sym super::mldsa::vg_mldsa_high_bits_avx2,
vg_mldsa_simple_bit_pack = sym super::mldsa::vg_mldsa_simple_bit_pack,
vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball,
vg_mldsa_add_avx2 = sym super::mldsa::vg_mldsa_add_avx2,
- vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt,
+ vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2,
vg_mldsa_sub_avx2 = sym super::mldsa::vg_mldsa_sub_avx2,
- vg_mldsa_low_bits = sym super::mldsa::vg_mldsa_low_bits,
- vg_mldsa_make_hint = sym super::mldsa::vg_mldsa_make_hint,
+ vg_mldsa_low_bits_avx2 = sym super::mldsa::vg_mldsa_low_bits_avx2,
+ vg_mldsa_make_hint_avx2 = sym super::mldsa::vg_mldsa_make_hint_avx2,
vg_mldsa_bit_pack = sym super::mldsa::vg_mldsa_bit_pack,
vg_mldsa_hint_bit_pack = sym super::mldsa::vg_mldsa_hint_bit_pack,
)
@@ -3385,7 +3385,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952
"mov rdi, rbx",
"add rdi, 16384",
"mov esi, 524092",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 688",
@@ -3398,7 +3398,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952
"mov rdi, rbx",
"add rdi, 17408",
"mov esi, 524092",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 1328",
@@ -3411,7 +3411,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952
"mov rdi, rbx",
"add rdi, 18432",
"mov esi, 524092",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 1968",
@@ -3424,7 +3424,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952
"mov rdi, rbx",
"add rdi, 19456",
"mov esi, 524092",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 2608",
@@ -3437,7 +3437,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952
"mov rdi, rbx",
"add rdi, 20480",
"mov esi, 524092",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"test r15d, r15d",
"jne 22f",
@@ -4471,7 +4471,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa65_verify_avx2(pk: *const [u8; 1952
"ret",
vg_mldsa_hint_bit_unpack = sym super::mldsa::vg_mldsa_hint_bit_unpack,
vg_mldsa_bit_unpack = sym super::mldsa::vg_mldsa_bit_unpack,
- vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt,
+ vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2,
vg_mldsa_rej_ntt_poly4_avx2 = sym super::mldsa::vg_mldsa_rej_ntt_poly4_avx2,
vg_mldsa_rej_ntt_poly = sym super::mldsa::vg_mldsa_rej_ntt_poly,
vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball,
diff --git a/src/asm/x86_64/mldsa87.rs b/src/asm/x86_64/mldsa87.rs
index f8dfeb95d..8e96a3ffe 100644
--- a/src/asm/x86_64/mldsa87.rs
+++ b/src/asm/x86_64/mldsa87.rs
@@ -3698,7 +3698,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -3711,7 +3711,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -3724,7 +3724,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -3737,7 +3737,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -3750,7 +3750,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -3763,7 +3763,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -3776,7 +3776,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -3789,7 +3789,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 6144",
- "call {vg_mldsa_high_bits}",
+ "call {vg_mldsa_high_bits_avx2}",
"mov rdi, rbx",
"add rdi, 6144",
"mov esi, 15",
@@ -3905,7 +3905,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 18432",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -3927,7 +3927,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 19456",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -3949,7 +3949,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 20480",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -3971,7 +3971,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 21504",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -3993,7 +3993,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 22528",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -4015,7 +4015,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 23552",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -4037,7 +4037,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 24576",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -4061,11 +4061,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261768",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -4089,11 +4089,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261768",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -4117,11 +4117,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261768",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -4145,11 +4145,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261768",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -4173,11 +4173,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261768",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -4201,11 +4201,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261768",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -4229,11 +4229,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261768",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 6144",
@@ -4257,11 +4257,11 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov esi, 261888",
"mov rdx, rbx",
"add rdx, 7168",
- "call {vg_mldsa_low_bits}",
+ "call {vg_mldsa_low_bits_avx2}",
"mov rdi, rbx",
"add rdi, 7168",
"mov esi, 261768",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 8192",
@@ -4278,7 +4278,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -4309,7 +4309,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 10240",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -4328,7 +4328,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -4359,7 +4359,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 11264",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -4378,7 +4378,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -4409,7 +4409,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 12288",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -4428,7 +4428,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -4459,7 +4459,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 13312",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -4478,7 +4478,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -4509,7 +4509,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 14336",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -4528,7 +4528,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -4559,7 +4559,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 15360",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -4578,7 +4578,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -4609,7 +4609,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 16384",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -4628,7 +4628,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov rdi, rbx",
"add rdi, 8192",
"mov esi, 261888",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, rbx",
"add rdi, 9216",
@@ -4659,7 +4659,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
"mov edx, 261888",
"mov rcx, rbx",
"add rcx, 17408",
- "call {vg_mldsa_make_hint}",
+ "call {vg_mldsa_make_hint_avx2}",
"mov ecx, DWORD PTR [rbx+904]",
"add ecx, eax",
"mov QWORD PTR [rbx+904], rcx",
@@ -4782,14 +4782,14 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_sign_avx2(sk: *const [u8; 4896],
vg_mldsa_multiply_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_ntt_avx2,
vg_mldsa_multiply_add_ntt_avx2 = sym super::mldsa::vg_mldsa_multiply_add_ntt_avx2,
vg_mldsa_inv_ntt_avx2 = sym super::mldsa::vg_mldsa_inv_ntt_avx2,
- vg_mldsa_high_bits = sym super::mldsa::vg_mldsa_high_bits,
+ vg_mldsa_high_bits_avx2 = sym super::mldsa::vg_mldsa_high_bits_avx2,
vg_mldsa_simple_bit_pack = sym super::mldsa::vg_mldsa_simple_bit_pack,
vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball,
vg_mldsa_add_avx2 = sym super::mldsa::vg_mldsa_add_avx2,
- vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt,
+ vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2,
vg_mldsa_sub_avx2 = sym super::mldsa::vg_mldsa_sub_avx2,
- vg_mldsa_low_bits = sym super::mldsa::vg_mldsa_low_bits,
- vg_mldsa_make_hint = sym super::mldsa::vg_mldsa_make_hint,
+ vg_mldsa_low_bits_avx2 = sym super::mldsa::vg_mldsa_low_bits_avx2,
+ vg_mldsa_make_hint_avx2 = sym super::mldsa::vg_mldsa_make_hint_avx2,
vg_mldsa_bit_pack = sym super::mldsa::vg_mldsa_bit_pack,
vg_mldsa_hint_bit_pack = sym super::mldsa::vg_mldsa_hint_bit_pack,
)
@@ -4852,7 +4852,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592
"mov rdi, rbx",
"add rdi, 16384",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 704",
@@ -4865,7 +4865,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592
"mov rdi, rbx",
"add rdi, 17408",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 1344",
@@ -4878,7 +4878,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592
"mov rdi, rbx",
"add rdi, 18432",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 1984",
@@ -4891,7 +4891,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592
"mov rdi, rbx",
"add rdi, 19456",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 2624",
@@ -4904,7 +4904,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592
"mov rdi, rbx",
"add rdi, 20480",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 3264",
@@ -4917,7 +4917,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592
"mov rdi, rbx",
"add rdi, 21504",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"mov rdi, r13",
"add rdi, 3904",
@@ -4930,7 +4930,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592
"mov rdi, rbx",
"add rdi, 22528",
"mov esi, 524168",
- "call {vg_mldsa_norm_lt}",
+ "call {vg_mldsa_norm_lt_avx2}",
"and r15d, eax",
"test r15d, r15d",
"jne 22f",
@@ -6456,7 +6456,7 @@ pub(crate) unsafe extern "sysv64" fn vg_mldsa87_verify_avx2(pk: *const [u8; 2592
"ret",
vg_mldsa_hint_bit_unpack = sym super::mldsa::vg_mldsa_hint_bit_unpack,
vg_mldsa_bit_unpack = sym super::mldsa::vg_mldsa_bit_unpack,
- vg_mldsa_norm_lt = sym super::mldsa::vg_mldsa_norm_lt,
+ vg_mldsa_norm_lt_avx2 = sym super::mldsa::vg_mldsa_norm_lt_avx2,
vg_mldsa_rej_ntt_poly4_avx2 = sym super::mldsa::vg_mldsa_rej_ntt_poly4_avx2,
vg_mldsa_sample_in_ball = sym super::mldsa::vg_mldsa_sample_in_ball,
vg_mldsa_ntt_avx2 = sym super::mldsa::vg_mldsa_ntt_avx2,