@@ -15,9 +15,7 @@ import Mathlib.Algebra.BigOperators.Fin
1515
1616/-!
1717# Bit operations on natural numbers
18-
1918-/
20-
2119namespace Nat
2220
2321-- Note: this is already done with `Nat.sub_add_eq_max`
@@ -26,6 +24,19 @@ theorem max_eq_add_sub {m n : Nat} : Nat.max m n = m + (n - m) := by
2624 · simp [Nat.sub_eq_zero_of_le, h]
2725 · simp only [Nat.max_eq_right (Nat.le_of_not_le h), Nat.add_sub_of_le (Nat.le_of_not_le h)]
2826
27+ theorem sub_add_eq_sub_sub_rev (a b c : Nat) (h1 : c ≤ b) (h2 : b ≤ a) :
28+ a - b + c = a - (b - c) := by
29+ conv =>
30+ rhs
31+ rw [← Nat.sub_add_cancel h2]
32+ rw [Nat.add_sub_assoc (Nat.sub_le b c)]
33+ rw [Nat.sub_sub_self h1]
34+
35+ @[simp]
36+ lemma lt_add_of_pos_right_of_le (a b c : ℕ) [NeZero c] (h : a ≤ b) : a < b + c := by
37+ apply Nat.lt_of_le_of_lt (n:=a) (m:=b) (k:=b + c) h
38+ apply Nat.lt_add_of_pos_right (by exact pos_of_neZero c)
39+
2940/--
3041Returns the `k`-th least significant bit of a natural number `n` as a natural number (in `{0, 1}`).
3142
@@ -135,7 +146,8 @@ lemma eq_iff_eq_all_getBits {n m : ℕ} : n = m ↔ ∀ k, getBit k n = getBit k
135146 rw [h_all_getBits k]
136147
137148lemma shiftRight_and_one_distrib {n m k : ℕ} :
138- (n &&& m) >>> k &&& 1 = ((n >>> k) &&& 1 ) &&& ((m >>> k) &&& 1 ) := by
149+ Nat.getBit k (n &&& m) = Nat.getBit k n &&& Nat.getBit k m := by
150+ unfold getBit
139151 rw [Nat.shiftRight_and_distrib]
140152 conv =>
141153 lhs
@@ -145,26 +157,26 @@ lemma shiftRight_and_one_distrib {n m k : ℕ} :
145157 rw [Nat.and_assoc]
146158
147159lemma and_eq_zero_iff_and_each_getBit_eq_zero {n m : ℕ} :
148- n &&& m = 0 ↔ ∀ k, ((n >>> k) &&& 1 ) &&& ((m >>> k) &&& 1 ) = 0 := by
160+ n &&& m = 0 ↔ ∀ k, Nat.getBit k n &&& Nat.getBit k m = 0 := by
149161 constructor
150162 · intro h_and_zero
151163 intro k
152164 have h_k := shiftRight_and_one_distrib (n := n) (m := m) (k := k)
153165 rw [←h_k]
154- rw [h_and_zero, Nat.zero_shiftRight, Nat.zero_and]
166+ rw [h_and_zero, getBit, Nat.zero_shiftRight, Nat.zero_and]
155167 · intro h_forall_k -- h_forall_k : ∀ (k : ℕ), n >>> k &&& 1 &&& (m >>> k &&& 1) = 0
156168 apply eq_iff_eq_all_getBits.mpr
157169 unfold getBit
158170 intro k
159171 -- ⊢ (n &&& m) >>> k &&& 1 = 0 >>> k &&& 1
160172 have h_forall_k_eq : ∀ k, ((n &&& m) >>> k) &&& 1 = 0 := by
161173 intro k
162- rw [shiftRight_and_one_distrib]
174+ rw [←getBit, shiftRight_and_one_distrib]
163175 exact h_forall_k k
164176 rw [h_forall_k_eq k]
165177 rw [Nat.zero_shiftRight, Nat.zero_and]
166178
167- lemma getBit_two_pow {i k: ℕ} : (getBit k (2 ^i) = if i == k then 1 else 0 ) := by
179+ lemma getBit_two_pow {i k : ℕ} : (getBit k (2 ^i) = if i == k then 1 else 0 ) := by
168180 have h_two_pow_i: 2 ^i = 1 <<< i := by
169181 simp only [Nat.shiftLeft_eq, one_mul]
170182 rw [getBit, h_two_pow_i]
@@ -215,19 +227,20 @@ lemma getBit_two_pow {i k: ℕ} : (getBit k (2^i) = if i == k then 1 else 0) :=
215227 omega
216228 rw [h_res]
217229
218- lemma and_two_pow_eq_zero_of_getBit_0 {n i : ℕ} (h_getBit: getBit i n = 0 ) : n &&& (2 ^ i) = 0 := by
230+ lemma and_two_pow_eq_zero_of_getBit_0 {n i : ℕ} (h_getBit : getBit i n = 0 )
231+ : n &&& (2 ^ i) = 0 := by
219232 apply and_eq_zero_iff_and_each_getBit_eq_zero.mpr
220233 intro k
221234 have h_getBit_two_pow := getBit_two_pow (i := i) (k := k)
222235 if h_k: k = i then
223236 simp only [h_k, BEq.rfl, ↓reduceIte] at h_getBit_two_pow
224237 rw [getBit, h_k.symm] at h_getBit
225- rw [h_getBit, Nat.zero_and]
238+ rw [getBit, h_getBit, Nat.zero_and]
226239 else
227240 push_neg at h_k
228241 simp only [beq_iff_eq, h_k.symm, ↓reduceIte] at h_getBit_two_pow
229242 rw [getBit] at h_getBit_two_pow
230- rw [h_getBit_two_pow]
243+ rw [getBit, getBit, h_getBit_two_pow]
231244 rw [Nat.and_zero]
232245
233246lemma and_two_pow_eq_two_pow_of_getBit_1 {n i : ℕ} (h_getBit: getBit i n = 1 ) :
@@ -1203,4 +1216,92 @@ lemma getBit_of_binaryFinMapToNat {n : ℕ} (m : Fin n → ℕ) (h_binary: ∀ j
12031216 simp only [ite_eq_right_iff, one_ne_zero, imp_false, ne_eq]
12041217 omega
12051218
1219+ /-- Middle bits: take `len` bits starting at `offset` from `n`. -/
1220+ def getMiddleBits (offset len n : ℕ) : ℕ :=
1221+ getLowBits (numLowBits:=len) (n:=n >>> offset)
1222+
1223+ /-- Bit-level characterization of middle bits. -/
1224+ lemma getBit_of_middleBits {n offset len k : ℕ} :
1225+ getBit k (getMiddleBits offset len n) =
1226+ if k < len then getBit (k + offset) n else 0 := by
1227+ unfold getMiddleBits
1228+ -- use existing lemmas
1229+ rw [getBit_of_lowBits, getBit_of_shiftRight]
1230+
1231+ /-- Middle bits are strictly less than `2^len`. -/
1232+ lemma getMiddleBits_lt_two_pow {n offset len : ℕ} :
1233+ getMiddleBits offset len n < 2 ^ len := by
1234+ unfold getMiddleBits
1235+ exact getLowBits_lt_two_pow (n := n >>> offset) len
1236+
1237+ /-- Middle bits as a modulus form. -/
1238+ lemma getMiddleBits_eq_mod {n offset len : ℕ} :
1239+ getMiddleBits offset len n = (n >>> offset) % (2 ^ len) := by
1240+ unfold getMiddleBits
1241+ exact getLowBits_eq_mod_two_pow (n := n >>> offset) (numLowBits := len)
1242+
1243+ lemma and_shl_eq_zero_of_lt_two_pow {a n b : ℕ} (hb : b < 2 ^ n) : (a <<< n) &&& b = 0 := by
1244+ apply Nat.and_eq_zero_iff_and_each_getBit_eq_zero.mpr
1245+ intro k
1246+ rw [getBit_of_shiftLeft]
1247+ rw [getBit_of_lt_two_pow (a := ⟨b, hb⟩)]
1248+ split_ifs with h_k_lt_n
1249+ · simp only [Nat.zero_and]
1250+ · simp only [Nat.and_zero]
1251+
1252+ /-- Concatenate high (length m) and low (length n) using shifts. -/
1253+ def joinBits {n m : ℕ} (low : Fin (2 ^ n)) (high : Fin (2 ^ m)) : Fin (2 ^ (m+n)) :=
1254+ ⟨(high.val <<< n) ||| low.val, by
1255+ have h_and_zero := and_shl_eq_zero_of_lt_two_pow (a := high.val) (b := low.val) (hb := low.isLt)
1256+ rw [←Nat.sum_of_and_eq_zero_is_or h_and_zero]
1257+ rw [Nat.shiftLeft_eq, mul_comm, Nat.pow_add]
1258+ -- ⊢ 2 ^ n * ↑high + ↑low < 2 ^ m * 2 ^ n
1259+ calc
1260+ 2 ^ n * high.val + low.val < 2 ^ n * high.val + 2 ^ n := by
1261+ exact Nat.add_lt_add_left low.isLt _
1262+ _ = 2 ^ n * (high.val + 1 ) := by rw [Nat.mul_add, Nat.mul_one]
1263+ _ ≤ 2 ^ n * (2 ^ m) := by -- `high.val < 2^m` implies `high.val + 1 ≤ 2^m`
1264+ exact Nat.mul_le_mul_left _ (Nat.succ_le_of_lt high.isLt)
1265+ _ = 2 ^ m * 2 ^ n := by rw [mul_comm]
1266+ ⟩
1267+
1268+ /-- Bit characterization: below cut use low, above cut use high. -/
1269+ lemma getBit_joinBits {n m k : ℕ} (low : Fin (2 ^ n)) (high : Fin (2 ^ m)) :
1270+ getBit k (joinBits low high).val =
1271+ if k < n then getBit k low.val else getBit (k - n) high.val := by
1272+ unfold joinBits
1273+ dsimp
1274+ rw [getBit_of_or]
1275+ rw [getBit_of_shiftLeft]
1276+ rw [getBit_of_lt_two_pow (a := low)]
1277+ split_ifs with h_k
1278+ · simp only [zero_or]
1279+ · simp only [Nat.or_zero]
1280+
1281+ /-- Low n bits of joinBits are exactly low. -/
1282+ lemma getLowBits_joinBits {n m : ℕ} (low : Fin (2 ^ n)) (high : Fin (2 ^ m)) :
1283+ getLowBits n (joinBits low high).val = low.val := by
1284+ unfold joinBits
1285+ dsimp
1286+ rw [getLowBits_eq_mod_two_pow]
1287+ have h_and_zero := and_shl_eq_zero_of_lt_two_pow (a := high.val) (b := low.val) (hb := low.isLt)
1288+ rw [←Nat.sum_of_and_eq_zero_is_or h_and_zero]
1289+ rw [Nat.shiftLeft_eq, mul_comm, add_mod, mul_mod, mod_self, zero_mul, zero_mod, zero_add]
1290+ rw [Nat.mod_mod]
1291+ exact Nat.mod_eq_of_lt low.isLt
1292+
1293+ /-- Dropping low n bits by shifting right recovers high. -/
1294+ lemma getHighBits_no_shl_joinBits {n m : ℕ} (low : Fin (2 ^ n)) (high : Fin (2 ^ m)) :
1295+ getHighBits_no_shl n (joinBits low high).val = high.val := by
1296+ unfold joinBits getHighBits_no_shl
1297+ dsimp
1298+ have h_and_zero := and_shl_eq_zero_of_lt_two_pow (a := high.val) (b := low.val) (hb := low.isLt)
1299+ rw [←Nat.sum_of_and_eq_zero_is_or h_and_zero]
1300+ rw [Nat.add_shiftRight_distrib h_and_zero]
1301+ rw [Nat.shiftLeft_shiftRight]
1302+ rw [Nat.shiftRight_eq_div_pow]
1303+ have h: low.val/2 ^n = 0 := by
1304+ apply Nat.div_eq_zero_iff_lt (x:=low) (k:=2 ^n) (h:=by exact Nat.two_pow_pos n).mpr (by omega)
1305+ simp only [h, add_zero]
1306+
12061307end Nat
0 commit comments