set up torch AoA estimation

This commit is contained in:
Christos Falas 2024-12-31 16:45:42 +00:00
parent 814321f524
commit d6072d6fd5
No known key found for this signature in database
7 changed files with 272 additions and 36 deletions

View File

@ -15,6 +15,7 @@ dependencies = [
"h5py>=3.12.1", "h5py>=3.12.1",
"pyyaml>=6.0.2", "pyyaml>=6.0.2",
"pydantic>=2.10.6", "pydantic>=2.10.6",
"torch"
] ]
[project.scripts] [project.scripts]

223
uv.lock
View File

@ -1,5 +1,9 @@
version = 1 version = 1
requires-python = ">=3.11" requires-python = ">=3.11"
resolution-markers = [
"python_full_version >= '3.12'",
"python_full_version < '3.12'",
]
[[package]] [[package]]
name = "annotated-types" name = "annotated-types"
@ -100,6 +104,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/e7/05/c19819d5e3d95294a6f5947fb9b9629efb316b96de511b418c53d245aae6/cycler-0.12.1-py3-none-any.whl", hash = "sha256:85cef7cff222d8644161529808465972e51340599459b8ac3ccbac5a854e0d30", size = 8321 }, { url = "https://files.pythonhosted.org/packages/e7/05/c19819d5e3d95294a6f5947fb9b9629efb316b96de511b418c53d245aae6/cycler-0.12.1-py3-none-any.whl", hash = "sha256:85cef7cff222d8644161529808465972e51340599459b8ac3ccbac5a854e0d30", size = 8321 },
] ]
[[package]]
name = "filelock"
version = "3.17.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/dc/9c/0b15fb47b464e1b663b1acd1253a062aa5feecb07d4e597daea542ebd2b5/filelock-3.17.0.tar.gz", hash = "sha256:ee4e77401ef576ebb38cd7f13b9b28893194acc20a8e68e18730ba9c0e54660e", size = 18027 }
wheels = [
{ url = "https://files.pythonhosted.org/packages/89/ec/00d68c4ddfedfe64159999e5f8a98fb8442729a63e2077eb9dcd89623d27/filelock-3.17.0-py3-none-any.whl", hash = "sha256:533dc2f7ba78dc2f0f531fc6c4940addf7b70a481e269a5a3b93be94ffbe8338", size = 16164 },
]
[[package]] [[package]]
name = "flask" name = "flask"
version = "3.1.0" version = "3.1.0"
@ -162,6 +175,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/1e/6a/6afc55d75036b8d3fe5ceaea2e8da2c04e8f3b298325de73a35f098cb9a8/fonttools-4.55.6-py3-none-any.whl", hash = "sha256:d20ab5a78d0536c26628eaadba661e7ae2427b1e5c748a0a510a44d914e1b155", size = 1112524 }, { url = "https://files.pythonhosted.org/packages/1e/6a/6afc55d75036b8d3fe5ceaea2e8da2c04e8f3b298325de73a35f098cb9a8/fonttools-4.55.6-py3-none-any.whl", hash = "sha256:d20ab5a78d0536c26628eaadba661e7ae2427b1e5c748a0a510a44d914e1b155", size = 1112524 },
] ]
[[package]]
name = "fsspec"
version = "2024.12.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/ee/11/de70dee31455c546fbc88301971ec03c328f3d1138cfba14263f651e9551/fsspec-2024.12.0.tar.gz", hash = "sha256:670700c977ed2fb51e0d9f9253177ed20cbde4a3e5c0283cc5385b5870c8533f", size = 291600 }
wheels = [
{ url = "https://files.pythonhosted.org/packages/de/86/5486b0188d08aa643e127774a99bac51ffa6cf343e3deb0583956dca5b22/fsspec-2024.12.0-py3-none-any.whl", hash = "sha256:b520aed47ad9804237ff878b504267a3b0b441e97508bd6d2d8774e3db85cee2", size = 183862 },
]
[[package]] [[package]]
name = "h11" name = "h11"
version = "0.14.0" version = "0.14.0"
@ -405,6 +427,24 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/b3/38/89ba8ad64ae25be8de66a6d463314cf1eb366222074cfda9ee839c56a4b4/mdurl-0.1.2-py3-none-any.whl", hash = "sha256:84008a41e51615a49fc9966191ff91509e3c40b939176e643fd50a5c2196b8f8", size = 9979 }, { url = "https://files.pythonhosted.org/packages/b3/38/89ba8ad64ae25be8de66a6d463314cf1eb366222074cfda9ee839c56a4b4/mdurl-0.1.2-py3-none-any.whl", hash = "sha256:84008a41e51615a49fc9966191ff91509e3c40b939176e643fd50a5c2196b8f8", size = 9979 },
] ]
[[package]]
name = "mpmath"
version = "1.3.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/e0/47/dd32fa426cc72114383ac549964eecb20ecfd886d1e5ccf5340b55b02f57/mpmath-1.3.0.tar.gz", hash = "sha256:7a28eb2a9774d00c7bc92411c19a89209d5da7c4c9a9e227be8330a23a25b91f", size = 508106 }
wheels = [
{ url = "https://files.pythonhosted.org/packages/43/e3/7d92a15f894aa0c9c4b49b8ee9ac9850d6e63b03c9c32c0367a13ae62209/mpmath-1.3.0-py3-none-any.whl", hash = "sha256:a0b2b9fe80bbcd81a6647ff13108738cfb482d481d826cc0e02f5b35e5c88d2c", size = 536198 },
]
[[package]]
name = "networkx"
version = "3.4.2"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/fd/1d/06475e1cd5264c0b870ea2cc6fdb3e37177c1e565c43f56ff17a10e3937f/networkx-3.4.2.tar.gz", hash = "sha256:307c3669428c5362aab27c8a1260aa8f47c4e91d3891f48be0141738d8d053e1", size = 2151368 }
wheels = [
{ url = "https://files.pythonhosted.org/packages/b9/54/dd730b32ea14ea797530a4479b2ed46a6fb250f682a9cfb997e968bf0261/networkx-3.4.2-py3-none-any.whl", hash = "sha256:df5d4365b724cf81b8c6a7312509d0c22386097011ad1abe274afd5e9d3bbc5f", size = 1723263 },
]
[[package]] [[package]]
name = "numpy" name = "numpy"
version = "2.2.2" version = "2.2.2"
@ -453,6 +493,116 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/80/94/cd9e9b04012c015cb6320ab3bf43bc615e248dddfeb163728e800a5d96f0/numpy-2.2.2-cp313-cp313t-win_amd64.whl", hash = "sha256:97b974d3ba0fb4612b77ed35d7627490e8e3dff56ab41454d9e8b23448940576", size = 12696208 }, { url = "https://files.pythonhosted.org/packages/80/94/cd9e9b04012c015cb6320ab3bf43bc615e248dddfeb163728e800a5d96f0/numpy-2.2.2-cp313-cp313t-win_amd64.whl", hash = "sha256:97b974d3ba0fb4612b77ed35d7627490e8e3dff56ab41454d9e8b23448940576", size = 12696208 },
] ]
[[package]]
name = "nvidia-cublas-cu12"
version = "12.4.5.8"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/ae/71/1c91302526c45ab494c23f61c7a84aa568b8c1f9d196efa5993957faf906/nvidia_cublas_cu12-12.4.5.8-py3-none-manylinux2014_x86_64.whl", hash = "sha256:2fc8da60df463fdefa81e323eef2e36489e1c94335b5358bcb38360adf75ac9b", size = 363438805 },
]
[[package]]
name = "nvidia-cuda-cupti-cu12"
version = "12.4.127"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/67/42/f4f60238e8194a3106d06a058d494b18e006c10bb2b915655bd9f6ea4cb1/nvidia_cuda_cupti_cu12-12.4.127-py3-none-manylinux2014_x86_64.whl", hash = "sha256:9dec60f5ac126f7bb551c055072b69d85392b13311fcc1bcda2202d172df30fb", size = 13813957 },
]
[[package]]
name = "nvidia-cuda-nvrtc-cu12"
version = "12.4.127"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/2c/14/91ae57cd4db3f9ef7aa99f4019cfa8d54cb4caa7e00975df6467e9725a9f/nvidia_cuda_nvrtc_cu12-12.4.127-py3-none-manylinux2014_x86_64.whl", hash = "sha256:a178759ebb095827bd30ef56598ec182b85547f1508941a3d560eb7ea1fbf338", size = 24640306 },
]
[[package]]
name = "nvidia-cuda-runtime-cu12"
version = "12.4.127"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/ea/27/1795d86fe88ef397885f2e580ac37628ed058a92ed2c39dc8eac3adf0619/nvidia_cuda_runtime_cu12-12.4.127-py3-none-manylinux2014_x86_64.whl", hash = "sha256:64403288fa2136ee8e467cdc9c9427e0434110899d07c779f25b5c068934faa5", size = 883737 },
]
[[package]]
name = "nvidia-cudnn-cu12"
version = "9.1.0.70"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "nvidia-cublas-cu12" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/9f/fd/713452cd72343f682b1c7b9321e23829f00b842ceaedcda96e742ea0b0b3/nvidia_cudnn_cu12-9.1.0.70-py3-none-manylinux2014_x86_64.whl", hash = "sha256:165764f44ef8c61fcdfdfdbe769d687e06374059fbb388b6c89ecb0e28793a6f", size = 664752741 },
]
[[package]]
name = "nvidia-cufft-cu12"
version = "11.2.1.3"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "nvidia-nvjitlink-cu12" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/27/94/3266821f65b92b3138631e9c8e7fe1fb513804ac934485a8d05776e1dd43/nvidia_cufft_cu12-11.2.1.3-py3-none-manylinux2014_x86_64.whl", hash = "sha256:f083fc24912aa410be21fa16d157fed2055dab1cc4b6934a0e03cba69eb242b9", size = 211459117 },
]
[[package]]
name = "nvidia-curand-cu12"
version = "10.3.5.147"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/8a/6d/44ad094874c6f1b9c654f8ed939590bdc408349f137f9b98a3a23ccec411/nvidia_curand_cu12-10.3.5.147-py3-none-manylinux2014_x86_64.whl", hash = "sha256:a88f583d4e0bb643c49743469964103aa59f7f708d862c3ddb0fc07f851e3b8b", size = 56305206 },
]
[[package]]
name = "nvidia-cusolver-cu12"
version = "11.6.1.9"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "nvidia-cublas-cu12" },
{ name = "nvidia-cusparse-cu12" },
{ name = "nvidia-nvjitlink-cu12" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/3a/e1/5b9089a4b2a4790dfdea8b3a006052cfecff58139d5a4e34cb1a51df8d6f/nvidia_cusolver_cu12-11.6.1.9-py3-none-manylinux2014_x86_64.whl", hash = "sha256:19e33fa442bcfd085b3086c4ebf7e8debc07cfe01e11513cc6d332fd918ac260", size = 127936057 },
]
[[package]]
name = "nvidia-cusparse-cu12"
version = "12.3.1.170"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "nvidia-nvjitlink-cu12" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/db/f7/97a9ea26ed4bbbfc2d470994b8b4f338ef663be97b8f677519ac195e113d/nvidia_cusparse_cu12-12.3.1.170-py3-none-manylinux2014_x86_64.whl", hash = "sha256:ea4f11a2904e2a8dc4b1833cc1b5181cde564edd0d5cd33e3c168eff2d1863f1", size = 207454763 },
]
[[package]]
name = "nvidia-nccl-cu12"
version = "2.21.5"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/df/99/12cd266d6233f47d00daf3a72739872bdc10267d0383508b0b9c84a18bb6/nvidia_nccl_cu12-2.21.5-py3-none-manylinux2014_x86_64.whl", hash = "sha256:8579076d30a8c24988834445f8d633c697d42397e92ffc3f63fa26766d25e0a0", size = 188654414 },
]
[[package]]
name = "nvidia-nvjitlink-cu12"
version = "12.4.127"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/ff/ff/847841bacfbefc97a00036e0fce5a0f086b640756dc38caea5e1bb002655/nvidia_nvjitlink_cu12-12.4.127-py3-none-manylinux2014_x86_64.whl", hash = "sha256:06b3b9b25bf3f8af351d664978ca26a16d2c5127dbd53c0497e28d1fb9611d57", size = 21066810 },
]
[[package]]
name = "nvidia-nvtx-cu12"
version = "12.4.127"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/87/20/199b8713428322a2f22b722c62b8cc278cc53dffa9705d744484b5035ee9/nvidia_nvtx_cu12-12.4.127-py3-none-manylinux2014_x86_64.whl", hash = "sha256:781e950d9b9f60d8241ccea575b32f5105a5baf4c2351cab5256a24869f12a1a", size = 99144 },
]
[[package]] [[package]]
name = "optype" name = "optype"
version = "0.9.0" version = "0.9.0"
@ -746,6 +896,15 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/71/1f/f1db21491a40923251244ee8de4f71ce95dee2ed83c2b49ad479c2ea5189/scipy_stubs-1.15.1.0-py3-none-any.whl", hash = "sha256:f6e8f8dfd2aaa343705c46ac1fdc69556ea8a0a9b0645cb3af4dd652665a452f", size = 454057 }, { url = "https://files.pythonhosted.org/packages/71/1f/f1db21491a40923251244ee8de4f71ce95dee2ed83c2b49ad479c2ea5189/scipy_stubs-1.15.1.0-py3-none-any.whl", hash = "sha256:f6e8f8dfd2aaa343705c46ac1fdc69556ea8a0a9b0645cb3af4dd652665a452f", size = 454057 },
] ]
[[package]]
name = "setuptools"
version = "75.8.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/92/ec/089608b791d210aec4e7f97488e67ab0d33add3efccb83a056cbafe3a2a6/setuptools-75.8.0.tar.gz", hash = "sha256:c5afc8f407c626b8313a86e10311dd3f661c6cd9c09d4bf8c15c0e11f9f2b0e6", size = 1343222 }
wheels = [
{ url = "https://files.pythonhosted.org/packages/69/8a/b9dc7678803429e4a3bc9ba462fa3dd9066824d3c607490235c6a796be5a/setuptools-75.8.0-py3-none-any.whl", hash = "sha256:e3982f444617239225d675215d51f6ba05f845d4eec313da4418fdbb56fb27e3", size = 1228782 },
]
[[package]] [[package]]
name = "shellingham" name = "shellingham"
version = "1.5.4" version = "1.5.4"
@ -776,6 +935,68 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/b7/ce/149a00dd41f10bc29e5921b496af8b574d8413afcd5e30dfa0ed46c2cc5e/six-1.17.0-py2.py3-none-any.whl", hash = "sha256:4721f391ed90541fddacab5acf947aa0d3dc7d27b2e1e8eda2be8970586c3274", size = 11050 }, { url = "https://files.pythonhosted.org/packages/b7/ce/149a00dd41f10bc29e5921b496af8b574d8413afcd5e30dfa0ed46c2cc5e/six-1.17.0-py2.py3-none-any.whl", hash = "sha256:4721f391ed90541fddacab5acf947aa0d3dc7d27b2e1e8eda2be8970586c3274", size = 11050 },
] ]
[[package]]
name = "sympy"
version = "1.13.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "mpmath" },
]
sdist = { url = "https://files.pythonhosted.org/packages/ca/99/5a5b6f19ff9f083671ddf7b9632028436167cd3d33e11015754e41b249a4/sympy-1.13.1.tar.gz", hash = "sha256:9cebf7e04ff162015ce31c9c6c9144daa34a93bd082f54fd8f12deca4f47515f", size = 7533040 }
wheels = [
{ url = "https://files.pythonhosted.org/packages/b2/fe/81695a1aa331a842b582453b605175f419fe8540355886031328089d840a/sympy-1.13.1-py3-none-any.whl", hash = "sha256:db36cdc64bf61b9b24578b6f7bab1ecdd2452cf008f34faa33776680c26d66f8", size = 6189177 },
]
[[package]]
name = "torch"
version = "2.5.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "filelock" },
{ name = "fsspec" },
{ name = "jinja2" },
{ name = "networkx" },
{ name = "nvidia-cublas-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-cuda-cupti-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-cuda-nvrtc-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-cuda-runtime-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-cudnn-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-cufft-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-curand-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-cusolver-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-cusparse-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-nccl-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-nvjitlink-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "nvidia-nvtx-cu12", marker = "platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "setuptools", marker = "python_full_version >= '3.12'" },
{ name = "sympy" },
{ name = "triton", marker = "python_full_version < '3.13' and platform_machine == 'x86_64' and sys_platform == 'linux'" },
{ name = "typing-extensions" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/d1/35/e8b2daf02ce933e4518e6f5682c72fd0ed66c15910ea1fb4168f442b71c4/torch-2.5.1-cp311-cp311-manylinux1_x86_64.whl", hash = "sha256:de5b7d6740c4b636ef4db92be922f0edc425b65ed78c5076c43c42d362a45457", size = 906474467 },
{ url = "https://files.pythonhosted.org/packages/40/04/bd91593a4ca178ece93ca55f27e2783aa524aaccbfda66831d59a054c31e/torch-2.5.1-cp311-cp311-manylinux2014_aarch64.whl", hash = "sha256:340ce0432cad0d37f5a31be666896e16788f1adf8ad7be481196b503dad675b9", size = 91919450 },
{ url = "https://files.pythonhosted.org/packages/0d/4a/e51420d46cfc90562e85af2fee912237c662ab31140ab179e49bd69401d6/torch-2.5.1-cp311-cp311-win_amd64.whl", hash = "sha256:603c52d2fe06433c18b747d25f5c333f9c1d58615620578c326d66f258686f9a", size = 203098237 },
{ url = "https://files.pythonhosted.org/packages/d0/db/5d9cbfbc7968d79c5c09a0bc0bc3735da079f2fd07cc10498a62b320a480/torch-2.5.1-cp311-none-macosx_11_0_arm64.whl", hash = "sha256:31f8c39660962f9ae4eeec995e3049b5492eb7360dd4f07377658ef4d728fa4c", size = 63884466 },
{ url = "https://files.pythonhosted.org/packages/8b/5c/36c114d120bfe10f9323ed35061bc5878cc74f3f594003854b0ea298942f/torch-2.5.1-cp312-cp312-manylinux1_x86_64.whl", hash = "sha256:ed231a4b3a5952177fafb661213d690a72caaad97d5824dd4fc17ab9e15cec03", size = 906389343 },
{ url = "https://files.pythonhosted.org/packages/6d/69/d8ada8b6e0a4257556d5b4ddeb4345ea8eeaaef3c98b60d1cca197c7ad8e/torch-2.5.1-cp312-cp312-manylinux2014_aarch64.whl", hash = "sha256:3f4b7f10a247e0dcd7ea97dc2d3bfbfc90302ed36d7f3952b0008d0df264e697", size = 91811673 },
{ url = "https://files.pythonhosted.org/packages/5f/ba/607d013b55b9fd805db2a5c2662ec7551f1910b4eef39653eeaba182c5b2/torch-2.5.1-cp312-cp312-win_amd64.whl", hash = "sha256:73e58e78f7d220917c5dbfad1a40e09df9929d3b95d25e57d9f8558f84c9a11c", size = 203046841 },
{ url = "https://files.pythonhosted.org/packages/57/6c/bf52ff061da33deb9f94f4121fde7ff3058812cb7d2036c97bc167793bd1/torch-2.5.1-cp312-none-macosx_11_0_arm64.whl", hash = "sha256:8c712df61101964eb11910a846514011f0b6f5920c55dbf567bff8a34163d5b1", size = 63858109 },
{ url = "https://files.pythonhosted.org/packages/69/72/20cb30f3b39a9face296491a86adb6ff8f1a47a897e4d14667e6cf89d5c3/torch-2.5.1-cp313-cp313-manylinux1_x86_64.whl", hash = "sha256:9b61edf3b4f6e3b0e0adda8b3960266b9009d02b37555971f4d1c8f7a05afed7", size = 906393265 },
]
[[package]]
name = "triton"
version = "3.1.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "filelock" },
]
wheels = [
{ url = "https://files.pythonhosted.org/packages/86/17/d9a5cf4fcf46291856d1e90762e36cbabd2a56c7265da0d1d9508c8e3943/triton-3.1.0-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:0f34f6e7885d1bf0eaaf7ba875a5f0ce6f3c13ba98f9503651c1e6dc6757ed5c", size = 209506424 },
{ url = "https://files.pythonhosted.org/packages/78/eb/65f5ba83c2a123f6498a3097746607e5b2f16add29e36765305e4ac7fdd8/triton-3.1.0-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:c8182f42fd8080a7d39d666814fa36c5e30cc00ea7eeeb1a2983dbb4c99a0fdc", size = 209551444 },
]
[[package]] [[package]]
name = "typer" name = "typer"
version = "0.15.1" version = "0.15.1"
@ -827,6 +1048,7 @@ dependencies = [
{ name = "pyyaml" }, { name = "pyyaml" },
{ name = "scipy" }, { name = "scipy" },
{ name = "scipy-stubs" }, { name = "scipy-stubs" },
{ name = "torch" },
{ name = "typer" }, { name = "typer" },
] ]
@ -842,6 +1064,7 @@ requires-dist = [
{ name = "pyyaml", specifier = ">=6.0.2" }, { name = "pyyaml", specifier = ">=6.0.2" },
{ name = "scipy" }, { name = "scipy" },
{ name = "scipy-stubs" }, { name = "scipy-stubs" },
{ name = "torch" },
{ name = "typer" }, { name = "typer" },
] ]

