refactoring into vqvae_util.py

master
David 2020-12-27 15:52:38 +10:30
parent 1c7fb55cc0
commit 4c954d8dea
3 changed files with 113 additions and 155 deletions

View File

@ -22,6 +22,7 @@ logging.getLogger('tensorflow').setLevel(logging.FATAL)
import tensorflow as tf
from vqvae_twostage import *
from vqvae_util import *
# Constants -------------------------------------------------
@ -39,9 +40,9 @@ parser.add_argument('--nb_samples', type=int, default=1000000, help='Number of f
parser.add_argument('--embedding_dim', type=int, default=16, help='dimension of embedding vectors')
parser.add_argument('--num_embedding', type=int, default=2048, help='number of embedded vectors')
parser.add_argument('--scale', type=float, default=0.125, help='apply this gain to features when read in')
parser.add_argument('--nnout', type=str, default="", help='Name of output NN we have trained')
parser.add_argument('--nnout', type=str, help='Name of output NN we have trained')
parser.add_argument('--mean', action='store_true', help='Extract mean from each chunk')
parser.add_argument('--mean_thresh', type=float, default=0, help='Discard chunks with less than this mean threshold')
parser.add_argument('--mean_thresh', type=float, default=0.0, help='Discard chunks with less than this mean threshold')
args = parser.parse_args()
dim = args.embedding_dim
nb_samples = args.nb_samples
@ -133,24 +134,15 @@ vqvae.get_layer('vq2').set_vq(vq_initial)
history = vqvae.fit(train, train_target, batch_size=batch_size, epochs=args.epochs,
validation_split=validation_split,callbacks=[CustomCallback()])
with open(args.nnout, 'wb') as f:
np.save(f, vqvae.get_layer("conv1d_a").get_weights(), allow_pickle=True)
np.save(f, vqvae.get_layer("conv1d_b").get_weights(), allow_pickle=True)
np.save(f, vqvae.get_layer("vq1").get_vq(), allow_pickle=True)
np.save(f, vqvae.get_layer("vq2").get_vq(), allow_pickle=True)
np.save(f, vqvae.get_layer("conv1d_c").get_weights(), allow_pickle=True)
np.save(f, vqvae.get_layer("conv1d_d").get_weights(), allow_pickle=True)
np.save(f, vqvae.get_layer("conv1d_e").get_weights(), allow_pickle=True)
'''
with open('test.npy', 'rb') as f:
vqvae.get_layer("conv1d_a").set_weights(np.load(f, allow_pickle=True))
vqvae.get_layer("conv1d_b").set_weights(np.load(f, allow_pickle=True))
vqvae.get_layer("vq1").set_vq(np.load(f, allow_pickle=True))
vqvae.get_layer("vq2").set_vq(np.load(f, allow_pickle=True))
vqvae.get_layer("conv1d_c").set_weights(np.load(f, allow_pickle=True))
vqvae.get_layer("conv1d_d").set_weights(np.load(f, allow_pickle=True))
vqvae.get_layer("conv1d_e").set_weights(np.load(f, allow_pickle=True))
'''
if args.nnout is not None:
with open(args.nnout, 'wb') as f:
np.save(f, vqvae.get_layer("conv1d_a").get_weights(), allow_pickle=True)
np.save(f, vqvae.get_layer("conv1d_b").get_weights(), allow_pickle=True)
np.save(f, vqvae.get_layer("vq1").get_vq(), allow_pickle=True)
np.save(f, vqvae.get_layer("vq2").get_vq(), allow_pickle=True)
np.save(f, vqvae.get_layer("conv1d_c").get_weights(), allow_pickle=True)
np.save(f, vqvae.get_layer("conv1d_d").get_weights(), allow_pickle=True)
np.save(f, vqvae.get_layer("conv1d_e").get_weights(), allow_pickle=True)
vq_weights = vqvae.get_layer('vq1').get_vq()
@ -184,42 +176,23 @@ plt.show(block=False)
# Calculate total mean square error and mse per frame
def calc_mse(train, train_est, nb_samples, nb_features, dec):
msepf = np.zeros(nb_samples-dec)
e1 = 0; n = 0
for i in range(nb_samples-dec):
e = (train_est[i,:] - train[i,:])**2
msepf[i] = np.mean(e)
e1 += np.sum(e); n += nb_features
mse = e1/n
return mse, msepf
print("mse",train_target.shape, train_est.shape)
mse,msepf = calc_mse(train_target, train_est, nb_samples, nb_features, 1)
print("mse: %4.2f dB*dB" % (mse))
worst_fr = np.argsort(-msepf);
worst_mse = np.sort(-msepf);
print(worst_fr[:10], worst_mse[:10]);
plt.figure(3)
plt.plot(msepf)
plt.title('Spectral Distortion dB*dB per frame')
plt.show(block=False)
def reject_outliers(data, m=2):
return data[abs(data - np.mean(data)) < m * np.std(data)]
plt.figure(4)
plt.title('Histogram of Spectral Distortion dB*dB out to 2*sigma')
plt.hist(reject_outliers(msepf), bins='fd')
plt.show(block=False)
# Count how many times each vector is used
def vector_count(x, vq, dim, nb_vecs):
# VQ search outside of Keras Backend
flat_inputs = np.reshape(x, (-1, dim))
distances = np.sum(flat_inputs**2, axis=1, keepdims=True) - 2* np.dot(flat_inputs, vq.T) + np.sum(vq.T ** 2, axis=0, keepdims=True)
encoding_indices = np.argmax(-distances, axis=1)
count = np.zeros(nb_vecs, dtype="int")
count[encoding_indices] += 1
return count
count = np.zeros(args.num_embedding, dtype="int")
for i in range(0, nb_samples, batch_size):
count += vector_count(encoder_out[i:i+batch_size], vq_weights, dim, args.num_embedding)
@ -235,23 +208,6 @@ plt.hist(train_mean, bins='fd')
plt.show(block=False)
plt.title('Mean of each chunk')
# use PCA to plot encoder space and VQ in 2D -----------------------------------------
# https://towardsdatascience.com/principal-component-analysis-pca-from-scratch-in-python-7f3e2a540c51
def find_pca(A):
# calculate the mean of each column
M = np.mean(A.T, axis=1)
# center columns by subtracting column means
C = A - M
# calculate covariance matrix of centered matrix
V = np.cov(C.T)
# eigendecomposition of covariance matrix
values, vectors = np.linalg.eig(V)
#print(vectors)
#print(values)
P = vectors.T.dot(C.T)
return P.T
fig,ax = plt.subplots()
encoder_pca=find_pca(encoder_out)
ax.hist2d(encoder_pca[:,0],encoder_pca[:,1], bins=(50,50))
@ -264,37 +220,4 @@ print("Press any key to start VQ pager....")
key = getch.getch()
plt.close('all')
# VQ Pager - plot input/output spectra to sanity check
nb_plots = 8
fs = 100;
key = ' '
while key != 'q':
frames=range(fs,fs+nb_plots)
nb_plotsy = np.floor(np.sqrt(nb_plots)); nb_plotsx=nb_plots/nb_plotsy;
plt.figure(8)
plt.clf()
plt.tight_layout()
plt.title('Rate K Amplitude Spectra')
for r in range(nb_plots):
plt.subplot(nb_plotsy,nb_plotsx,r+1)
f = frames[r];
plt.plot(train_target[f,:],'g')
plt.plot(train_est[f,:],'r')
plt.ylim(0,80)
a_mse = np.mean((train_target[f,:]-train_est[f,:])**2)
t = "f: %d %3.1f" % (f, a_mse)
plt.title(t)
plt.show(block=False)
plt.pause(0.0001)
print("n-next b-back s-save_png q-quit", end='\r', flush=True);
key = getch.getch()
if key == 'n':
fs += nb_plots
if key == 'b':
fs -= nb_plots
if key == 's':
plt.savefig('vqvae_spectra.png')
plt.close()
vqvae_pager(8,0,train_target,train_est,worst_fr)

View File

@ -7,7 +7,7 @@
$ ./vq_vae_kmeans_conv1d.py dev-clean-8k-K20.f32 --eband_K 20 --epochs 5 --scale 0.005 --nnout test.npy
-> 11.48dB*dB
$ sox -t .sw -r 8000 ~/Downloads/train_8k.sw -t .sw - trim 0 2.5 | ~/codec2/build_linux/src/c2sim - --rateK --rateKout test.f32
$ ./vq_vae_kmeans_conv1d_out.py test.npy test.f32 test_out.f32 --eband_K 20 --scale 0.005
$ ./vq_vae_kmeans_conv1d_out.py test.npy test.f32 --featurefile_out test_out.f32 --eband_K 20 --scale 0.005
-> 21.6 dB*dB
$ sox -t .sw -c 1 -r 8000 ~/Downloads/train_8k.sw -t .sw - trim 0 2.5 | ~/codec2/build_linux/src/c2sim - --rateK --rateKin test_out.f32 -o test1.raw
@ -24,6 +24,7 @@ logging.getLogger('tensorflow').setLevel(logging.FATAL)
import tensorflow as tf
from vqvae_twostage import *
from vqvae_util import *
# Constants -------------------------------------------------
@ -36,7 +37,7 @@ nb_timesteps = 4
parser = argparse.ArgumentParser(description='Two stage VQ-VAE for rate K vectors')
parser.add_argument('ampnn', help='model weights and VQs in .npy')
parser.add_argument('featurefile', help='input f32 file of spectral mag vectors, each element is 10*log10(energy), i.e. dB')
parser.add_argument('featurefile_out', help='output f32 file of spectral mag vectors, each element is 10*log10(energy), i.e. dB')
parser.add_argument('--featurefile_out', help='output f32 file of spectral mag vectors, each element is 10*log10(energy), i.e. dB')
parser.add_argument('--nb_samples', type=int, default=1000000, help='Number of frames to train on')
parser.add_argument('--eband_K', type=int, default=14, help='Length of eband vector')
parser.add_argument('--embedding_dim', type=int, default=16, help='dimension of embedding vectors')
@ -122,29 +123,23 @@ encoder_out = encoder_out.reshape(-1, dim)
print("target_est", target_est.shape, nb_samples)
# make output file the same size despite chunking
target_est = np.concatenate((target_est, np.zeros((nb_samples_file-nb_samples,nb_features))))
print(target_est.shape, nb_samples_file)
target_est_out = target_est.astype(np.float32);
print(features.shape, target_est_out.shape, target_est_out.dtype);
target_est_out.tofile(args.featurefile_out)
if args.featurefile_out is not None:
target_est = np.concatenate((target_est, np.zeros((nb_samples_file-nb_samples,nb_features))))
print(target_est.shape, nb_samples_file)
target_est_out = target_est.astype(np.float32);
print(features.shape, target_est_out.shape, target_est_out.dtype);
target_est_out.tofile(args.featurefile_out)
# Plot training results -------------------------
# Calculate total mean square error and mse per frame
def calc_mse(train, train_est, nb_samples, nb_features, dec):
msepf = np.zeros(nb_samples-dec)
e1 = 0; n = 0
for i in range(nb_samples-dec):
e = (train_est[i,:] - train[i,:])**2
msepf[i] = np.mean(e)
e1 += np.sum(e); n += nb_features
mse = e1/n
return mse, msepf
print("mse",target.shape, target_est.shape)
mse,msepf = calc_mse(target, target_est, nb_samples, nb_features, 1)
print("mse: %4.2f dB*dB" % (mse))
worst_fr = np.argsort(-msepf);
worst_mse = np.sort(-msepf);
print(worst_fr[:10], worst_mse[:10]);
if args.plots == False:
quit();
@ -153,23 +148,11 @@ plt.plot(msepf)
plt.title('Spectral Distortion dB*dB per frame')
plt.show(block=False)
def reject_outliers(data, m=2):
return data[abs(data - np.mean(data)) < m * np.std(data)]
plt.figure(2)
plt.title('Histogram of Spectral Distortion dB*dB out to 2*sigma')
plt.hist(reject_outliers(msepf), bins='fd')
plt.show(block=False)
# Count how many times each vector is used
def vector_count(x, vq, dim, nb_vecs):
# VQ search outside of Keras Backend
flat_inputs = np.reshape(x, (-1, dim))
distances = np.sum(flat_inputs**2, axis=1, keepdims=True) - 2* np.dot(flat_inputs, vq.T) + np.sum(vq.T ** 2, axis=0, keepdims=True)
encoding_indices = np.argmax(-distances, axis=1)
count = np.zeros(nb_vecs, dtype="int")
count[encoding_indices] += 1
return count
count = np.zeros(args.num_embedding, dtype="int")
for i in range(0, nb_samples, batch_size):
count += vector_count(encoder_out[i:i+batch_size], vq_weights, dim, args.num_embedding)
@ -185,37 +168,5 @@ plt.hist(target_mean, bins='fd')
plt.show(block=False)
plt.title('Mean of each chunk')
# VQ Pager - plot input/output spectra to sanity check
nb_plots = 8
fs = 0;
key = ' '
while key != 'q':
frames=range(fs,fs+nb_plots)
nb_plotsy = np.floor(np.sqrt(nb_plots)); nb_plotsx=nb_plots/nb_plotsy;
plt.figure(5)
plt.clf()
plt.tight_layout()
plt.title('Rate K Amplitude Spectra')
for r in range(nb_plots):
plt.subplot(nb_plotsy,nb_plotsx,r+1)
f = frames[r];
plt.plot(target[f,:],'g')
plt.plot(target_est[f,:],'r')
plt.ylim(0,80)
a_mse = np.mean((target[f,:]-target_est[f,:])**2)
t = "f: %d %3.1f" % (f, a_mse)
plt.title(t)
plt.show(block=False)
plt.pause(0.0001)
print("n-next b-back s-save_png q-quit", end='\r', flush=True);
key = getch.getch()
if key == 'n':
fs += nb_plots
if key == 'b':
fs -= nb_plots
if key == 's':
plt.savefig('vqvae_spectra.png')
plt.close()
vqvae_pager(5,0,target,target_est,worst_fr)

84
vqvae_util.py 100644
View File

@ -0,0 +1,84 @@
'''
VQVAE utility functions
'''
import numpy as np
from matplotlib import pyplot as plt
import getch
# VQ Pager - plot input/output spectra to sanity check
def vqvae_pager(fig,fs,target,est,worst_fr):
nb_plots = 8
w = 0;
key = ' '
while key != 'q':
frames=range(fs,fs+nb_plots)
nb_plotsy = np.floor(np.sqrt(nb_plots)); nb_plotsx=nb_plots/nb_plotsy;
plt.figure(fig)
plt.clf()
plt.tight_layout()
plt.title('Rate K Amplitude Spectra')
for r in range(nb_plots):
plt.subplot(nb_plotsy,nb_plotsx,r+1)
f = frames[r];
plt.plot(target[f,:],'g')
plt.plot(est[f,:],'r')
plt.ylim(0,80)
a_mse = np.mean((target[f,:]-est[f,:])**2)
t = "f: %d %3.1f" % (f, a_mse)
plt.title(t)
plt.show(block=False)
plt.pause(0.0001)
print("n-next b-back w-worst s-save_png q-quit", end='\r', flush=True);
key = getch.getch()
if key == 'n':
fs += nb_plots
if key == 'b':
fs -= nb_plots
if key == 's':
plt.savefig('vqvae_spectra.png')
if key == 's':
plt.savefig('vqvae_spectra.png')
if key == 'w':
fs = min(worst_fr[w],target.shape[0]-nb_plots)
w += 1;
plt.close()
# Calculate total mean square error and mse per frame
def calc_mse(train, train_est, nb_samples, nb_features, dec):
msepf = np.zeros(nb_samples-dec)
e1 = 0; n = 0
for i in range(nb_samples-dec):
e = (train_est[i,:] - train[i,:])**2
msepf[i] = np.mean(e)
e1 += np.sum(e); n += nb_features
mse = e1/n
return mse, msepf
def reject_outliers(data, m=2):
return data[abs(data - np.mean(data)) < m * np.std(data)]
# Count how many times each vector is used
def vector_count(x, vq, dim, nb_vecs):
# VQ search outside of Keras Backend
flat_inputs = np.reshape(x, (-1, dim))
distances = np.sum(flat_inputs**2, axis=1, keepdims=True) - 2* np.dot(flat_inputs, vq.T) + np.sum(vq.T ** 2, axis=0, keepdims=True)
encoding_indices = np.argmax(-distances, axis=1)
count = np.zeros(nb_vecs, dtype="int")
count[encoding_indices] += 1
return count
# use PCA to plot encoder space and VQ in 2D
# https://towardsdatascience.com/principal-component-analysis-pca-from-scratch-in-python-7f3e2a540c51
def find_pca(A):
# calculate the mean of each column
M = np.mean(A.T, axis=1)
# center columns by subtracting column means
C = A - M
# calculate covariance matrix of centered matrix
V = np.cov(C.T)
# eigendecomposition of covariance matrix
values, vectors = np.linalg.eig(V)
P = vectors.T.dot(C.T)
return P.T