138 KiB
138 KiB
In [1]:
%matplotlib inline
import numpy as np
import math
from scipy import signal
import matplotlib.pyplot as plt
# number of points
n = 500
# start and final times
t0 = 0.0
tn = 1.0
# Period
t = np.linspace(t0, tn, n, endpoint=False)
SqrSignal = np.zeros(n)
SqrSignal = 1.0+signal.square(2*np.pi*5*t)
plt.plot(t, SqrSignal)
plt.ylim(-0.5, 2.5)
plt.show()In [2]:
import numpy as np
import math
from scipy import signal
import matplotlib.pyplot as plt
# number of points
n = 500
# start and final times
t0 = 0.0
tn = 1.0
# Period
T =0.2
# Max value of square signal
Fmax= 2.0
# Width of signal
Width = 0.1
t = np.linspace(t0, tn, n, endpoint=False)
SqrSignal = np.zeros(n)
FourierSeriesSignal = np.zeros(n)
SqrSignal = 1.0+signal.square(2*np.pi*5*t+np.pi*Width/T)
a0 = Fmax*Width/T
FourierSeriesSignal = a0
Factor = 2.0*Fmax/np.pi
for i in range(1,500):
FourierSeriesSignal += Factor/(i)*np.sin(np.pi*i*Width/T)*np.cos(i*t*2*np.pi/T)
plt.plot(t, SqrSignal)
plt.plot(t, FourierSeriesSignal)
plt.ylim(-0.5, 2.5)
plt.show()In [3]:
# import necessary packages
import numpy as np
import matplotlib.pyplot as plt
from sklearn import datasets
# ensure the same random numbers appear every time
np.random.seed(0)
# display images in notebook
%matplotlib inline
plt.rcParams['figure.figsize'] = (12,12)
# download MNIST dataset
digits = datasets.load_digits()
# define inputs and labels
inputs = digits.images
labels = digits.target
# RGB images have a depth of 3
# our images are grayscale so they should have a depth of 1
inputs = inputs[:,:,:,np.newaxis]
print("inputs = (n_inputs, pixel_width, pixel_height, depth) = " + str(inputs.shape))
print("labels = (n_inputs) = " + str(labels.shape))
# choose some random images to display
n_inputs = len(inputs)
indices = np.arange(n_inputs)
random_indices = np.random.choice(indices, size=5)
for i, image in enumerate(digits.images[random_indices]):
plt.subplot(1, 5, i+1)
plt.axis('off')
plt.imshow(image, cmap=plt.cm.gray_r, interpolation='nearest')
plt.title("Label: %d" % digits.target[random_indices[i]])
plt.show()inputs = (n_inputs, pixel_width, pixel_height, depth) = (1797, 8, 8, 1) labels = (n_inputs) = (1797,)
In [4]:
from tensorflow.keras import datasets, layers, models
from tensorflow.keras.layers import Input
from tensorflow.keras.models import Sequential #This allows appending layers to existing models
from tensorflow.keras.layers import Dense #This allows defining the characteristics of a particular layer
from tensorflow.keras import optimizers #This allows using whichever optimiser we want (sgd,adam,RMSprop)
from tensorflow.keras import regularizers #This allows using whichever regularizer we want (l1,l2,l1_l2)
from tensorflow.keras.utils import to_categorical #This allows using categorical cross entropy as the cost function
#from tensorflow.keras import Conv2D
#from tensorflow.keras import MaxPooling2D
#from tensorflow.keras import Flatten
from sklearn.model_selection import train_test_split
# representation of labels
labels = to_categorical(labels)
# split into train and test data
# one-liner from scikit-learn library
train_size = 0.8
test_size = 1 - train_size
X_train, X_test, Y_train, Y_test = train_test_split(inputs, labels, train_size=train_size,
test_size=test_size)/Users/mhjensen/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/_src/lib/__init__.py:32: UserWarning: JAX on Mac ARM machines is experimental and minimally tested. Please see https://github.com/google/jax/issues/5501 in the event of problems.
warnings.warn("JAX on Mac ARM machines is experimental and minimally tested. "
[0;31m---------------------------------------------------------------------------[0m [0;31mAttributeError[0m Traceback (most recent call last) Input [0;32mIn [4][0m, in [0;36m<cell line: 1>[0;34m()[0m [0;32m----> 1[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mkeras[39;00m [38;5;28;01mimport[39;00m datasets, layers, models [1;32m 2[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mkeras[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlayers[39;00m [38;5;28;01mimport[39;00m Input [1;32m 3[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mkeras[39;00m[38;5;21;01m.[39;00m[38;5;21;01mmodels[39;00m [38;5;28;01mimport[39;00m Sequential [38;5;66;03m#This allows appending layers to existing models[39;00m File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/__init__.py:51[0m, in [0;36m<module>[0;34m[0m [1;32m 49[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m [38;5;28;01mimport[39;00m autograph [1;32m 50[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m [38;5;28;01mimport[39;00m bitwise [0;32m---> 51[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m [38;5;28;01mimport[39;00m compat [1;32m 52[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m [38;5;28;01mimport[39;00m config [1;32m 53[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m [38;5;28;01mimport[39;00m data File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/__init__.py:37[0m, in [0;36m<module>[0;34m[0m [1;32m 3[0m [38;5;124;03m"""Compatibility functions.[39;00m [1;32m 4[0m [1;32m 5[0m [38;5;124;03mThe `tf.compat` module contains two sets of compatibility functions.[39;00m [0;32m (...)[0m [1;32m 32[0m [1;32m 33[0m [38;5;124;03m"""[39;00m [1;32m 35[0m [38;5;28;01mimport[39;00m [38;5;21;01msys[39;00m [38;5;28;01mas[39;00m [38;5;21;01m_sys[39;00m [0;32m---> 37[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m v1 [1;32m 38[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m v2 [1;32m 39[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m [38;5;28;01mimport[39;00m forward_compatibility_horizon File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/__init__.py:30[0m, in [0;36m<module>[0;34m[0m [1;32m 28[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m autograph [1;32m 29[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m bitwise [0;32m---> 30[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m compat [1;32m 31[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m config [1;32m 32[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m data File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/compat/__init__.py:37[0m, in [0;36m<module>[0;34m[0m [1;32m 3[0m [38;5;124;03m"""Compatibility functions.[39;00m [1;32m 4[0m [1;32m 5[0m [38;5;124;03mThe `tf.compat` module contains two sets of compatibility functions.[39;00m [0;32m (...)[0m [1;32m 32[0m [1;32m 33[0m [38;5;124;03m"""[39;00m [1;32m 35[0m [38;5;28;01mimport[39;00m [38;5;21;01msys[39;00m [38;5;28;01mas[39;00m [38;5;21;01m_sys[39;00m [0;32m---> 37[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m v1 [1;32m 38[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m v2 [1;32m 39[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m [38;5;28;01mimport[39;00m forward_compatibility_horizon File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/compat/v1/__init__.py:47[0m, in [0;36m<module>[0;34m[0m [1;32m 45[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv1[39;00m [38;5;28;01mimport[39;00m layers [1;32m 46[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv1[39;00m [38;5;28;01mimport[39;00m linalg [0;32m---> 47[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv1[39;00m [38;5;28;01mimport[39;00m lite [1;32m 48[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv1[39;00m [38;5;28;01mimport[39;00m logging [1;32m 49[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_api[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv2[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcompat[39;00m[38;5;21;01m.[39;00m[38;5;21;01mv1[39;00m [38;5;28;01mimport[39;00m lookup File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/lite/__init__.py:9[0m, in [0;36m<module>[0;34m[0m [1;32m 6[0m [38;5;28;01mimport[39;00m [38;5;21;01msys[39;00m [38;5;28;01mas[39;00m [38;5;21;01m_sys[39;00m [1;32m 8[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m constants [0;32m----> 9[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m experimental [1;32m 10[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m [38;5;28;01mimport[39;00m Interpreter [1;32m 11[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m [38;5;28;01mimport[39;00m OpHint File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/lite/experimental/__init__.py:8[0m, in [0;36m<module>[0;34m[0m [1;32m 3[0m [38;5;124;03m"""Public API for tf.lite.experimental namespace.[39;00m [1;32m 4[0m [38;5;124;03m"""[39;00m [1;32m 6[0m [38;5;28;01mimport[39;00m [38;5;21;01msys[39;00m [38;5;28;01mas[39;00m [38;5;21;01m_sys[39;00m [0;32m----> 8[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m authoring [1;32m 9[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01manalyzer[39;00m [38;5;28;01mimport[39;00m ModelAnalyzer [38;5;28;01mas[39;00m Analyzer [1;32m 10[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m [38;5;28;01mimport[39;00m OpResolverType File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/_api/v2/compat/v1/lite/experimental/authoring/__init__.py:8[0m, in [0;36m<module>[0;34m[0m [1;32m 3[0m [38;5;124;03m"""Public API for tf.lite.experimental.authoring namespace.[39;00m [1;32m 4[0m [38;5;124;03m"""[39;00m [1;32m 6[0m [38;5;28;01mimport[39;00m [38;5;21;01msys[39;00m [38;5;28;01mas[39;00m [38;5;21;01m_sys[39;00m [0;32m----> 8[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mauthoring[39;00m[38;5;21;01m.[39;00m[38;5;21;01mauthoring[39;00m [38;5;28;01mimport[39;00m compatible File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/lite/python/authoring/authoring.py:43[0m, in [0;36m<module>[0;34m[0m [1;32m 39[0m [38;5;28;01mimport[39;00m [38;5;21;01mfunctools[39;00m [1;32m 42[0m [38;5;66;03m# pylint: disable=g-import-not-at-top[39;00m [0;32m---> 43[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m [38;5;28;01mimport[39;00m convert [1;32m 44[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m [38;5;28;01mimport[39;00m lite [1;32m 45[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mmetrics[39;00m [38;5;28;01mimport[39;00m converter_error_data_pb2 File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/lite/python/convert.py:29[0m, in [0;36m<module>[0;34m[0m [1;32m 26[0m [38;5;28;01mimport[39;00m [38;5;21;01msix[39;00m [1;32m 28[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m [38;5;28;01mimport[39;00m lite_constants [0;32m---> 29[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m [38;5;28;01mimport[39;00m util [1;32m 30[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m [38;5;28;01mimport[39;00m wrap_toco [1;32m 31[0m [38;5;28;01mfrom[39;00m [38;5;21;01mtensorflow[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlite[39;00m[38;5;21;01m.[39;00m[38;5;21;01mpython[39;00m[38;5;21;01m.[39;00m[38;5;21;01mconvert_phase[39;00m [38;5;28;01mimport[39;00m Component File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/tensorflow/lite/python/util.py:51[0m, in [0;36m<module>[0;34m[0m [1;32m 47[0m [38;5;66;03m# Jax functions used by TFLite[39;00m [1;32m 48[0m [38;5;66;03m# pylint: disable=g-import-not-at-top[39;00m [1;32m 49[0m [38;5;66;03m# pylint: disable=unused-import[39;00m [1;32m 50[0m [38;5;28;01mtry[39;00m: [0;32m---> 51[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m [38;5;28;01mimport[39;00m xla_computation [38;5;28;01mas[39;00m _xla_computation [1;32m 52[0m [38;5;28;01mexcept[39;00m [38;5;167;01mImportError[39;00m: [1;32m 53[0m _xla_computation [38;5;241m=[39m [38;5;28;01mNone[39;00m File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/__init__.py:116[0m, in [0;36m<module>[0;34m[0m [1;32m 40[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mconfig[39;00m [38;5;28;01mimport[39;00m ( [1;32m 41[0m config [38;5;28;01mas[39;00m config, [1;32m 42[0m enable_checks [38;5;28;01mas[39;00m enable_checks, [0;32m (...)[0m [1;32m 51[0m numpy_rank_promotion [38;5;28;01mas[39;00m numpy_rank_promotion, [1;32m 52[0m ) [1;32m 53[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mapi[39;00m [38;5;28;01mimport[39;00m ( [1;32m 54[0m ad, [38;5;66;03m# TODO(phawkins): update users to avoid this.[39;00m [1;32m 55[0m checkpoint [38;5;28;01mas[39;00m checkpoint, [0;32m (...)[0m [1;32m 114[0m xla_computation [38;5;28;01mas[39;00m xla_computation, [1;32m 115[0m ) [0;32m--> 116[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01mexperimental[39;00m[38;5;21;01m.[39;00m[38;5;21;01mmaps[39;00m [38;5;28;01mimport[39;00m soft_pmap [38;5;28;01mas[39;00m soft_pmap [1;32m 117[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01mversion[39;00m [38;5;28;01mimport[39;00m __version__ [38;5;28;01mas[39;00m __version__ [1;32m 119[0m [38;5;66;03m# These submodules are separate because they are in an import cycle with[39;00m [1;32m 120[0m [38;5;66;03m# jax and rely on the names imported above.[39;00m File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/experimental/maps.py:26[0m, in [0;36m<module>[0;34m[0m [1;32m 23[0m [38;5;28;01mfrom[39;00m [38;5;21;01mfunctools[39;00m [38;5;28;01mimport[39;00m wraps, partial, partialmethod [1;32m 24[0m [38;5;28;01mfrom[39;00m [38;5;21;01menum[39;00m [38;5;28;01mimport[39;00m Enum [0;32m---> 26[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m numpy [38;5;28;01mas[39;00m jnp [1;32m 27[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m core [1;32m 28[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m[38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m linear_util [38;5;28;01mas[39;00m lu File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/numpy/__init__.py:19[0m, in [0;36m<module>[0;34m[0m [1;32m 1[0m [38;5;66;03m# Copyright 2018 Google LLC[39;00m [1;32m 2[0m [38;5;66;03m#[39;00m [1;32m 3[0m [38;5;66;03m# Licensed under the Apache License, Version 2.0 (the "License");[39;00m [0;32m (...)[0m [1;32m 17[0m [1;32m 18[0m [38;5;66;03m# flake8: noqa: F401[39;00m [0;32m---> 19[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m fft [38;5;28;01mas[39;00m fft [1;32m 20[0m [38;5;28;01mfrom[39;00m [38;5;21;01m.[39;00m [38;5;28;01mimport[39;00m linalg [38;5;28;01mas[39;00m linalg [1;32m 22[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01minterpreters[39;00m[38;5;21;01m.[39;00m[38;5;21;01mxla[39;00m [38;5;28;01mimport[39;00m DeviceArray [38;5;28;01mas[39;00m DeviceArray File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/numpy/fft.py:17[0m, in [0;36m<module>[0;34m[0m [1;32m 1[0m [38;5;66;03m# Copyright 2020 Google LLC[39;00m [1;32m 2[0m [38;5;66;03m#[39;00m [1;32m 3[0m [38;5;66;03m# Licensed under the Apache License, Version 2.0 (the "License");[39;00m [0;32m (...)[0m [1;32m 14[0m [1;32m 15[0m [38;5;66;03m# flake8: noqa: F401[39;00m [0;32m---> 17[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mnumpy[39;00m[38;5;21;01m.[39;00m[38;5;21;01mfft[39;00m [38;5;28;01mimport[39;00m ( [1;32m 18[0m ifft [38;5;28;01mas[39;00m ifft, [1;32m 19[0m ifft2 [38;5;28;01mas[39;00m ifft2, [1;32m 20[0m ifftn [38;5;28;01mas[39;00m ifftn, [1;32m 21[0m ifftshift [38;5;28;01mas[39;00m ifftshift, [1;32m 22[0m ihfft [38;5;28;01mas[39;00m ihfft, [1;32m 23[0m irfft [38;5;28;01mas[39;00m irfft, [1;32m 24[0m irfft2 [38;5;28;01mas[39;00m irfft2, [1;32m 25[0m irfftn [38;5;28;01mas[39;00m irfftn, [1;32m 26[0m fft [38;5;28;01mas[39;00m fft, [1;32m 27[0m fft2 [38;5;28;01mas[39;00m fft2, [1;32m 28[0m fftfreq [38;5;28;01mas[39;00m fftfreq, [1;32m 29[0m fftn [38;5;28;01mas[39;00m fftn, [1;32m 30[0m fftshift [38;5;28;01mas[39;00m fftshift, [1;32m 31[0m hfft [38;5;28;01mas[39;00m hfft, [1;32m 32[0m rfft [38;5;28;01mas[39;00m rfft, [1;32m 33[0m rfft2 [38;5;28;01mas[39;00m rfft2, [1;32m 34[0m rfftfreq [38;5;28;01mas[39;00m rfftfreq, [1;32m 35[0m rfftn [38;5;28;01mas[39;00m rfftn, [1;32m 36[0m ) [1;32m 38[0m [38;5;66;03m# Module initialization is encapsulated in a function to avoid accidental[39;00m [1;32m 39[0m [38;5;66;03m# namespace pollution.[39;00m [1;32m 40[0m _NOT_IMPLEMENTED [38;5;241m=[39m [] File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/_src/numpy/fft.py:19[0m, in [0;36m<module>[0;34m[0m [1;32m 16[0m [38;5;28;01mimport[39;00m [38;5;21;01moperator[39;00m [1;32m 17[0m [38;5;28;01mimport[39;00m [38;5;21;01mnumpy[39;00m [38;5;28;01mas[39;00m [38;5;21;01mnp[39;00m [0;32m---> 19[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m [38;5;28;01mimport[39;00m lax [1;32m 20[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlib[39;00m [38;5;28;01mimport[39;00m xla_client [1;32m 21[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mutil[39;00m [38;5;28;01mimport[39;00m safe_zip File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/lax/__init__.py:332[0m, in [0;36m<module>[0;34m[0m [1;32m 299[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m [38;5;28;01mimport[39;00m (_reduce_sum, _reduce_max, _reduce_min, _reduce_or, [1;32m 300[0m _reduce_and, _reduce_window_sum, _reduce_window_max, [1;32m 301[0m _reduce_window_min, _reduce_window_prod, [0;32m (...)[0m [1;32m 306[0m _upcast_fp16_for_computation, _broadcasting_shape_rule, [1;32m 307[0m _eye, _tri, _delta, _ones, _zeros, _dilate_shape) [1;32m 308[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m[38;5;21;01m.[39;00m[38;5;21;01mcontrol_flow[39;00m [38;5;28;01mimport[39;00m ( [1;32m 309[0m associative_scan [38;5;28;01mas[39;00m associative_scan, [1;32m 310[0m cond [38;5;28;01mas[39;00m cond, [0;32m (...)[0m [1;32m 330[0m while_p [38;5;28;01mas[39;00m while_p, [1;32m 331[0m ) [0;32m--> 332[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m[38;5;21;01m.[39;00m[38;5;21;01mfft[39;00m [38;5;28;01mimport[39;00m ( [1;32m 333[0m fft [38;5;28;01mas[39;00m fft, [1;32m 334[0m fft_p [38;5;28;01mas[39;00m fft_p, [1;32m 335[0m ) [1;32m 336[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m[38;5;21;01m.[39;00m[38;5;21;01mparallel[39;00m [38;5;28;01mimport[39;00m ( [1;32m 337[0m all_gather [38;5;28;01mas[39;00m all_gather, [1;32m 338[0m all_to_all [38;5;28;01mas[39;00m all_to_all, [0;32m (...)[0m [1;32m 355[0m xeinsum [38;5;28;01mas[39;00m xeinsum, [1;32m 356[0m ) [1;32m 357[0m [38;5;28;01mfrom[39;00m [38;5;21;01mjax[39;00m[38;5;21;01m.[39;00m[38;5;21;01m_src[39;00m[38;5;21;01m.[39;00m[38;5;21;01mlax[39;00m[38;5;21;01m.[39;00m[38;5;21;01mother[39;00m [38;5;28;01mimport[39;00m ( [1;32m 358[0m conv_general_dilated_patches [38;5;28;01mas[39;00m conv_general_dilated_patches [1;32m 359[0m ) File [0;32m~/miniforge3/envs/myenv/lib/python3.9/site-packages/jax/_src/lax/fft.py:145[0m, in [0;36m<module>[0;34m[0m [1;32m 143[0m batching[38;5;241m.[39mprimitive_batchers[fft_p] [38;5;241m=[39m fft_batching_rule [1;32m 144[0m [38;5;28;01mif[39;00m pocketfft: [0;32m--> 145[0m xla[38;5;241m.[39mbackend_specific_translations[[38;5;124m'[39m[38;5;124mcpu[39m[38;5;124m'[39m][fft_p] [38;5;241m=[39m [43mpocketfft[49m[38;5;241;43m.[39;49m[43mpocketfft[49m [0;31mAttributeError[0m: module 'jaxlib.pocketfft' has no attribute 'pocketfft'
In [5]:
def create_convolutional_neural_network_keras(input_shape, receptive_field,
n_filters, n_neurons_connected, n_categories,
eta, lmbd):
model = Sequential()
model.add(layers.Conv2D(n_filters, (receptive_field, receptive_field), input_shape=input_shape, padding='same',
activation='relu', kernel_regularizer=regularizers.l2(lmbd)))
model.add(layers.MaxPooling2D(pool_size=(2, 2)))
model.add(layers.Flatten())
model.add(layers.Dense(n_neurons_connected, activation='relu', kernel_regularizer=regularizers.l2(lmbd)))
model.add(layers.Dense(n_categories, activation='softmax', kernel_regularizer=regularizers.l2(lmbd)))
sgd = optimizers.SGD(lr=eta)
model.compile(loss='categorical_crossentropy', optimizer=sgd, metrics=['accuracy'])
return model
epochs = 100
batch_size = 100
input_shape = X_train.shape[1:4]
receptive_field = 3
n_filters = 10
n_neurons_connected = 50
n_categories = 10
eta_vals = np.logspace(-5, 1, 7)
lmbd_vals = np.logspace(-5, 1, 7)In [6]:
CNN_keras = np.zeros((len(eta_vals), len(lmbd_vals)), dtype=object)
for i, eta in enumerate(eta_vals):
for j, lmbd in enumerate(lmbd_vals):
CNN = create_convolutional_neural_network_keras(input_shape, receptive_field,
n_filters, n_neurons_connected, n_categories,
eta, lmbd)
CNN.fit(X_train, Y_train, epochs=epochs, batch_size=batch_size, verbose=0)
scores = CNN.evaluate(X_test, Y_test)
CNN_keras[i][j] = CNN
print("Learning rate = ", eta)
print("Lambda = ", lmbd)
print("Test accuracy: %.3f" % scores[1])
print()In [7]:
# visual representation of grid search
# uses seaborn heatmap, could probably do this in matplotlib
import seaborn as sns
sns.set()
train_accuracy = np.zeros((len(eta_vals), len(lmbd_vals)))
test_accuracy = np.zeros((len(eta_vals), len(lmbd_vals)))
for i in range(len(eta_vals)):
for j in range(len(lmbd_vals)):
CNN = CNN_keras[i][j]
train_accuracy[i][j] = CNN.evaluate(X_train, Y_train)[1]
test_accuracy[i][j] = CNN.evaluate(X_test, Y_test)[1]
fig, ax = plt.subplots(figsize = (10, 10))
sns.heatmap(train_accuracy, annot=True, ax=ax, cmap="viridis")
ax.set_title("Training Accuracy")
ax.set_ylabel("$\eta$")
ax.set_xlabel("$\lambda$")
plt.show()
fig, ax = plt.subplots(figsize = (10, 10))
sns.heatmap(test_accuracy, annot=True, ax=ax, cmap="viridis")
ax.set_title("Test Accuracy")
ax.set_ylabel("$\eta$")
ax.set_xlabel("$\lambda$")
plt.show()In [8]:
import tensorflow as tf
from tensorflow.keras import datasets, layers, models
import matplotlib.pyplot as plt
# We import the data set
(train_images, train_labels), (test_images, test_labels) = datasets.cifar10.load_data()
# Normalize pixel values to be between 0 and 1 by dividing by 255.
train_images, test_images = train_images / 255.0, test_images / 255.0In [9]:
class_names = ['airplane', 'automobile', 'bird', 'cat', 'deer',
'dog', 'frog', 'horse', 'ship', 'truck']
plt.figure(figsize=(10,10))
for i in range(25):
plt.subplot(5,5,i+1)
plt.xticks([])
plt.yticks([])
plt.grid(False)
plt.imshow(train_images[i], cmap=plt.cm.binary)
# The CIFAR labels happen to be arrays,
# which is why you need the extra index
plt.xlabel(class_names[train_labels[i][0]])
plt.show()In [10]:
model = models.Sequential()
model.add(layers.Conv2D(32, (3, 3), activation='relu', input_shape=(32, 32, 3)))
model.add(layers.MaxPooling2D((2, 2)))
model.add(layers.Conv2D(64, (3, 3), activation='relu'))
model.add(layers.MaxPooling2D((2, 2)))
model.add(layers.Conv2D(64, (3, 3), activation='relu'))
# Let's display the architecture of our model so far.
model.summary()In [11]:
model.add(layers.Flatten())
model.add(layers.Dense(64, activation='relu'))
model.add(layers.Dense(10))
Here's the complete architecture of our model.
model.summary()In [12]:
model.compile(optimizer='adam',
loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),
metrics=['accuracy'])
history = model.fit(train_images, train_labels, epochs=10,
validation_data=(test_images, test_labels))In [13]:
plt.plot(history.history['accuracy'], label='accuracy')
plt.plot(history.history['val_accuracy'], label = 'val_accuracy')
plt.xlabel('Epoch')
plt.ylabel('Accuracy')
plt.ylim([0.5, 1])
plt.legend(loc='lower right')
test_loss, test_acc = model.evaluate(test_images, test_labels, verbose=2)
print(test_acc)
