improvement by precomputing du and dv

This commit is contained in:
alu committed 2024-11-17 00:03:34 +00:00
1 parent 55cc4ef781
commit e51a540330
1 file changed
+39 -24
+39 -24
View File
@@ -2,7 +2,7 @@
"cells": [
{
"cell_type": "code",
"execution_count": 44,
"execution_count": 11,
"metadata": {},
"outputs": [],
"source": [
@@ -19,7 +19,7 @@
},
{
"cell_type": "code",
"execution_count": 45,
"execution_count": 12,
"metadata": {},
"outputs": [],
"source": [
@@ -29,7 +29,7 @@
},
{
"cell_type": "code",
"execution_count": 46,
"execution_count": 13,
"metadata": {},
"outputs": [
{
@@ -64,26 +64,28 @@
},
{
"cell_type": "code",
"execution_count": 47,
"execution_count": 14,
"metadata": {},
"outputs": [],
"source": [
"# Generate and encrypt query vector\n",
"x = np.random.rand(encoding_width)\n",
"#cx = np.array([HE_client.encrypt(x[j]) for j in range(len(x))])\n",
"# precompute 1 - |x|^2 for query vector\n",
"du = 1 - x @ x\n",
"cx = HE_client.encrypt(x)"
]
},
{
"cell_type": "code",
"execution_count": 48,
"execution_count": 15,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"[Client] sending HE_client=<ckks Pyfhel obj at 0x779133ff2df0, [pk:Y, sk:Y, rtk:Y, rlk:Y, contx(n=32768, t=0, sec=128, qi=[60, 30, 30, 30, 60], scale=1073741824.0, )]> and cx=<Pyfhel Ciphertext at 0x77913c6fc4f0, scheme=ckks, size=2/2, scale_bits=30, mod_level=0>\n"
"[Client] sending HE_client=<ckks Pyfhel obj at 0x7d957f4dfb70, [pk:Y, sk:Y, rtk:Y, rlk:Y, contx(n=16384, t=0, sec=128, qi=[60, 30, 30, 30, 60], scale=1073741824.0, )]> and cx=<Pyfhel Ciphertext at 0x7d9585e874a0, scheme=ckks, size=2/2, scale_bits=30, mod_level=0>\n"
]
}
],
@@ -109,38 +111,45 @@
},
{
"cell_type": "code",
"execution_count": 49,
"execution_count": 16,
"metadata": {},
"outputs": [],
"source": [
"def hyperbolic_distance_parts(u, v): # returns only the numerator and denominator of the hyperbolic distance formula\n",
" diff = u - v\n",
" du = -(1 - u @ u) # for some reason we need to negate this\n",
" dv = -(1 - v @ v) # for some reason we need to negate this\n",
" return diff @ diff, du * dv # returns the numerator and denominator\n"
" #du = -(1 - u @ u) # for some reason we need to negate this\n",
" #dv = -(1 - v @ v) # for some reason we need to negate this\n",
" #return diff @ diff, du * dv # returns the numerator and denominator\n",
" return diff @ diff # returns the numerator and denominator\n"
]
},
{
"cell_type": "code",
"execution_count": 50,
"execution_count": 17,
"metadata": {},
"outputs": [],
"source": [
"# document matrix containing rows of document encoding vectors\n",
"D = np.random.rand(database_size, encoding_width)"
"D = np.random.rand(database_size, encoding_width)\n",
"# precompute 1 - |D|^2 for each row vector in D\n",
"dv = []\n",
"for i in range(len(D)):\n",
" v = D[i]\n",
" dv.append(1 - v @ v)\n",
"dv = np.array(dv)"
]
},
{
"cell_type": "code",
"execution_count": 51,
"execution_count": 18,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"[Server] received HE_server=<ckks Pyfhel obj at 0x779133f424e0, [pk:Y, sk:-, rtk:Y, rlk:Y, contx(n=32768, t=0, sec=128, qi=[60, 30, 30, 30, 60], scale=1073741824.0, )]> and cx=<Pyfhel Ciphertext at 0x77914a0b1400, scheme=ckks, size=2/2, scale_bits=30, mod_level=0>\n",
"[Server] Distances computed! Responding: res=[(<Pyfhel Ciphertext at 0x779133fef540, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae450, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae310, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae180, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae1d0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae4f0, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae3b0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae130, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae040, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae4a0, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae400, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae540, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae590, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae5e0, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae680, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae360, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae6d0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae720, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae770, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae7c0, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae810, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae860, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae900, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0ae630, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae950, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aeb30, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae9f0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aea40, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0aebd0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aea90, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae9a0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aed10, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0aeae0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aec20, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0ae8b0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aecc0, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0aed60, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aeb80, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0aedb0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aef40, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0aee50, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aec70, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0aa1d0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aa450, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0aaf90, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aa4a0, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0aa590, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aa5e0, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0aa360, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aa2c0, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0aa680, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aa720, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0aa270, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x77914a0aa540, scheme=ckks, size=3/3, scale_bits=60, mod_level=3>), (<Pyfhel Ciphertext at 0x77914a0Line truncated
"[Server] received HE_server=<ckks Pyfhel obj at 0x7d957f4dd3f0, [pk:Y, sk:-, rtk:Y, rlk:Y, contx(n=16384, t=0, sec=128, qi=[60, 30, 30, 30, 60], scale=1073741824.0, )]> and cx=<Pyfhel Ciphertext at 0x7d957f4d7bd0, scheme=ckks, size=2/2, scale_bits=30, mod_level=0>\n",
"[Server] Distances computed! Responding: res=[<Pyfhel Ciphertext at 0x7d9584b26810, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4dfef0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea040, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d9584b26cc0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d9585f052c0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea1d0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea130, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea220, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea270, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea2c0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea310, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea360, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea3b0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea400, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea450, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea4a0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea4f0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea540, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea590, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea5e0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea630, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea680, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea6d0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea720, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea770, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea7c0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea810, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea860, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea8b0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea900, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea950, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea9a0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ea9f0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eaa40, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eaa90, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eaae0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eab30, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eab80, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eabd0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eac20, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eac70, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eacc0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ead10, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ead60, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eadb0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eae00, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eae50, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eaea0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eaef0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eaf40, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4eaf90, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ec040, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ec090, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ec0e0, scheme=ckks, size=2/4, scale_bits=60, mod_level=2>, <Pyfhel Ciphertext at 0x7d957f4ec130, scheme=ckks, size=2/4, scale_bits=60, mod_level=Line truncated
]
}
],
@@ -164,7 +173,7 @@
" # Compute distance bewteen recieved query and D[i]\n",
" res.append(hyperbolic_distance_parts(cx, cd))\n",
"\n",
"s_res = [(res[j][0].to_bytes(), res[j][1].to_bytes()) for j in range(len(res))]\n",
"s_res = [res[j].to_bytes() for j in range(len(res))]\n",
"\n",
"print(f\"[Server] Distances computed! Responding: res={res}\")"
]
@@ -185,7 +194,7 @@
},
{
"cell_type": "code",
"execution_count": 52,
"execution_count": 19,
"metadata": {},
"outputs": [],
"source": [
@@ -197,7 +206,7 @@
},
{
"cell_type": "code",
"execution_count": 53,
"execution_count": 20,
"metadata": {},
"outputs": [],
"source": [
@@ -205,11 +214,14 @@
"#res = HE_client.decrypt(c_res)\n",
"c_res = []\n",
"for i in range(len(s_res)):\n",
" c_num = PyCtxt(pyfhel=HE_server, bytestring=s_res[i][0])\n",
" c_den = PyCtxt(pyfhel=HE_server, bytestring=s_res[i][1])\n",
" c_num = PyCtxt(pyfhel=HE_server, bytestring=s_res[i])\n",
" #c_den = PyCtxt(pyfhel=HE_server, bytestring=s_res[i][1])\n",
" p_num = HE_client.decrypt(c_num)[0]\n",
" p_den = HE_client.decrypt(c_den)[0]\n",
" dist = np.arccosh(1 + 2 * (p_num / p_den))\n",
" #p_den = HE_client.decrypt(c_den)[0]\n",
" #dist = np.arccosh(1 + 2 * (p_num / p_den))\n",
" # compute final score\n",
" dist = np.arccosh(1 + 2 * (p_num / (du * dv[i])))\n",
" #print(dist)\n",
" c_res.append(dist)\n",
"\n",
"# Checking result\n",
@@ -218,8 +230,11 @@
"for i in range(len(c_res)):\n",
" result = c_res[i]\n",
" expected = expected_res[i]\n",
" #print(f\"got: {result}, expected: {expected}\")\n",
" assert np.abs(result - expected) < 1e-3"
" if np.abs(result - expected) < 1e-3:\n",
" pass\n",
" else:\n",
" print(f\"got: {result}, expected: {expected}\")\n",
" assert False"
]
}
],