View File

@ -3,6 +3,7 @@ import multiprocessing as mp
import numpy as np import numpy as np
import numpy.typing as npt import numpy.typing as npt
import torch
import typer import typer
from .. import visualise from .. import visualise
@ -13,6 +14,7 @@ from . import file, globals
app = typer.Typer(callback=globals.main) app = typer.Typer(callback=globals.main)
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
@app.command() @app.command()
@ -42,11 +44,12 @@ def heatmap() -> None:
webapp = mp.Process(target=visualise.start, args=(webapp_queue,)) webapp = mp.Process(target=visualise.start, args=(webapp_queue,))
webapp.start() webapp.start()
def callback(antenna_data: npt.NDArray[np.complex128]) -> None: def callback(antenna_data: npt.NDArray[np.complex64]) -> None:
logger.info(f"Got final CSI data with shape {antenna_data.shape}") logger.info(f"Got final CSI data with shape {antenna_data.shape}")
processed = preprocessor.preprocess(antenna_data) processed = preprocessor.preprocess(antenna_data)
processed_tensor = torch.tensor(processed, device=device)
# visualise.add_data(all_data, processed) # visualise.add_data(all_data, processed)
aoa.update(processed) aoa.update(processed_tensor)
if not webapp_queue.full(): if not webapp_queue.full():
webapp_queue.put(aoa) webapp_queue.put(aoa)

