set up torch AoA estimation
This commit is contained in:
parent
66e041cf73
commit
a2f94737f8
1
Pipfile
1
Pipfile
@ -11,6 +11,7 @@ pytest = "*"
|
||||
matplotlib = "*"
|
||||
scipy = "*"
|
||||
scipy-stubs = "*"
|
||||
torch = "*"
|
||||
|
||||
[dev-packages]
|
||||
|
||||
|
||||
192
Pipfile.lock
generated
192
Pipfile.lock
generated
@ -1,7 +1,7 @@
|
||||
{
|
||||
"_meta": {
|
||||
"hash": {
|
||||
"sha256": "c6ec3616be9da08134fd0a1eb5319ad8a9020abd1a929dd1196bbf5bf077f2e5"
|
||||
"sha256": "c6fdcd549f16bcc17f8af735f5682046bed934fc24446e2356cf932c853ca21f"
|
||||
},
|
||||
"pipfile-spec": 6,
|
||||
"requires": {
|
||||
@ -100,6 +100,14 @@
|
||||
"markers": "python_version >= '3.8'",
|
||||
"version": "==0.12.1"
|
||||
},
|
||||
"filelock": {
|
||||
"hashes": [
|
||||
"sha256:2082e5703d51fbf98ea75855d9d5527e33d8ff23099bec374a134febee6946b0",
|
||||
"sha256:c249fbfcd5db47e5e2d6d62198e565475ee65e4831e2561c8e313fa7eb961435"
|
||||
],
|
||||
"markers": "python_version >= '3.8'",
|
||||
"version": "==3.16.1"
|
||||
},
|
||||
"flask": {
|
||||
"hashes": [
|
||||
"sha256:5f873c5184c897c8d9d1b05df1e3d01b14910ce69607a117bd3277098a5836ac",
|
||||
@ -174,6 +182,14 @@
|
||||
"markers": "python_version >= '3.8'",
|
||||
"version": "==4.55.3"
|
||||
},
|
||||
"fsspec": {
|
||||
"hashes": [
|
||||
"sha256:670700c977ed2fb51e0d9f9253177ed20cbde4a3e5c0283cc5385b5870c8533f",
|
||||
"sha256:b520aed47ad9804237ff878b504267a3b0b441e97508bd6d2d8774e3db85cee2"
|
||||
],
|
||||
"markers": "python_version >= '3.8'",
|
||||
"version": "==2024.12.0"
|
||||
},
|
||||
"h11": {
|
||||
"hashes": [
|
||||
"sha256:8f19fbbe99e72420ff35c00b27a34cb9937e902a8b810e2c88300c6f0a3b699d",
|
||||
@ -400,6 +416,21 @@
|
||||
"markers": "python_version >= '3.10'",
|
||||
"version": "==3.10.0"
|
||||
},
|
||||
"mpmath": {
|
||||
"hashes": [
|
||||
"sha256:7a28eb2a9774d00c7bc92411c19a89209d5da7c4c9a9e227be8330a23a25b91f",
|
||||
"sha256:a0b2b9fe80bbcd81a6647ff13108738cfb482d481d826cc0e02f5b35e5c88d2c"
|
||||
],
|
||||
"version": "==1.3.0"
|
||||
},
|
||||
"networkx": {
|
||||
"hashes": [
|
||||
"sha256:307c3669428c5362aab27c8a1260aa8f47c4e91d3891f48be0141738d8d053e1",
|
||||
"sha256:df5d4365b724cf81b8c6a7312509d0c22386097011ad1abe274afd5e9d3bbc5f"
|
||||
],
|
||||
"markers": "python_version >= '3.10'",
|
||||
"version": "==3.4.2"
|
||||
},
|
||||
"numpy": {
|
||||
"hashes": [
|
||||
"sha256:059e6a747ae84fce488c3ee397cee7e5f905fd1bda5fb18c66bc41807ff119b2",
|
||||
@ -462,13 +493,118 @@
|
||||
"markers": "python_version >= '3.10'",
|
||||
"version": "==2.2.1"
|
||||
},
|
||||
"nvidia-cublas-cu12": {
|
||||
"hashes": [
|
||||
"sha256:0f8aa1706812e00b9f19dfe0cdb3999b092ccb8ca168c0db5b8ea712456fd9b3",
|
||||
"sha256:2fc8da60df463fdefa81e323eef2e36489e1c94335b5358bcb38360adf75ac9b",
|
||||
"sha256:5a796786da89203a0657eda402bcdcec6180254a8ac22d72213abc42069522dc"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==12.4.5.8"
|
||||
},
|
||||
"nvidia-cuda-cupti-cu12": {
|
||||
"hashes": [
|
||||
"sha256:5688d203301ab051449a2b1cb6690fbe90d2b372f411521c86018b950f3d7922",
|
||||
"sha256:79279b35cf6f91da114182a5ce1864997fd52294a87a16179ce275773799458a",
|
||||
"sha256:9dec60f5ac126f7bb551c055072b69d85392b13311fcc1bcda2202d172df30fb"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==12.4.127"
|
||||
},
|
||||
"nvidia-cuda-nvrtc-cu12": {
|
||||
"hashes": [
|
||||
"sha256:0eedf14185e04b76aa05b1fea04133e59f465b6f960c0cbf4e37c3cb6b0ea198",
|
||||
"sha256:a178759ebb095827bd30ef56598ec182b85547f1508941a3d560eb7ea1fbf338",
|
||||
"sha256:a961b2f1d5f17b14867c619ceb99ef6fcec12e46612711bcec78eb05068a60ec"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==12.4.127"
|
||||
},
|
||||
"nvidia-cuda-runtime-cu12": {
|
||||
"hashes": [
|
||||
"sha256:09c2e35f48359752dfa822c09918211844a3d93c100a715d79b59591130c5e1e",
|
||||
"sha256:64403288fa2136ee8e467cdc9c9427e0434110899d07c779f25b5c068934faa5",
|
||||
"sha256:961fe0e2e716a2a1d967aab7caee97512f71767f852f67432d572e36cb3a11f3"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==12.4.127"
|
||||
},
|
||||
"nvidia-cudnn-cu12": {
|
||||
"hashes": [
|
||||
"sha256:165764f44ef8c61fcdfdfdbe769d687e06374059fbb388b6c89ecb0e28793a6f",
|
||||
"sha256:6278562929433d68365a07a4a1546c237ba2849852c0d4b2262a486e805b977a"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==9.1.0.70"
|
||||
},
|
||||
"nvidia-cufft-cu12": {
|
||||
"hashes": [
|
||||
"sha256:5dad8008fc7f92f5ddfa2101430917ce2ffacd86824914c82e28990ad7f00399",
|
||||
"sha256:d802f4954291101186078ccbe22fc285a902136f974d369540fd4a5333d1440b",
|
||||
"sha256:f083fc24912aa410be21fa16d157fed2055dab1cc4b6934a0e03cba69eb242b9"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==11.2.1.3"
|
||||
},
|
||||
"nvidia-curand-cu12": {
|
||||
"hashes": [
|
||||
"sha256:1f173f09e3e3c76ab084aba0de819c49e56614feae5c12f69883f4ae9bb5fad9",
|
||||
"sha256:a88f583d4e0bb643c49743469964103aa59f7f708d862c3ddb0fc07f851e3b8b",
|
||||
"sha256:f307cc191f96efe9e8f05a87096abc20d08845a841889ef78cb06924437f6771"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==10.3.5.147"
|
||||
},
|
||||
"nvidia-cusolver-cu12": {
|
||||
"hashes": [
|
||||
"sha256:19e33fa442bcfd085b3086c4ebf7e8debc07cfe01e11513cc6d332fd918ac260",
|
||||
"sha256:d338f155f174f90724bbde3758b7ac375a70ce8e706d70b018dd3375545fc84e",
|
||||
"sha256:e77314c9d7b694fcebc84f58989f3aa4fb4cb442f12ca1a9bde50f5e8f6d1b9c"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==11.6.1.9"
|
||||
},
|
||||
"nvidia-cusparse-cu12": {
|
||||
"hashes": [
|
||||
"sha256:9bc90fb087bc7b4c15641521f31c0371e9a612fc2ba12c338d3ae032e6b6797f",
|
||||
"sha256:9d32f62896231ebe0480efd8a7f702e143c98cfaa0e8a76df3386c1ba2b54df3",
|
||||
"sha256:ea4f11a2904e2a8dc4b1833cc1b5181cde564edd0d5cd33e3c168eff2d1863f1"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==12.3.1.170"
|
||||
},
|
||||
"nvidia-nccl-cu12": {
|
||||
"hashes": [
|
||||
"sha256:8579076d30a8c24988834445f8d633c697d42397e92ffc3f63fa26766d25e0a0"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==2.21.5"
|
||||
},
|
||||
"nvidia-nvjitlink-cu12": {
|
||||
"hashes": [
|
||||
"sha256:06b3b9b25bf3f8af351d664978ca26a16d2c5127dbd53c0497e28d1fb9611d57",
|
||||
"sha256:4abe7fef64914ccfa909bc2ba39739670ecc9e820c83ccc7a6ed414122599b83",
|
||||
"sha256:fd9020c501d27d135f983c6d3e244b197a7ccad769e34df53a42e276b0e25fa1"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==12.4.127"
|
||||
},
|
||||
"nvidia-nvtx-cu12": {
|
||||
"hashes": [
|
||||
"sha256:641dccaaa1139f3ffb0d3164b4b84f9d253397e38246a4f2f36728b48566d485",
|
||||
"sha256:781e950d9b9f60d8241ccea575b32f5105a5baf4c2351cab5256a24869f12a1a",
|
||||
"sha256:7959ad635db13edf4fc65c06a6e9f9e55fc2f92596db928d169c0bb031e88ef3"
|
||||
],
|
||||
"markers": "python_version >= '3'",
|
||||
"version": "==12.4.127"
|
||||
},
|
||||
"optype": {
|
||||
"hashes": [
|
||||
"sha256:51c8dd104ac197457059bcee5e256160a641ca72c6c852012d952790a5e0cac0",
|
||||
"sha256:856416484131038799e0e9cefc19d0ef37e7b4fde2144f25e8bb0e8981ebbe95"
|
||||
"sha256:8cbfd452d6f06c7c70502048f38a0d5451bc601054d3a577dd09c7d6363950e1",
|
||||
"sha256:90a7760177f2e7feae379a60445fceec37b932b75a00c3d96067497573c5e84d"
|
||||
],
|
||||
"markers": "python_version >= '3.10'",
|
||||
"version": "==0.7.3"
|
||||
"version": "==0.8.0"
|
||||
},
|
||||
"packaging": {
|
||||
"hashes": [
|
||||
@ -641,6 +777,14 @@
|
||||
"markers": "python_version >= '3.10'",
|
||||
"version": "==1.14.1.6"
|
||||
},
|
||||
"setuptools": {
|
||||
"hashes": [
|
||||
"sha256:8199222558df7c86216af4f84c30e9b34a61d8ba19366cc914424cdbd28252f6",
|
||||
"sha256:ce74b49e8f7110f9bf04883b730f4765b774ef3ef28f722cce7c273d253aaf7d"
|
||||
],
|
||||
"markers": "python_version >= '3.9'",
|
||||
"version": "==75.6.0"
|
||||
},
|
||||
"simple-websocket": {
|
||||
"hashes": [
|
||||
"sha256:4af6069630a38ed6c561010f0e11a5bc0d4ca569b36306eb257cd9a192497c8c",
|
||||
@ -657,6 +801,46 @@
|
||||
"markers": "python_version >= '2.7' and python_version not in '3.0, 3.1, 3.2'",
|
||||
"version": "==1.17.0"
|
||||
},
|
||||
"sympy": {
|
||||
"hashes": [
|
||||
"sha256:9cebf7e04ff162015ce31c9c6c9144daa34a93bd082f54fd8f12deca4f47515f",
|
||||
"sha256:db36cdc64bf61b9b24578b6f7bab1ecdd2452cf008f34faa33776680c26d66f8"
|
||||
],
|
||||
"markers": "python_version >= '3.8'",
|
||||
"version": "==1.13.1"
|
||||
},
|
||||
"torch": {
|
||||
"hashes": [
|
||||
"sha256:1f3b7fb3cf7ab97fae52161423f81be8c6b8afac8d9760823fd623994581e1a3",
|
||||
"sha256:23d062bf70776a3d04dbe74db950db2a5245e1ba4f27208a87f0d743b0d06e86",
|
||||
"sha256:31f8c39660962f9ae4eeec995e3049b5492eb7360dd4f07377658ef4d728fa4c",
|
||||
"sha256:32a037bd98a241df6c93e4c789b683335da76a2ac142c0973675b715102dc5fa",
|
||||
"sha256:340ce0432cad0d37f5a31be666896e16788f1adf8ad7be481196b503dad675b9",
|
||||
"sha256:34bfa1a852e5714cbfa17f27c49d8ce35e1b7af5608c4bc6e81392c352dbc601",
|
||||
"sha256:3f4b7f10a247e0dcd7ea97dc2d3bfbfc90302ed36d7f3952b0008d0df264e697",
|
||||
"sha256:46c817d3ea33696ad3b9df5e774dba2257e9a4cd3c4a3afbf92f6bb13ac5ce2d",
|
||||
"sha256:603c52d2fe06433c18b747d25f5c333f9c1d58615620578c326d66f258686f9a",
|
||||
"sha256:71328e1bbe39d213b8721678f9dcac30dfc452a46d586f1d514a6aa0a99d4744",
|
||||
"sha256:73e58e78f7d220917c5dbfad1a40e09df9929d3b95d25e57d9f8558f84c9a11c",
|
||||
"sha256:7974e3dce28b5a21fb554b73e1bc9072c25dde873fa00d54280861e7a009d7dc",
|
||||
"sha256:8046768b7f6d35b85d101b4b38cba8aa2f3cd51952bc4c06a49580f2ce682291",
|
||||
"sha256:8c712df61101964eb11910a846514011f0b6f5920c55dbf567bff8a34163d5b1",
|
||||
"sha256:9b61edf3b4f6e3b0e0adda8b3960266b9009d02b37555971f4d1c8f7a05afed7",
|
||||
"sha256:de5b7d6740c4b636ef4db92be922f0edc425b65ed78c5076c43c42d362a45457",
|
||||
"sha256:ed231a4b3a5952177fafb661213d690a72caaad97d5824dd4fc17ab9e15cec03"
|
||||
],
|
||||
"index": "pypi",
|
||||
"markers": "python_full_version >= '3.8.0'",
|
||||
"version": "==2.5.1"
|
||||
},
|
||||
"typing-extensions": {
|
||||
"hashes": [
|
||||
"sha256:04e5ca0351e0f3f85c6853954072df659d0d13fac324d0072316b67d7794700d",
|
||||
"sha256:1a7ead55c7e559dd4dee8856e3a88b41225abfe1ce8df57b7c13915fe121ffb8"
|
||||
],
|
||||
"markers": "python_version >= '3.8'",
|
||||
"version": "==4.12.2"
|
||||
},
|
||||
"werkzeug": {
|
||||
"hashes": [
|
||||
"sha256:54b78bf3716d19a65be4fceccc0d1d7b89e608834989dfae50ea87564639213e",
|
||||
|
||||
110
src/aoa.py
110
src/aoa.py
@ -1,5 +1,6 @@
|
||||
import numpy as np
|
||||
import numpy.typing as npt
|
||||
import torch
|
||||
import torch.linalg
|
||||
from . import config
|
||||
from datetime import datetime
|
||||
|
||||
@ -7,16 +8,19 @@ import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
torch.set_default_device(device)
|
||||
|
||||
|
||||
class AoA:
|
||||
def __init__(self):
|
||||
self.historical_autocorr = np.array([])
|
||||
self.historical_autocorr = torch.tensor([], dtype=torch.complex64)
|
||||
self.N_subcarriers = -1
|
||||
self.N_rx = -1
|
||||
self.timestamp = datetime.now()
|
||||
pass
|
||||
|
||||
def smooth(self, data: npt.NDArray[np.complex128]):
|
||||
def smooth(self, data: torch.Tensor):
|
||||
assert len(data.shape) == 3
|
||||
|
||||
M = data.shape[0] # Number of subcarriers
|
||||
@ -31,32 +35,32 @@ class AoA:
|
||||
# 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"
|
||||
|
||||
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 j in range(M // 2):
|
||||
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 = np.vstack(H_sm_rows)
|
||||
H_sm_rows = [torch.hstack(list(H_n[i : i + N // 2 + 1])) for i in range(N // 2)]
|
||||
H_sm = torch.vstack(H_sm_rows)
|
||||
|
||||
logger.debug(f"Smoothed: {H_sm.shape}")
|
||||
|
||||
return H_sm
|
||||
|
||||
def update(self, data: npt.NDArray[np.complex128]):
|
||||
def update(self, data: torch.Tensor):
|
||||
self.timestamp = datetime.now()
|
||||
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.
|
||||
# Therefore, all of its eigenvectors are orthogonal.
|
||||
|
||||
if self.historical_autocorr.size == 0:
|
||||
self.historical_autocorr = np.expand_dims(auto_corr, 0)
|
||||
if len(self.historical_autocorr.shape) <= 1:
|
||||
self.historical_autocorr = torch.unsqueeze(auto_corr, 0)
|
||||
else:
|
||||
self.historical_autocorr = np.append(
|
||||
self.historical_autocorr, np.expand_dims(auto_corr, 0), axis=0
|
||||
self.historical_autocorr = torch.cat(
|
||||
(self.historical_autocorr, torch.unsqueeze(auto_corr, 0))
|
||||
)
|
||||
|
||||
WINDOW_SIZE = config.AOA_SLIDING_WINDOW_SIZE
|
||||
@ -64,72 +68,78 @@ class AoA:
|
||||
self.historical_autocorr = self.historical_autocorr[-WINDOW_SIZE:]
|
||||
|
||||
# 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,
|
||||
# and the largest span the signal subspace.
|
||||
eigvals, eigvecs = np.linalg.eigh(R)
|
||||
self.E_n = eigvecs[:, np.abs(eigvals) < config.EIGVAL_THRESHOLD]
|
||||
|
||||
omega_base = np.exp(-2j * np.pi * config.DELTA_F)
|
||||
phi_base = np.exp(
|
||||
2j * np.pi * config.CENTRAL_FREQUENCY_HZ * config.ANTENNA_SPACING / config.C
|
||||
)
|
||||
eigvals, eigvecs = torch.linalg.eigh(R)
|
||||
logging.debug(f"Eigenvalues: {eigvals}")
|
||||
self.E_n = eigvecs[:, torch.abs(eigvals) < config.EIGVAL_THRESHOLD]
|
||||
|
||||
def steering_vector(self, theta: float, tof: float):
|
||||
omega_t = np.exp(-2j * np.pi * config.DELTA_F * tof)
|
||||
phi_theta = np.exp(
|
||||
2j
|
||||
* np.pi
|
||||
* config.CENTRAL_FREQUENCY_HZ
|
||||
* config.ANTENNA_SPACING
|
||||
* (1 - np.cos(theta))
|
||||
/ config.C
|
||||
omega_t = torch.exp(
|
||||
torch.tensor([-2j * torch.pi * config.DELTA_F * tof], dtype=torch.complex64)
|
||||
)
|
||||
phi_theta = torch.exp(
|
||||
torch.tensor(
|
||||
[
|
||||
2j
|
||||
* np.pi
|
||||
* config.CENTRAL_FREQUENCY_HZ
|
||||
* config.ANTENNA_SPACING
|
||||
* (1 - np.cos(theta))
|
||||
/ config.C
|
||||
],
|
||||
dtype=torch.complex64,
|
||||
)
|
||||
)
|
||||
|
||||
omega_t = np.expand_dims(omega_t, axis=-1)
|
||||
phi_theta = np.expand_dims(phi_theta, axis=-1)
|
||||
omega_t = torch.unsqueeze(omega_t, dim=-1)
|
||||
phi_theta = torch.unsqueeze(phi_theta, dim=-1)
|
||||
|
||||
antenna_v = omega_t ** np.arange(self.N_subcarriers // 2)
|
||||
phis = phi_theta ** np.arange(self.N_rx // 2)
|
||||
antenna_v = np.expand_dims(antenna_v, axis=-1)
|
||||
steering = antenna_v * phis
|
||||
antenna_v = omega_t ** torch.arange(self.N_subcarriers // 2)
|
||||
phis = phi_theta ** torch.arange(self.N_rx // 2)
|
||||
antenna_v = torch.unsqueeze(antenna_v, dim=-1)
|
||||
print(antenna_v.shape, phis.shape)
|
||||
steering = antenna_v[0] * phis
|
||||
print(steering.shape)
|
||||
return steering.T.reshape(-1)
|
||||
|
||||
def evaluate(self, theta: float, tof: float):
|
||||
try:
|
||||
steering = self.steering_vector(theta, tof)
|
||||
steering_h = np.conj(steering).T
|
||||
steering_h = torch.conj(steering).T
|
||||
except Exception as e:
|
||||
logger.exception(e)
|
||||
return 0
|
||||
|
||||
assert isinstance(self.E_n, torch.Tensor)
|
||||
E_n = self.E_n
|
||||
E_n_H = np.conj(E_n).T
|
||||
E_n_H = torch.conj(E_n).T
|
||||
c = 1 / (0.001 + (steering_h @ E_n @ E_n_H @ steering))
|
||||
return np.abs(c.real)
|
||||
return torch.abs(c.real)
|
||||
|
||||
|
||||
def test_smoothing():
|
||||
row, col = np.indices((4, 2))
|
||||
row, col = torch.indices((6, 4))
|
||||
data = row + 1j * col
|
||||
data = np.expand_dims(data, axis=2)
|
||||
np.set_printoptions(linewidth=200)
|
||||
print(data.shape)
|
||||
aoa = AoA()
|
||||
aoa.N_subcarriers = 6
|
||||
aoa.N_rx = 4
|
||||
smoothed = aoa.smooth(data)
|
||||
H_0 = np.array([[0 + 0j, 0 + 1j, 0 + 2j], [0 + 1j, 0 + 2j, 0 + 3j]])
|
||||
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)
|
||||
print(smoothed)
|
||||
|
||||
pass
|
||||
|
||||
|
||||
def test_steering_vector():
|
||||
aoa = AoA()
|
||||
aoa.N_subcarriers = 10
|
||||
aoa.N_rx = 2
|
||||
print(aoa.omega_base)
|
||||
print(aoa.phi_base)
|
||||
tau = 1
|
||||
theta = 0
|
||||
aoa.N_subcarriers = 6
|
||||
aoa.N_rx = 4
|
||||
tau = 1e-8
|
||||
theta = np.pi / 2
|
||||
print(aoa.steering_vector(theta, tau))
|
||||
assert False
|
||||
|
||||
@ -94,13 +94,13 @@ class CSIHeader:
|
||||
class CSI:
|
||||
@staticmethod
|
||||
def parseCsiData(data: bytes, header: CSIHeader):
|
||||
csi_matrix: npt.NDArray[np.complex128] = np.zeros(
|
||||
csi_matrix: npt.NDArray[np.complex64] = np.zeros(
|
||||
(
|
||||
header.num_subcarriers,
|
||||
header.num_rx,
|
||||
header.num_tx,
|
||||
),
|
||||
dtype=np.complex128,
|
||||
dtype=np.complex64,
|
||||
)
|
||||
pos = 0
|
||||
for j in range(header.num_rx):
|
||||
|
||||
@ -3,6 +3,7 @@ import time
|
||||
import socket
|
||||
import multiprocessing as mp
|
||||
import struct
|
||||
import torch
|
||||
from datetime import datetime
|
||||
import numpy as np
|
||||
|
||||
@ -12,6 +13,7 @@ from .preprocess import Preprocessor
|
||||
from . import config
|
||||
|
||||
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
Host = tuple[str, int]
|
||||
|
||||
|
||||
@ -123,8 +125,10 @@ class CSIProcessor:
|
||||
self.logger.info(f"Got final CSI data with shape {all_data.shape}")
|
||||
|
||||
processed = self.preprocess.preprocess(all_data)
|
||||
visualise.add_data(all_data, processed)
|
||||
self.aoa.update(processed)
|
||||
processed_tensor = torch.tensor(processed, device=device)
|
||||
# visualise.add_data(all_data_tensor, processed)
|
||||
self.aoa.update(processed_tensor)
|
||||
self.logger.info("Processed data")
|
||||
if not self.webserver.full():
|
||||
self.webserver.put(self.aoa)
|
||||
|
||||
|
||||
@ -15,9 +15,9 @@ np.seterr(invalid="ignore")
|
||||
|
||||
class Preprocessor:
|
||||
def __init__(self):
|
||||
self.prev_entries: Queue[npt.NDArray[np.complex128]] = Queue(maxsize=100)
|
||||
self.short_term_avg = np.zeros((1,), dtype=np.complex128)
|
||||
self.long_term_avg = np.zeros((1,), dtype=np.complex128)
|
||||
self.prev_entries: Queue[npt.NDArray[np.complex64]] = Queue(maxsize=100)
|
||||
self.short_term_avg = np.zeros((1,), dtype=np.complex64)
|
||||
self.long_term_avg = np.zeros((1,), dtype=np.complex64)
|
||||
self.filter = butter(
|
||||
5,
|
||||
[
|
||||
@ -29,7 +29,7 @@ class Preprocessor:
|
||||
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.
|
||||
h_hat = np.where(
|
||||
np.expand_dims(h[:, 0, 0] == 0, axis=(1, 2)),
|
||||
@ -49,7 +49,7 @@ class Preprocessor:
|
||||
|
||||
# Assume that all csi matrices will have the same 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 * (1 - config.PREPROCESSING_LONG_TERM_ALPHA)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user