set up torch AoA estimation

This commit is contained in:
Christos Falas 2024-12-31 16:45:42 +00:00
parent 66e041cf73
commit a2f94737f8
No known key found for this signature in database
6 changed files with 262 additions and 63 deletions

View File

@ -11,6 +11,7 @@ pytest = "*"
matplotlib = "*" matplotlib = "*"
scipy = "*" scipy = "*"
scipy-stubs = "*" scipy-stubs = "*"
torch = "*"
[dev-packages] [dev-packages]

192
Pipfile.lock generated
View File

@ -1,7 +1,7 @@
{ {
"_meta": { "_meta": {
"hash": { "hash": {
"sha256": "c6ec3616be9da08134fd0a1eb5319ad8a9020abd1a929dd1196bbf5bf077f2e5" "sha256": "c6fdcd549f16bcc17f8af735f5682046bed934fc24446e2356cf932c853ca21f"
}, },
"pipfile-spec": 6, "pipfile-spec": 6,
"requires": { "requires": {
@ -100,6 +100,14 @@
"markers": "python_version >= '3.8'", "markers": "python_version >= '3.8'",
"version": "==0.12.1" "version": "==0.12.1"
}, },
"filelock": {
"hashes": [
"sha256:2082e5703d51fbf98ea75855d9d5527e33d8ff23099bec374a134febee6946b0",
"sha256:c249fbfcd5db47e5e2d6d62198e565475ee65e4831e2561c8e313fa7eb961435"
],
"markers": "python_version >= '3.8'",
"version": "==3.16.1"
},
"flask": { "flask": {
"hashes": [ "hashes": [
"sha256:5f873c5184c897c8d9d1b05df1e3d01b14910ce69607a117bd3277098a5836ac", "sha256:5f873c5184c897c8d9d1b05df1e3d01b14910ce69607a117bd3277098a5836ac",
@ -174,6 +182,14 @@
"markers": "python_version >= '3.8'", "markers": "python_version >= '3.8'",
"version": "==4.55.3" "version": "==4.55.3"
}, },
"fsspec": {
"hashes": [
"sha256:670700c977ed2fb51e0d9f9253177ed20cbde4a3e5c0283cc5385b5870c8533f",
"sha256:b520aed47ad9804237ff878b504267a3b0b441e97508bd6d2d8774e3db85cee2"
],
"markers": "python_version >= '3.8'",
"version": "==2024.12.0"
},
"h11": { "h11": {
"hashes": [ "hashes": [
"sha256:8f19fbbe99e72420ff35c00b27a34cb9937e902a8b810e2c88300c6f0a3b699d", "sha256:8f19fbbe99e72420ff35c00b27a34cb9937e902a8b810e2c88300c6f0a3b699d",
@ -400,6 +416,21 @@
"markers": "python_version >= '3.10'", "markers": "python_version >= '3.10'",
"version": "==3.10.0" "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": { "numpy": {
"hashes": [ "hashes": [
"sha256:059e6a747ae84fce488c3ee397cee7e5f905fd1bda5fb18c66bc41807ff119b2", "sha256:059e6a747ae84fce488c3ee397cee7e5f905fd1bda5fb18c66bc41807ff119b2",
@ -462,13 +493,118 @@
"markers": "python_version >= '3.10'", "markers": "python_version >= '3.10'",
"version": "==2.2.1" "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": { "optype": {
"hashes": [ "hashes": [
"sha256:51c8dd104ac197457059bcee5e256160a641ca72c6c852012d952790a5e0cac0", "sha256:8cbfd452d6f06c7c70502048f38a0d5451bc601054d3a577dd09c7d6363950e1",
"sha256:856416484131038799e0e9cefc19d0ef37e7b4fde2144f25e8bb0e8981ebbe95" "sha256:90a7760177f2e7feae379a60445fceec37b932b75a00c3d96067497573c5e84d"
], ],
"markers": "python_version >= '3.10'", "markers": "python_version >= '3.10'",
"version": "==0.7.3" "version": "==0.8.0"
}, },
"packaging": { "packaging": {
"hashes": [ "hashes": [
@ -641,6 +777,14 @@
"markers": "python_version >= '3.10'", "markers": "python_version >= '3.10'",
"version": "==1.14.1.6" "version": "==1.14.1.6"
}, },
"setuptools": {
"hashes": [
"sha256:8199222558df7c86216af4f84c30e9b34a61d8ba19366cc914424cdbd28252f6",
"sha256:ce74b49e8f7110f9bf04883b730f4765b774ef3ef28f722cce7c273d253aaf7d"
],
"markers": "python_version >= '3.9'",
"version": "==75.6.0"
},
"simple-websocket": { "simple-websocket": {
"hashes": [ "hashes": [
"sha256:4af6069630a38ed6c561010f0e11a5bc0d4ca569b36306eb257cd9a192497c8c", "sha256:4af6069630a38ed6c561010f0e11a5bc0d4ca569b36306eb257cd9a192497c8c",
@ -657,6 +801,46 @@
"markers": "python_version >= '2.7' and python_version not in '3.0, 3.1, 3.2'", "markers": "python_version >= '2.7' and python_version not in '3.0, 3.1, 3.2'",
"version": "==1.17.0" "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": { "werkzeug": {
"hashes": [ "hashes": [
"sha256:54b78bf3716d19a65be4fceccc0d1d7b89e608834989dfae50ea87564639213e", "sha256:54b78bf3716d19a65be4fceccc0d1d7b89e608834989dfae50ea87564639213e",

View File

@ -1,5 +1,6 @@
import numpy as np import numpy as np
import numpy.typing as npt import torch
import torch.linalg
from . import config from . import config
from datetime import datetime from datetime import datetime
@ -7,16 +8,19 @@ import logging
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): def __init__(self):
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]): def smooth(self, data: 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
@ -31,32 +35,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]): def update(self, data: torch.Tensor):
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.AOA_SLIDING_WINDOW_SIZE WINDOW_SIZE = config.AOA_SLIDING_WINDOW_SIZE
@ -64,72 +68,78 @@ 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.
eigvals, eigvecs = np.linalg.eigh(R) eigvals, eigvecs = torch.linalg.eigh(R)
self.E_n = eigvecs[:, np.abs(eigvals) < config.EIGVAL_THRESHOLD] logging.debug(f"Eigenvalues: {eigvals}")
self.E_n = eigvecs[:, torch.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
)
def steering_vector(self, theta: float, tof: float): def steering_vector(self, theta: float, tof: float):
omega_t = np.exp(-2j * np.pi * config.DELTA_F * tof) omega_t = torch.exp(
phi_theta = np.exp( torch.tensor([-2j * torch.pi * config.DELTA_F * tof], dtype=torch.complex64)
2j )
* np.pi phi_theta = torch.exp(
* config.CENTRAL_FREQUENCY_HZ torch.tensor(
* config.ANTENNA_SPACING [
* (1 - np.cos(theta)) 2j
/ config.C * 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) omega_t = torch.unsqueeze(omega_t, dim=-1)
phi_theta = np.expand_dims(phi_theta, axis=-1) phi_theta = torch.unsqueeze(phi_theta, dim=-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 = torch.unsqueeze(antenna_v, dim=-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): def evaluate(self, theta: float, tof: float):
try: try:
steering = self.steering_vector(theta, tof) steering = self.steering_vector(theta, tof)
steering_h = np.conj(steering).T steering_h = torch.conj(steering).T
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 = torch.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 torch.abs(c.real)
def test_smoothing(): def test_smoothing():
row, col = np.indices((4, 2)) row, col = torch.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)
pass pass
def test_steering_vector(): def test_steering_vector():
aoa = AoA() aoa = AoA()
aoa.N_subcarriers = 10 aoa.N_subcarriers = 6
aoa.N_rx = 2 aoa.N_rx = 4
print(aoa.omega_base) tau = 1e-8
print(aoa.phi_base) theta = np.pi / 2
tau = 1
theta = 0
print(aoa.steering_vector(theta, tau)) print(aoa.steering_vector(theta, tau))
assert False assert False

View File

@ -94,13 +94,13 @@ class CSIHeader:
class CSI: class CSI:
@staticmethod @staticmethod
def parseCsiData(data: bytes, header: CSIHeader): 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_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,6 +3,7 @@ import time
import socket import socket
import multiprocessing as mp import multiprocessing as mp
import struct import struct
import torch
from datetime import datetime from datetime import datetime
import numpy as np import numpy as np
@ -12,6 +13,7 @@ from .preprocess import Preprocessor
from . import config from . import config
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
Host = tuple[str, int] Host = tuple[str, int]
@ -123,8 +125,10 @@ class CSIProcessor:
self.logger.info(f"Got final CSI data with shape {all_data.shape}") self.logger.info(f"Got final CSI data with shape {all_data.shape}")
processed = self.preprocess.preprocess(all_data) processed = self.preprocess.preprocess(all_data)
visualise.add_data(all_data, processed) processed_tensor = torch.tensor(processed, device=device)
self.aoa.update(processed) # visualise.add_data(all_data_tensor, processed)
self.aoa.update(processed_tensor)
self.logger.info("Processed data")
if not self.webserver.full(): if not self.webserver.full():
self.webserver.put(self.aoa) self.webserver.put(self.aoa)

View File

@ -15,9 +15,9 @@ np.seterr(invalid="ignore")
class Preprocessor: class Preprocessor:
def __init__(self): def __init__(self):
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,
[ [
@ -29,7 +29,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)),
@ -49,7 +49,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_LONG_TERM_ALPHA) self.long_term_avg * (1 - config.PREPROCESSING_LONG_TERM_ALPHA)