Repository navigation
Expand file tree
/
Copy pathML_process.py
More file actions
105 lines (88 loc) · 3.08 KB
/
Copy pathML_process.py
File metadata and controls
105 lines (88 loc) · 3.08 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
import taichi as ti
import numpy as np
import zmq
ti.init(ti.gpu)
# Create ZMQ context and SUB socket
ctx = zmq.Context()
socket = ctx.socket(zmq.SUB)
socket.connect("tcp://127.0.0.1:5555") # must match your GRC sink
socket.setsockopt(zmq.SUBSCRIBE, b"") # subscribe to all topics
def flush_socket():
"""Drain any queued messages."""
try:
while True:
socket.recv(flags=zmq.NOBLOCK)
except zmq.Again:
pass # queue is empty
N =2048
RES = (N//2,N//4)
window = ti.ui.Window("ML Process",res=RES)
canvas = window.get_canvas()
gui = window.get_gui()
screen = ti.Vector.field(3,ti.f32,shape=RES)
screen_np = screen.to_numpy()
print(screen_np.shape)
X = ti.Vector.field(2,ti.f32,shape=N)
X_ref = 0.75
gains = ti.field(ti.f32,shape=2)
PS = ti.Vector.field(2,ti.f32,shape=N)
PS_ref = 0.5
def init_signal():
X_np = X.to_numpy()
X_np[:,0] = np.linspace(0,1,N,endpoint=True)
X_np[:,1] = X_ref
X.from_numpy(X_np)
PS_np = PS.to_numpy()
PS_np[:,0] = np.linspace(0,1,N,endpoint=True)
PS_np[:,1] = PS_ref
PS.from_numpy(PS_np)
gains[0],gains[1] = 1.0,10.0
init_signal()
def softmax(x):
e_x = np.exp(x - np.max(x)) # stable
return e_x / np.sum(e_x)
def sample(on):
if on:
try:
msg = socket.recv(flags=zmq.NOBLOCK) # Receive a message (bytes)
np_recieved_signal = np.frombuffer(msg, dtype=np.float32)
time_step = np_recieved_signal.shape[0]
#gui.text(f'{time_step}')
X_np = X.to_numpy()
X_np[:,1] = np.roll(X_np[:,1],time_step)
X_np[:time_step,1] = X_ref + gains[0]/ 100 * np_recieved_signal
X.from_numpy(X_np)
PS_np = PS.to_numpy()
spectrum = np.fft.rfft((X_np[:,1] - X_ref)*100,n = 2*N)[:N]
power_spectrum = np.abs(spectrum**2) / N
log_spectrum = np.log1p(power_spectrum)
cepstrum = np.fft.irfft(log_spectrum,n = 2*N)[:N] * (np.linspace(1,20,N)) * gains[1]
cepstrum = 1 - 2 / (1 + np.exp(-cepstrum))
PS_np[:,1] = cepstrum + PS_ref
PS.from_numpy(PS_np)
screen_np = screen.to_numpy()
screen_np = np.roll(screen_np, shift=1, axis=0)
screen_np[0,:,0] = 0
screen_np[0,:,0] = cepstrum[:RES[1]] * 10
screen_np[0,:,1] = screen_np[0,:,0]
screen_np[0,:,2] = screen_np[0,:,0]
#screen_np[RES[0]-1,:,:] *= 0
screen.from_numpy(screen_np)
except zmq.Again:
pass # no new message yet
else:
flush_socket() # actively discard any buffered messages
def handler():
gains[0] = gui.slider_float("signal gain",gains[0],1e-9,20.0)
gains[1] = gui.slider_float("power spectrum gain",gains[1],1e-9,10.0)
rec_on = 1
while window.running:
handler()
rec_on = gui.checkbox('record on',rec_on)
for i in range(1):
sample(rec_on)
sample(rec_on)
canvas.set_image(screen)
canvas.circles(X,2e-3,color=(0.1,1,0.1))
canvas.circles(PS,2e-3,color=(0.5,0.5,0.9))
window.show()