View File

@ -93,14 +93,14 @@ class CSIHeader:
class CSI: class CSI:
@staticmethod @staticmethod
def parseCsiData(data: bytes, header: CSIHeader) -> npt.NDArray[np.complex128]: def parseCsiData(data: bytes, header: CSIHeader) -> npt.NDArray[np.complex64]:
csi_matrix: npt.NDArray[np.complex128] = np.zeros( csi_matrix: npt.NDArray[np.complex64] = np.zeros(
( (
header.num_subcarriers, header.num_subcarriers,
header.num_rx, header.num_rx,
header.num_tx, header.num_tx,
), ),
dtype=np.complex128, dtype=np.complex64,
) )
pos = 0 pos = 0
for j in range(header.num_rx): for j in range(header.num_rx):

View File

@ -3,7 +3,7 @@ from typing import Callable, Protocol
import numpy as np import numpy as np
import numpy.typing as npt import numpy.typing as npt
CSICallback = Callable[[npt.NDArray[np.complex128]], None] CSICallback = Callable[[npt.NDArray[np.complex64]], None]
class CSIProducer(Protocol): class CSIProducer(Protocol):

View File

@ -3,21 +3,26 @@ from datetime import datetime
import numpy as np import numpy as np
import numpy.typing as npt import numpy.typing as npt
import torch
import torch.linalg
from ..config import config from ..config import config
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.set_default_device(device)
class AoA: class AoA:
def __init__(self) -> None: def __init__(self) -> None:
self.historical_autocorr = np.array([]) self.historical_autocorr = torch.tensor([], dtype=torch.complex64)
self.N_subcarriers = -1 self.N_subcarriers = -1
self.N_rx = -1 self.N_rx = -1
self.timestamp = datetime.now() self.timestamp = datetime.now()
pass pass
def smooth(self, data: npt.NDArray[np.complex128]) -> npt.NDArray[np.complex128]: def smooth(self, data: torch.Tensor) -> torch.Tensor:
assert len(data.shape) == 3 assert len(data.shape) == 3
M = data.shape[0] # Number of subcarriers M = data.shape[0] # Number of subcarriers
@ -32,32 +37,32 @@ class AoA:
# This only works with 1 TX antenna (i.e. no MIMO) - see #4 for more details # This only works with 1 TX antenna (i.e. no MIMO) - see #4 for more details
assert T == 1, "The current implementation only supports 1 TX antenna" assert T == 1, "The current implementation only supports 1 TX antenna"
H_n = np.zeros((N, M // 2, M // 2 + 1), dtype=np.complex128) H_n = torch.zeros((N, M // 2, M // 2 + 1), dtype=torch.complex64)
for i in range(N): for i in range(N):
for j in range(M // 2): for j in range(M // 2):
H_n[i, j] = data[j : j + M // 2 + 1, i, 0] H_n[i, j] = data[j : j + M // 2 + 1, i, 0]
H_sm_rows = [np.hstack(H_n[i : i + N // 2 + 1]) for i in range(N // 2)] H_sm_rows = [torch.hstack(list(H_n[i : i + N // 2 + 1])) for i in range(N // 2)]
H_sm = np.vstack(H_sm_rows) H_sm = torch.vstack(H_sm_rows)
logger.debug(f"Smoothed: {H_sm.shape}") logger.debug(f"Smoothed: {H_sm.shape}")
return H_sm return H_sm
def update(self, data: npt.NDArray[np.complex128]) -> None: def update(self, data: torch.Tensor) -> None:
self.timestamp = datetime.now() self.timestamp = datetime.now()
H_sm = self.smooth(data) H_sm = self.smooth(data)
auto_corr = np.matmul(H_sm, np.conj(H_sm).T) auto_corr = H_sm @ torch.conj(H_sm).T
# This matrix is by definition Hermitian. # This matrix is by definition Hermitian.
# Therefore, all of its eigenvectors are orthogonal. # Therefore, all of its eigenvectors are orthogonal.
if self.historical_autocorr.size == 0: if len(self.historical_autocorr.shape) <= 1:
self.historical_autocorr = np.expand_dims(auto_corr, 0) self.historical_autocorr = torch.unsqueeze(auto_corr, 0)
else: else:
self.historical_autocorr = np.append( self.historical_autocorr = torch.cat(
self.historical_autocorr, np.expand_dims(auto_corr, 0), axis=0 (self.historical_autocorr, torch.unsqueeze(auto_corr, 0))
) )
WINDOW_SIZE = config.music.window_size WINDOW_SIZE = config.music.window_size
@ -65,7 +70,7 @@ class AoA:
self.historical_autocorr = self.historical_autocorr[-WINDOW_SIZE:] self.historical_autocorr = self.historical_autocorr[-WINDOW_SIZE:]
# Is the moving average also Hermitian? # Is the moving average also Hermitian?
R = np.mean(self.historical_autocorr, axis=0) R = torch.mean(self.historical_autocorr, dim=0)
# The smallest eigenvectors span the noise subspace, # The smallest eigenvectors span the noise subspace,
# and the largest span the signal subspace. # and the largest span the signal subspace.
@ -75,8 +80,8 @@ class AoA:
def steering_vector( def steering_vector(
self, theta: float, tof: float self, theta: float, tof: float
) -> npt.NDArray[np.complexfloating]: ) -> npt.NDArray[np.complexfloating]:
omega_t: npt.NDArray[np.complex128] = np.exp(-2j * np.pi * config.delta_f * tof) omega_t: npt.NDArray[np.complex64] = np.exp(-2j * np.pi * config.delta_f * tof)
phi_theta: npt.NDArray[np.complex128] = np.exp( phi_theta: npt.NDArray[np.complex64] = np.exp(
2j 2j
* np.pi * np.pi
* config.central_freq_hz * config.central_freq_hz
@ -88,10 +93,12 @@ class AoA:
omega_t = np.expand_dims(omega_t, axis=-1) omega_t = np.expand_dims(omega_t, axis=-1)
phi_theta = np.expand_dims(phi_theta, axis=-1) phi_theta = np.expand_dims(phi_theta, axis=-1)
antenna_v = omega_t ** np.arange(self.N_subcarriers // 2) antenna_v = omega_t ** torch.arange(self.N_subcarriers // 2)
phis = phi_theta ** np.arange(self.N_rx // 2) phis = phi_theta ** torch.arange(self.N_rx // 2)
antenna_v = np.expand_dims(antenna_v, axis=-1) antenna_v = np.expand_dims(antenna_v, axis=-1)
steering = antenna_v * phis print(antenna_v.shape, phis.shape)
steering = antenna_v[0] * phis
print(steering.shape)
return steering.T.reshape(-1) return steering.T.reshape(-1)
def evaluate(self, theta: float, tof: float) -> float: def evaluate(self, theta: float, tof: float) -> float:
@ -101,23 +108,25 @@ class AoA:
except Exception as e: except Exception as e:
logger.exception(e) logger.exception(e)
return 0 return 0
assert isinstance(self.E_n, torch.Tensor)
E_n = self.E_n E_n = self.E_n
E_n_H = np.conj(E_n).T E_n_H = np.conj(E_n).T
c = 1 / (0.001 + (steering_h @ E_n @ E_n_H @ steering)) c = 1 / (0.001 + (steering_h @ E_n @ E_n_H @ steering))
return np.abs(c.real) return np.abs(c.real).item()
def test_smoothing() -> None: def test_smoothing() -> None:
row, col = np.indices((4, 2)) row, col = np.indices((6, 4))
data = row + 1j * col data = row + 1j * col
data = np.expand_dims(data, axis=2)
np.set_printoptions(linewidth=200)
print(data.shape)
aoa = AoA() aoa = AoA()
aoa.N_subcarriers = 6
aoa.N_rx = 4
smoothed = aoa.smooth(data) smoothed = aoa.smooth(data)
H_0 = np.array([[0 + 0j, 0 + 1j, 0 + 2j], [0 + 1j, 0 + 2j, 0 + 3j]]) print(smoothed)
H_01 = np.vstack([H_0, H_0 + 1])
H_12 = np.vstack([H_0 + 1, H_0 + 2])
expected = np.hstack([H_01, H_12])
print(expected)
assert np.allclose(smoothed, expected)
def test_steering_vector() -> None: def test_steering_vector() -> None:

View File

@ -15,9 +15,9 @@ np.seterr(invalid="ignore")
class Preprocessor: class Preprocessor:
def __init__(self) -> None: def __init__(self) -> None:
self.prev_entries: Queue[npt.NDArray[np.complex128]] = Queue(maxsize=100) self.prev_entries: Queue[npt.NDArray[np.complex64]] = Queue(maxsize=100)
self.short_term_avg = np.zeros((1,), dtype=np.complex128) self.short_term_avg = np.zeros((1,), dtype=np.complex64)
self.long_term_avg = np.zeros((1,), dtype=np.complex128) self.long_term_avg = np.zeros((1,), dtype=np.complex64)
self.filter = butter( self.filter = butter(
5, 5,
config.preprocessing.bandpass.bounds, config.preprocessing.bandpass.bounds,
@ -26,7 +26,7 @@ class Preprocessor:
output="sos", output="sos",
) )
def preprocess(self, h: npt.NDArray[np.complex128]) -> npt.NDArray[np.complex128]: def preprocess(self, h: npt.NDArray[np.complex64]) -> npt.NDArray[np.complex64]:
# CSI data is not available for pilot subcarriers. # CSI data is not available for pilot subcarriers.
h_hat = np.where( h_hat = np.where(
np.expand_dims(h[:, 0, 0] == 0, axis=(1, 2)), np.expand_dims(h[:, 0, 0] == 0, axis=(1, 2)),
@ -46,7 +46,7 @@ class Preprocessor:
# Assume that all csi matrices will have the same shape # Assume that all csi matrices will have the same shape
if self.long_term_avg.shape != h_hat.shape: if self.long_term_avg.shape != h_hat.shape:
self.long_term_avg = np.zeros(h_hat.shape, dtype=np.complex128) self.long_term_avg = np.zeros(h_hat.shape, dtype=np.complex64)
self.long_term_avg = ( self.long_term_avg = (
self.long_term_avg * (1 - config.preprocessing.moving_average_alpha) self.long_term_avg * (1 - config.preprocessing.moving_average_alpha)