import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
import matplotlib.pyplot as plt
import plotly.io as pio
import importlib.util, shutil
if shutil.which('google-chrome') and importlib.util.find_spec('kaleido'): # a static PNG copy needs kaleido + Chrome
pio.renderers.default = 'plotly_mimetype+png' # interactive where supported, plus a static PNG
from PIL import Image, ImageDraw
device = 'cuda' if torch.cuda.is_available() else 'cpu'
_ = torch.manual_seed(0); np.random.seed(0)
torch.set_num_threads(max(1, torch.get_num_threads()))
plt.rcParams.update({'figure.dpi': 130, 'savefig.dpi': 130, 'image.interpolation': 'nearest', 'axes.grid': False})
print('device:', device)
# ---- the toy colored-shapes dataset (Figure 30.4 of the book) ----
COLORS = {'red':(220,50,50),'orange':(240,150,30),'yellow':(235,220,40),'green':(60,180,75),
'cyan':(40,200,200),'blue':(60,90,220),'purple':(150,60,200),'pink':(240,110,180)}
CNAMES = list(COLORS); SHAPES = ['circle','square','triangle']
def _one(rng, S=32):
ci, si = rng.integers(8), rng.integers(3)
base = np.array(COLORS[CNAMES[ci]], float)
col = tuple(int(np.clip(c + rng.normal(0, 12), 0, 255)) for c in base) # small color jitter
img = Image.new('RGB', (S, S), (20, 20, 20)); d = ImageDraw.Draw(img)
r = rng.integers(int(S*0.22), int(S*0.36)) # random size
cx, cy = rng.integers(r, S-r), rng.integers(r, S-r) # random position
ang = rng.uniform(0, 360) # random rotation
if si == 0:
d.ellipse([cx-r, cy-r, cx+r, cy+r], fill=col)
else:
n = 4 if si == 1 else 3
a0 = np.deg2rad(ang) + (np.pi/4 if si == 1 else -np.pi/2)
pts = [(cx+r*np.cos(a0+2*np.pi*k/n), cy+r*np.sin(a0+2*np.pi*k/n)) for k in range(n)]
d.polygon(pts, fill=col)
return np.asarray(img, np.float32)/255.0, si, ci
def make_dataset(n, seed, S=32):
rng = np.random.default_rng(seed)
X = np.zeros((n, S, S, 3), np.float32); ys = np.zeros(n, int); yc = np.zeros(n, int)
for i in range(n): X[i], ys[i], yc[i] = _one(rng, S)
return X, ys, yc
Xtr, ys_tr, yc_tr = make_dataset(3000, seed=0)
Xte, ys_te, yc_te = make_dataset(600, seed=99)
Xtr_t = torch.tensor(Xtr).permute(0, 3, 1, 2); Xte_t = torch.tensor(Xte).permute(0, 3, 1, 2)
print('train', Xtr.shape, ' test', Xte.shape)