import numpy as np
import matplotlib.pyplot as plt
from scipy.fftpack import fft, ifft
from scipy.io.wavfile import read


def find_argmax(arr, num, win):
    cp = np.copy(arr)
    res = np.zeros(num, dtype=int)
    for ii in range(num):
        res[ii] = np.argmax(cp)
        cp[max(0, res[ii] - win): min(res[ii] + win, cp.shape[0] - 1)] = np.zeros_like(
            cp[max(0, res[ii] - win): min(res[ii] + win, cp.shape[0] - 1)])
    return res


num_freq = 4

input_signal = read('signal.wav')[1]
half_size = input_signal.shape[0] // 2
inp_fft = fft(input_signal)[:half_size]

fig, ax = plt.subplots()
ax.plot(np.abs(inp_fft[:1000]))
ax.set_title('Спектр сигнала')
ax.set_xlabel('Гц')
ax.grid()
plt.savefig('spec.png', dpi=500)
plt.show()

arg = find_argmax(inp_fft, num_freq, 10)
print("obtained freqs >> ", arg)

width = 100
im_shape = [10, 9, 8]

fig, ax = plt.subplots(3, 2, figsize=(10, 10))

for i in range(num_freq):
    ll = arg[i] - width
    rr = arg[i] + width
    new_signal = np.zeros_like(inp_fft)
    new_signal[ll: rr] = inp_fft[ll: rr]
    new_signal[half_size - ll: half_size - rr] = inp_fft[half_size - ll:  half_size - rr]
    new_signal_ifft = np.abs(ifft(new_signal))
    ax[i][0].plot(new_signal_ifft)
    ax[i][0].set_title(str(arg[i]) + ' Гц')

    image = np.zeros(shape=(im_shape[i], im_shape[i]), dtype=int)
    row = int(new_signal_ifft.shape[0] / im_shape[i] / im_shape[i])
    for j in range(im_shape[i]):
        for k in range(im_shape[i]):
            image[j, k] = np.mean(new_signal_ifft[(j * im_shape[i] + k) * row: (j * im_shape[i] + k + 1) * row])
    ax[i][1].imshow(image, extent='lower')

plt.savefig('res.png', dpi=1000, bbox_inches='tight')
plt.show()
