mirror of
https://github.com/wassname/mapping_mnist.git
synced 2026-09-09 11:26:35 +08:00
455 KiB
455 KiB
In [67]:
from __future__ import print_function
import argparse
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torchvision import datasets, transforms
from torch.autograd import Variable
from tqdm import tqdm_notebook as tqdm
%pylab inline
import pickle
import itertools
import datetimePopulating the interactive namespace from numpy and matplotlib
/home/isisilon/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/IPython/core/magics/pylab.py:160: UserWarning: pylab import has clobbered these variables: ['seed'] `%matplotlib` prevents importing * from pylab and numpy "\n`%matplotlib` prevents importing * from pylab and numpy"
In [68]:
import matplotlib.pyplot as plt
from tqdm import tqdm_notebook as tqdm
from mpl_toolkits.mplot3d import Axes3D
from matplotlib.colors import LogNorm
from matplotlib import animation
from IPython.display import HTML
from collections import defaultdict
from itertools import zip_longest
from functools import partialIn [69]:
cuda=False
seed=0
batch_size=512
ts = datetime.datetime.utcnow().strftime('%Y%m%d_%H-%M-%S')
model_path = 'data/model_%s.pickle' % ts
points_file = 'data/points_%s.pickle' % ts
torch.manual_seed(seed)
if cuda:
torch.cuda.manual_seed(seed)In [19]:
train_loader = torch.utils.data.DataLoader(
datasets.MNIST('../data', train=True, download=True,
transform=transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.1307,), (0.3081,))
])),
batch_size=batch_size, shuffle=True)
# test_loader = torch.utils.data.DataLoader(
# datasets.MNIST('../data', train=False, transform=transforms.Compose([
# transforms.ToTensor(),
# transforms.Normalize((0.1307,), (0.3081,))
# ])),
# batch_size=args.test_batch_size, shuffle=True, **kwargs)In [426]:
# https://github.com/pytorch/examples/blob/master/mnist/main.py
class Net(nn.Module):
def __init__(self):
super(Net, self).__init__()
self.conv1 = nn.Conv2d(1, 2, kernel_size=28, stride=(28,28))
# self.conv2 = nn.Conv2d(10, 2, kernel_size=5)
# self.fc1 = nn.Linear(320, 2)
self.fc2 = nn.Linear(2, 1)
def forward(self, x):
x = F.tanh(self.conv1(x))
# x = F.relu(F.max_pool2d(self.conv2(x), 2))
x = x.view(-1, 2)
# x = F.tanh(self.fc1(x))
# x = F.dropout(x, training=self.training)
x = self.fc2(x)
return F.sigmoid(x)
model = Net()
optimizer = optim.SGD(model.parameters(), lr=1e-3)
In [204]:
# nn.Conv2d?In [205]:
# I got from 0.7 to 0.4 in 100 epochs
epochs=200
for epoch in range(1, epochs + 1):
model.train()
for batch_idx, (data, target) in enumerate(train_loader):
if cuda:
data, target = data.cuda(), target.cuda()
data, target = Variable(data), Variable(target)
# reduce this to a binary problem
target = (target>5).type(torch.FloatTensor)
optimizer.zero_grad()
output = model(data)
loss = F.binary_cross_entropy(output, target)
loss.backward()
optimizer.step()
print('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format(
epoch, batch_idx * len(data), len(train_loader.dataset),
100. * batch_idx / len(train_loader), loss.data[0]))/home/isisilon/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/torch/nn/functional.py:767: UserWarning: Using a target size (torch.Size([512])) that is different to the input size (torch.Size([512, 1])) is deprecated. Please ensure they have the same size. "Please ensure they have the same size.".format(target.size(), input.size())) /home/isisilon/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/torch/nn/functional.py:767: UserWarning: Using a target size (torch.Size([96])) that is different to the input size (torch.Size([96, 1])) is deprecated. Please ensure they have the same size. "Please ensure they have the same size.".format(target.size(), input.size()))
Train Epoch: 1 [11232/60000 (99%)] Loss: 0.636986 Train Epoch: 2 [11232/60000 (99%)] Loss: 0.622902 Train Epoch: 3 [11232/60000 (99%)] Loss: 0.602751 Train Epoch: 4 [11232/60000 (99%)] Loss: 0.579772 Train Epoch: 5 [11232/60000 (99%)] Loss: 0.581861
[0;31m---------------------------------------------------------------------------[0m [0;31mKeyboardInterrupt[0m Traceback (most recent call last) [0;32m<ipython-input-205-68b33f9dff4c>[0m in [0;36m<module>[0;34m()[0m [1;32m 3[0m [0;32mfor[0m [0mepoch[0m [0;32min[0m [0mrange[0m[0;34m([0m[0;36m1[0m[0;34m,[0m [0mepochs[0m [0;34m+[0m [0;36m1[0m[0;34m)[0m[0;34m:[0m[0;34m[0m[0m [1;32m 4[0m [0mmodel[0m[0;34m.[0m[0mtrain[0m[0;34m([0m[0;34m)[0m[0;34m[0m[0m [0;32m----> 5[0;31m [0;32mfor[0m [0mbatch_idx[0m[0;34m,[0m [0;34m([0m[0mdata[0m[0;34m,[0m [0mtarget[0m[0;34m)[0m [0;32min[0m [0menumerate[0m[0;34m([0m[0mtrain_loader[0m[0;34m)[0m[0;34m:[0m[0;34m[0m[0m [0m[1;32m 6[0m [0;32mif[0m [0mcuda[0m[0;34m:[0m[0;34m[0m[0m [1;32m 7[0m [0mdata[0m[0;34m,[0m [0mtarget[0m [0;34m=[0m [0mdata[0m[0;34m.[0m[0mcuda[0m[0;34m([0m[0;34m)[0m[0;34m,[0m [0mtarget[0m[0;34m.[0m[0mcuda[0m[0;34m([0m[0;34m)[0m[0;34m[0m[0m [0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/torch/utils/data/dataloader.py[0m in [0;36m__next__[0;34m(self)[0m [1;32m 177[0m [0;32mif[0m [0mself[0m[0;34m.[0m[0mnum_workers[0m [0;34m==[0m [0;36m0[0m[0;34m:[0m [0;31m# same-process loading[0m[0;34m[0m[0m [1;32m 178[0m [0mindices[0m [0;34m=[0m [0mnext[0m[0;34m([0m[0mself[0m[0;34m.[0m[0msample_iter[0m[0;34m)[0m [0;31m# may raise StopIteration[0m[0;34m[0m[0m [0;32m--> 179[0;31m [0mbatch[0m [0;34m=[0m [0mself[0m[0;34m.[0m[0mcollate_fn[0m[0;34m([0m[0;34m[[0m[0mself[0m[0;34m.[0m[0mdataset[0m[0;34m[[0m[0mi[0m[0;34m][0m [0;32mfor[0m [0mi[0m [0;32min[0m [0mindices[0m[0;34m][0m[0;34m)[0m[0;34m[0m[0m [0m[1;32m 180[0m [0;32mif[0m [0mself[0m[0;34m.[0m[0mpin_memory[0m[0;34m:[0m[0;34m[0m[0m [1;32m 181[0m [0mbatch[0m [0;34m=[0m [0mpin_memory_batch[0m[0;34m([0m[0mbatch[0m[0;34m)[0m[0;34m[0m[0m [0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/torch/utils/data/dataloader.py[0m in [0;36m<listcomp>[0;34m(.0)[0m [1;32m 177[0m [0;32mif[0m [0mself[0m[0;34m.[0m[0mnum_workers[0m [0;34m==[0m [0;36m0[0m[0;34m:[0m [0;31m# same-process loading[0m[0;34m[0m[0m [1;32m 178[0m [0mindices[0m [0;34m=[0m [0mnext[0m[0;34m([0m[0mself[0m[0;34m.[0m[0msample_iter[0m[0;34m)[0m [0;31m# may raise StopIteration[0m[0;34m[0m[0m [0;32m--> 179[0;31m [0mbatch[0m [0;34m=[0m [0mself[0m[0;34m.[0m[0mcollate_fn[0m[0;34m([0m[0;34m[[0m[0mself[0m[0;34m.[0m[0mdataset[0m[0;34m[[0m[0mi[0m[0;34m][0m [0;32mfor[0m [0mi[0m [0;32min[0m [0mindices[0m[0;34m][0m[0;34m)[0m[0;34m[0m[0m [0m[1;32m 180[0m [0;32mif[0m [0mself[0m[0;34m.[0m[0mpin_memory[0m[0;34m:[0m[0;34m[0m[0m [1;32m 181[0m [0mbatch[0m [0;34m=[0m [0mpin_memory_batch[0m[0;34m([0m[0mbatch[0m[0;34m)[0m[0;34m[0m[0m [0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/torchvision/datasets/mnist.py[0m in [0;36m__getitem__[0;34m(self, index)[0m [1;32m 50[0m [0;31m# doing this so that it is consistent with all other datasets[0m[0;34m[0m[0;34m[0m[0m [1;32m 51[0m [0;31m# to return a PIL Image[0m[0;34m[0m[0;34m[0m[0m [0;32m---> 52[0;31m [0mimg[0m [0;34m=[0m [0mImage[0m[0;34m.[0m[0mfromarray[0m[0;34m([0m[0mimg[0m[0;34m.[0m[0mnumpy[0m[0;34m([0m[0;34m)[0m[0;34m,[0m [0mmode[0m[0;34m=[0m[0;34m'L'[0m[0;34m)[0m[0;34m[0m[0m [0m[1;32m 53[0m [0;34m[0m[0m [1;32m 54[0m [0;32mif[0m [0mself[0m[0;34m.[0m[0mtransform[0m [0;32mis[0m [0;32mnot[0m [0;32mNone[0m[0;34m:[0m[0;34m[0m[0m [0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/PIL/Image.py[0m in [0;36mfromarray[0;34m(obj, mode)[0m [1;32m 2436[0m [0mobj[0m [0;34m=[0m [0mobj[0m[0;34m.[0m[0mtostring[0m[0;34m([0m[0;34m)[0m[0;34m[0m[0m [1;32m 2437[0m [0;34m[0m[0m [0;32m-> 2438[0;31m [0;32mreturn[0m [0mfrombuffer[0m[0;34m([0m[0mmode[0m[0;34m,[0m [0msize[0m[0;34m,[0m [0mobj[0m[0;34m,[0m [0;34m"raw"[0m[0;34m,[0m [0mrawmode[0m[0;34m,[0m [0;36m0[0m[0;34m,[0m [0;36m1[0m[0;34m)[0m[0;34m[0m[0m [0m[1;32m 2439[0m [0;34m[0m[0m [1;32m 2440[0m [0;34m[0m[0m [0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/PIL/Image.py[0m in [0;36mfrombuffer[0;34m(mode, size, data, decoder_name, *args)[0m [1;32m 2386[0m [0mcore[0m[0;34m.[0m[0mmap_buffer[0m[0;34m([0m[0mdata[0m[0;34m,[0m [0msize[0m[0;34m,[0m [0mdecoder_name[0m[0;34m,[0m [0;32mNone[0m[0;34m,[0m [0;36m0[0m[0;34m,[0m [0margs[0m[0;34m)[0m[0;34m[0m[0m [1;32m 2387[0m ) [0;32m-> 2388[0;31m [0mim[0m[0;34m.[0m[0mreadonly[0m [0;34m=[0m [0;36m1[0m[0;34m[0m[0m [0m[1;32m 2389[0m [0;32mreturn[0m [0mim[0m[0;34m[0m[0m [1;32m 2390[0m [0;34m[0m[0m [0;31mKeyboardInterrupt[0m:
In [425]:
minima = model.fc2.weight.data.numpy()
minima_z = loss.data.numpy()[0]
minima, minima_zOut [425]:
(array([[ 0.22772104, 0.45285982]], dtype=float32), 5.7115297)
In [427]:
torch.save(model, model_path)/home/isisilon/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/torch/serialization.py:147: UserWarning: Couldn't retrieve source code for container of type Net. It won't be checked for correctness upon loading. "type " + obj.__name__ + ". It won't be checked "
In [428]:
model2 = torch.load(model_path)
# freeze all except last layer
mlayers = [
model2.conv1,
# model2.conv2,
# model2.fc1
]
for layer in mlayers:
for param in layer.parameters():
param.requires_grad = FalseIn [429]:
model2.fc2.weight.data.size()Out [429]:
torch.Size([1, 2])
In [ ]:
In [430]:
xm, ym = torch.Tensor(minima.T)
xm, ym = xm.numpy()[0], ym.numpy()[0]In [431]:
xs=torch.arange(xm-40,xm+0,0.3)
# xs = torch.Tensor(np.sort(np.random.normal(xm,100,100)))
ys=torch.arange(ym-40,xm+40,0.3)
# ys = torch.Tensor(np.sort(np.random.normal(ym,100,100)))
xys = list(itertools.product(xs,ys))In [434]:
xs=torch.arange(xm-20,xm+20,1)
# xs = torch.Tensor(np.sort(np.random.normal(xm,100,100)))
ys=torch.arange(ym-20,xm+20,1)
# ys = torch.Tensor(np.sort(np.random.normal(ym,100,100)))
xys = list(itertools.product(xs,ys))In [ ]:
# grid search, run this overnight
points = []
for x,y in tqdm(xys):
torch.manual_seed(seed)
if cuda:
torch.cuda.manual_seed(seed)
# set last layer weights
model2.fc2.weight.data = torch.Tensor([[y,x]])
zs=[]
dzs=[]
batches=0
for batch_idx, (data, target) in enumerate(train_loader):
if cuda:
data, target = data.cuda(), target.cuda()
data, target = Variable(data), Variable(target)
# reduce this to a binary problem
target = (target>5).type(torch.FloatTensor)
optimizer.zero_grad()
output = model2(data)
loss = F.binary_cross_entropy(output, target)
loss.backward()
optimizer.step()
dzs.append(model2.fc2.weight.grad.data.numpy())
zs.append(loss.data.numpy()[0])
batches += 1
if batches>6:
break
output = model2(data)
loss = F.binary_cross_entropy(output, target)
z = np.mean(zs)
z_var = np.std(zs)
dz = np.mean(dzs,0)
dz_var = np.std(dzs,0)
points.append([x,y,z, z_var, dz, dz_var, batches])[Data output - unsupported data type map[string]interface {} for mime type application/vnd.jupyter.widget-view+json]
/home/isisilon/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/torch/nn/functional.py:767: UserWarning: Using a target size (torch.Size([512])) that is different to the input size (torch.Size([512, 1])) is deprecated. Please ensure they have the same size. "Please ensure they have the same size.".format(target.size(), input.size()))
In [415]:
x = np.array([p[0] for p in points])
y = np.array([p[1] for p in points])
z = np.array([p[2] for p in points])
zv = np.array([p[3] for p in points])
dz = np.array([p[4] for p in points])[:,0,:]
dzv = np.array([p[5] for p in points])[:,0,:]In [416]:
plt.figure(figsize=(10,8))
# x,y,z,zv,dz,dzv, n = np.array(points).T
plt.quiver(x, y, dz[:,0], dz[:,1], angles='xy', scale_units='xy', scale=1)
plt.show()In [417]:
plt.figure(figsize=(10,8))
# x,y,z,zv,dz,dzv, n = np.array(points).T
plt.quiver(x, y, dzv[:,0], dzv[:,1], angles='xy', scale_units='xy', scale=1)
plt.show()In [418]:
plt.figure(figsize=(10,8))
# x,y,z,zv,dz,dzv, n = np.array(points).T
plt.scatter(x,y,c=dz[:,0])
plt.show()In [419]:
plt.figure(figsize=(10,8))
plt.scatter(x,y,c=z)
plt.title('loss')
plt.colorbar()Out [419]:
<matplotlib.colorbar.Colorbar at 0x7f0ad88e1898>
In [420]:
plt.figure(figsize=(10,8))
plt.scatter(x,y,c=z)
plt.title('log loss')
plt.colorbar()Out [420]:
<matplotlib.colorbar.Colorbar at 0x7f0adad59518>
In [ ]:
In [421]:
plt.figure(figsize=(10,8))
plt.title('log std')
plt.scatter(x,y,c=np.log(v))
plt.colorbar()[0;31m---------------------------------------------------------------------------[0m
[0;31mKeyError[0m Traceback (most recent call last)
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/colors.py[0m in [0;36mto_rgba[0;34m(c, alpha)[0m
[1;32m 140[0m [0;32mtry[0m[0;34m:[0m[0;34m[0m[0m
[0;32m--> 141[0;31m [0mrgba[0m [0;34m=[0m [0m_colors_full_map[0m[0;34m.[0m[0mcache[0m[0;34m[[0m[0mc[0m[0;34m,[0m [0malpha[0m[0;34m][0m[0;34m[0m[0m
[0m[1;32m 142[0m [0;32mexcept[0m [0;34m([0m[0mKeyError[0m[0;34m,[0m [0mTypeError[0m[0;34m)[0m[0;34m:[0m [0;31m# Not in cache, or unhashable.[0m[0;34m[0m[0m
[0;31mKeyError[0m: (-0.93098905452915248, None)
During handling of the above exception, another exception occurred:
[0;31mValueError[0m Traceback (most recent call last)
[0;32m<ipython-input-421-03316127cc9a>[0m in [0;36m<module>[0;34m()[0m
[1;32m 1[0m [0mplt[0m[0;34m.[0m[0mfigure[0m[0;34m([0m[0mfigsize[0m[0;34m=[0m[0;34m([0m[0;36m10[0m[0;34m,[0m[0;36m8[0m[0;34m)[0m[0;34m)[0m[0;34m[0m[0m
[1;32m 2[0m [0mplt[0m[0;34m.[0m[0mtitle[0m[0;34m([0m[0;34m'log std'[0m[0;34m)[0m[0;34m[0m[0m
[0;32m----> 3[0;31m [0mplt[0m[0;34m.[0m[0mscatter[0m[0;34m([0m[0mx[0m[0;34m,[0m[0my[0m[0;34m,[0m[0mc[0m[0;34m=[0m[0mnp[0m[0;34m.[0m[0mlog[0m[0;34m([0m[0mv[0m[0;34m)[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 4[0m [0mplt[0m[0;34m.[0m[0mcolorbar[0m[0;34m([0m[0;34m)[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/pyplot.py[0m in [0;36mscatter[0;34m(x, y, s, c, marker, cmap, norm, vmin, vmax, alpha, linewidths, verts, edgecolors, hold, data, **kwargs)[0m
[1;32m 3432[0m [0mvmin[0m[0;34m=[0m[0mvmin[0m[0;34m,[0m [0mvmax[0m[0;34m=[0m[0mvmax[0m[0;34m,[0m [0malpha[0m[0;34m=[0m[0malpha[0m[0;34m,[0m[0;34m[0m[0m
[1;32m 3433[0m [0mlinewidths[0m[0;34m=[0m[0mlinewidths[0m[0;34m,[0m [0mverts[0m[0;34m=[0m[0mverts[0m[0;34m,[0m[0;34m[0m[0m
[0;32m-> 3434[0;31m edgecolors=edgecolors, data=data, **kwargs)
[0m[1;32m 3435[0m [0;32mfinally[0m[0;34m:[0m[0;34m[0m[0m
[1;32m 3436[0m [0max[0m[0;34m.[0m[0m_hold[0m [0;34m=[0m [0mwashold[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/__init__.py[0m in [0;36minner[0;34m(ax, *args, **kwargs)[0m
[1;32m 1896[0m warnings.warn(msg % (label_namer, func.__name__),
[1;32m 1897[0m RuntimeWarning, stacklevel=2)
[0;32m-> 1898[0;31m [0;32mreturn[0m [0mfunc[0m[0;34m([0m[0max[0m[0;34m,[0m [0;34m*[0m[0margs[0m[0;34m,[0m [0;34m**[0m[0mkwargs[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 1899[0m [0mpre_doc[0m [0;34m=[0m [0minner[0m[0;34m.[0m[0m__doc__[0m[0;34m[0m[0m
[1;32m 1900[0m [0;32mif[0m [0mpre_doc[0m [0;32mis[0m [0;32mNone[0m[0;34m:[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/axes/_axes.py[0m in [0;36mscatter[0;34m(self, x, y, s, c, marker, cmap, norm, vmin, vmax, alpha, linewidths, verts, edgecolors, **kwargs)[0m
[1;32m 4032[0m [0moffsets[0m[0;34m=[0m[0moffsets[0m[0;34m,[0m[0;34m[0m[0m
[1;32m 4033[0m [0mtransOffset[0m[0;34m=[0m[0mkwargs[0m[0;34m.[0m[0mpop[0m[0;34m([0m[0;34m'transform'[0m[0;34m,[0m [0mself[0m[0;34m.[0m[0mtransData[0m[0;34m)[0m[0;34m,[0m[0;34m[0m[0m
[0;32m-> 4034[0;31m [0malpha[0m[0;34m=[0m[0malpha[0m[0;34m[0m[0m
[0m[1;32m 4035[0m )
[1;32m 4036[0m [0mcollection[0m[0;34m.[0m[0mset_transform[0m[0;34m([0m[0mmtransforms[0m[0;34m.[0m[0mIdentityTransform[0m[0;34m([0m[0;34m)[0m[0;34m)[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/collections.py[0m in [0;36m__init__[0;34m(self, paths, sizes, **kwargs)[0m
[1;32m 900[0m """
[1;32m 901[0m [0;34m[0m[0m
[0;32m--> 902[0;31m [0mCollection[0m[0;34m.[0m[0m__init__[0m[0;34m([0m[0mself[0m[0;34m,[0m [0;34m**[0m[0mkwargs[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 903[0m [0mself[0m[0;34m.[0m[0mset_paths[0m[0;34m([0m[0mpaths[0m[0;34m)[0m[0;34m[0m[0m
[1;32m 904[0m [0mself[0m[0;34m.[0m[0mset_sizes[0m[0;34m([0m[0msizes[0m[0;34m)[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/collections.py[0m in [0;36m__init__[0;34m(self, edgecolors, facecolors, linewidths, linestyles, antialiaseds, offsets, transOffset, norm, cmap, pickradius, hatch, urls, offset_position, zorder, **kwargs)[0m
[1;32m 138[0m [0;34m[0m[0m
[1;32m 139[0m [0mself[0m[0;34m.[0m[0m_hatch_color[0m [0;34m=[0m [0mmcolors[0m[0;34m.[0m[0mto_rgba[0m[0;34m([0m[0mmpl[0m[0;34m.[0m[0mrcParams[0m[0;34m[[0m[0;34m'hatch.color'[0m[0;34m][0m[0;34m)[0m[0;34m[0m[0m
[0;32m--> 140[0;31m [0mself[0m[0;34m.[0m[0mset_facecolor[0m[0;34m([0m[0mfacecolors[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 141[0m [0mself[0m[0;34m.[0m[0mset_edgecolor[0m[0;34m([0m[0medgecolors[0m[0;34m)[0m[0;34m[0m[0m
[1;32m 142[0m [0mself[0m[0;34m.[0m[0mset_linewidth[0m[0;34m([0m[0mlinewidths[0m[0;34m)[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/collections.py[0m in [0;36mset_facecolor[0;34m(self, c)[0m
[1;32m 679[0m """
[1;32m 680[0m [0mself[0m[0;34m.[0m[0m_original_facecolor[0m [0;34m=[0m [0mc[0m[0;34m[0m[0m
[0;32m--> 681[0;31m [0mself[0m[0;34m.[0m[0m_set_facecolor[0m[0;34m([0m[0mc[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 682[0m [0;34m[0m[0m
[1;32m 683[0m [0;32mdef[0m [0mset_facecolors[0m[0;34m([0m[0mself[0m[0;34m,[0m [0mc[0m[0;34m)[0m[0;34m:[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/collections.py[0m in [0;36m_set_facecolor[0;34m(self, c)[0m
[1;32m 664[0m [0;32mexcept[0m [0mAttributeError[0m[0;34m:[0m[0;34m[0m[0m
[1;32m 665[0m [0;32mpass[0m[0;34m[0m[0m
[0;32m--> 666[0;31m [0mself[0m[0;34m.[0m[0m_facecolors[0m [0;34m=[0m [0mmcolors[0m[0;34m.[0m[0mto_rgba_array[0m[0;34m([0m[0mc[0m[0;34m,[0m [0mself[0m[0;34m.[0m[0m_alpha[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 667[0m [0mself[0m[0;34m.[0m[0mstale[0m [0;34m=[0m [0;32mTrue[0m[0;34m[0m[0m
[1;32m 668[0m [0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/colors.py[0m in [0;36mto_rgba_array[0;34m(c, alpha)[0m
[1;32m 237[0m [0mresult[0m [0;34m=[0m [0mnp[0m[0;34m.[0m[0mempty[0m[0;34m([0m[0;34m([0m[0mlen[0m[0;34m([0m[0mc[0m[0;34m)[0m[0;34m,[0m [0;36m4[0m[0;34m)[0m[0;34m,[0m [0mfloat[0m[0;34m)[0m[0;34m[0m[0m
[1;32m 238[0m [0;32mfor[0m [0mi[0m[0;34m,[0m [0mcc[0m [0;32min[0m [0menumerate[0m[0;34m([0m[0mc[0m[0;34m)[0m[0;34m:[0m[0;34m[0m[0m
[0;32m--> 239[0;31m [0mresult[0m[0;34m[[0m[0mi[0m[0;34m][0m [0;34m=[0m [0mto_rgba[0m[0;34m([0m[0mcc[0m[0;34m,[0m [0malpha[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 240[0m [0;32mreturn[0m [0mresult[0m[0;34m[0m[0m
[1;32m 241[0m [0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/colors.py[0m in [0;36mto_rgba[0;34m(c, alpha)[0m
[1;32m 141[0m [0mrgba[0m [0;34m=[0m [0m_colors_full_map[0m[0;34m.[0m[0mcache[0m[0;34m[[0m[0mc[0m[0;34m,[0m [0malpha[0m[0;34m][0m[0;34m[0m[0m
[1;32m 142[0m [0;32mexcept[0m [0;34m([0m[0mKeyError[0m[0;34m,[0m [0mTypeError[0m[0;34m)[0m[0;34m:[0m [0;31m# Not in cache, or unhashable.[0m[0;34m[0m[0m
[0;32m--> 143[0;31m [0mrgba[0m [0;34m=[0m [0m_to_rgba_no_colorcycle[0m[0;34m([0m[0mc[0m[0;34m,[0m [0malpha[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 144[0m [0;32mtry[0m[0;34m:[0m[0;34m[0m[0m
[1;32m 145[0m [0m_colors_full_map[0m[0;34m.[0m[0mcache[0m[0;34m[[0m[0mc[0m[0;34m,[0m [0malpha[0m[0;34m][0m [0;34m=[0m [0mrgba[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/colors.py[0m in [0;36m_to_rgba_no_colorcycle[0;34m(c, alpha)[0m
[1;32m 192[0m [0;31m# float)` and `np.array(...).astype(float)` all convert "0.5" to 0.5.[0m[0;34m[0m[0;34m[0m[0m
[1;32m 193[0m [0;31m# Test dimensionality to reject single floats.[0m[0;34m[0m[0;34m[0m[0m
[0;32m--> 194[0;31m [0;32mraise[0m [0mValueError[0m[0;34m([0m[0;34m"Invalid RGBA argument: {!r}"[0m[0;34m.[0m[0mformat[0m[0;34m([0m[0morig_c[0m[0;34m)[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 195[0m [0;31m# Return a tuple to prevent the cached value from being modified.[0m[0;34m[0m[0;34m[0m[0m
[1;32m 196[0m [0mc[0m [0;34m=[0m [0mtuple[0m[0;34m([0m[0mc[0m[0;34m.[0m[0mastype[0m[0;34m([0m[0mfloat[0m[0;34m)[0m[0;34m)[0m[0;34m[0m[0m
[0;31mValueError[0m: Invalid RGBA argument: -0.93098905452915248In [422]:
plt.figure(figsize=(10,8))
plt.title('std')
plt.scatter(x,y,c=v)
plt.colorbar()[0;31m---------------------------------------------------------------------------[0m
[0;31mKeyError[0m Traceback (most recent call last)
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/colors.py[0m in [0;36mto_rgba[0;34m(c, alpha)[0m
[1;32m 140[0m [0;32mtry[0m[0;34m:[0m[0;34m[0m[0m
[0;32m--> 141[0;31m [0mrgba[0m [0;34m=[0m [0m_colors_full_map[0m[0;34m.[0m[0mcache[0m[0;34m[[0m[0mc[0m[0;34m,[0m [0malpha[0m[0;34m][0m[0;34m[0m[0m
[0m[1;32m 142[0m [0;32mexcept[0m [0;34m([0m[0mKeyError[0m[0;34m,[0m [0mTypeError[0m[0;34m)[0m[0;34m:[0m [0;31m# Not in cache, or unhashable.[0m[0;34m[0m[0m
[0;31mKeyError[0m: (0.39416366815567017, None)
During handling of the above exception, another exception occurred:
[0;31mValueError[0m Traceback (most recent call last)
[0;32m<ipython-input-422-0554dbdab3a9>[0m in [0;36m<module>[0;34m()[0m
[1;32m 1[0m [0mplt[0m[0;34m.[0m[0mfigure[0m[0;34m([0m[0mfigsize[0m[0;34m=[0m[0;34m([0m[0;36m10[0m[0;34m,[0m[0;36m8[0m[0;34m)[0m[0;34m)[0m[0;34m[0m[0m
[1;32m 2[0m [0mplt[0m[0;34m.[0m[0mtitle[0m[0;34m([0m[0;34m'std'[0m[0;34m)[0m[0;34m[0m[0m
[0;32m----> 3[0;31m [0mplt[0m[0;34m.[0m[0mscatter[0m[0;34m([0m[0mx[0m[0;34m,[0m[0my[0m[0;34m,[0m[0mc[0m[0;34m=[0m[0mv[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 4[0m [0mplt[0m[0;34m.[0m[0mcolorbar[0m[0;34m([0m[0;34m)[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/pyplot.py[0m in [0;36mscatter[0;34m(x, y, s, c, marker, cmap, norm, vmin, vmax, alpha, linewidths, verts, edgecolors, hold, data, **kwargs)[0m
[1;32m 3432[0m [0mvmin[0m[0;34m=[0m[0mvmin[0m[0;34m,[0m [0mvmax[0m[0;34m=[0m[0mvmax[0m[0;34m,[0m [0malpha[0m[0;34m=[0m[0malpha[0m[0;34m,[0m[0;34m[0m[0m
[1;32m 3433[0m [0mlinewidths[0m[0;34m=[0m[0mlinewidths[0m[0;34m,[0m [0mverts[0m[0;34m=[0m[0mverts[0m[0;34m,[0m[0;34m[0m[0m
[0;32m-> 3434[0;31m edgecolors=edgecolors, data=data, **kwargs)
[0m[1;32m 3435[0m [0;32mfinally[0m[0;34m:[0m[0;34m[0m[0m
[1;32m 3436[0m [0max[0m[0;34m.[0m[0m_hold[0m [0;34m=[0m [0mwashold[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/__init__.py[0m in [0;36minner[0;34m(ax, *args, **kwargs)[0m
[1;32m 1896[0m warnings.warn(msg % (label_namer, func.__name__),
[1;32m 1897[0m RuntimeWarning, stacklevel=2)
[0;32m-> 1898[0;31m [0;32mreturn[0m [0mfunc[0m[0;34m([0m[0max[0m[0;34m,[0m [0;34m*[0m[0margs[0m[0;34m,[0m [0;34m**[0m[0mkwargs[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 1899[0m [0mpre_doc[0m [0;34m=[0m [0minner[0m[0;34m.[0m[0m__doc__[0m[0;34m[0m[0m
[1;32m 1900[0m [0;32mif[0m [0mpre_doc[0m [0;32mis[0m [0;32mNone[0m[0;34m:[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/axes/_axes.py[0m in [0;36mscatter[0;34m(self, x, y, s, c, marker, cmap, norm, vmin, vmax, alpha, linewidths, verts, edgecolors, **kwargs)[0m
[1;32m 4032[0m [0moffsets[0m[0;34m=[0m[0moffsets[0m[0;34m,[0m[0;34m[0m[0m
[1;32m 4033[0m [0mtransOffset[0m[0;34m=[0m[0mkwargs[0m[0;34m.[0m[0mpop[0m[0;34m([0m[0;34m'transform'[0m[0;34m,[0m [0mself[0m[0;34m.[0m[0mtransData[0m[0;34m)[0m[0;34m,[0m[0;34m[0m[0m
[0;32m-> 4034[0;31m [0malpha[0m[0;34m=[0m[0malpha[0m[0;34m[0m[0m
[0m[1;32m 4035[0m )
[1;32m 4036[0m [0mcollection[0m[0;34m.[0m[0mset_transform[0m[0;34m([0m[0mmtransforms[0m[0;34m.[0m[0mIdentityTransform[0m[0;34m([0m[0;34m)[0m[0;34m)[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/collections.py[0m in [0;36m__init__[0;34m(self, paths, sizes, **kwargs)[0m
[1;32m 900[0m """
[1;32m 901[0m [0;34m[0m[0m
[0;32m--> 902[0;31m [0mCollection[0m[0;34m.[0m[0m__init__[0m[0;34m([0m[0mself[0m[0;34m,[0m [0;34m**[0m[0mkwargs[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 903[0m [0mself[0m[0;34m.[0m[0mset_paths[0m[0;34m([0m[0mpaths[0m[0;34m)[0m[0;34m[0m[0m
[1;32m 904[0m [0mself[0m[0;34m.[0m[0mset_sizes[0m[0;34m([0m[0msizes[0m[0;34m)[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/collections.py[0m in [0;36m__init__[0;34m(self, edgecolors, facecolors, linewidths, linestyles, antialiaseds, offsets, transOffset, norm, cmap, pickradius, hatch, urls, offset_position, zorder, **kwargs)[0m
[1;32m 138[0m [0;34m[0m[0m
[1;32m 139[0m [0mself[0m[0;34m.[0m[0m_hatch_color[0m [0;34m=[0m [0mmcolors[0m[0;34m.[0m[0mto_rgba[0m[0;34m([0m[0mmpl[0m[0;34m.[0m[0mrcParams[0m[0;34m[[0m[0;34m'hatch.color'[0m[0;34m][0m[0;34m)[0m[0;34m[0m[0m
[0;32m--> 140[0;31m [0mself[0m[0;34m.[0m[0mset_facecolor[0m[0;34m([0m[0mfacecolors[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 141[0m [0mself[0m[0;34m.[0m[0mset_edgecolor[0m[0;34m([0m[0medgecolors[0m[0;34m)[0m[0;34m[0m[0m
[1;32m 142[0m [0mself[0m[0;34m.[0m[0mset_linewidth[0m[0;34m([0m[0mlinewidths[0m[0;34m)[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/collections.py[0m in [0;36mset_facecolor[0;34m(self, c)[0m
[1;32m 679[0m """
[1;32m 680[0m [0mself[0m[0;34m.[0m[0m_original_facecolor[0m [0;34m=[0m [0mc[0m[0;34m[0m[0m
[0;32m--> 681[0;31m [0mself[0m[0;34m.[0m[0m_set_facecolor[0m[0;34m([0m[0mc[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 682[0m [0;34m[0m[0m
[1;32m 683[0m [0;32mdef[0m [0mset_facecolors[0m[0;34m([0m[0mself[0m[0;34m,[0m [0mc[0m[0;34m)[0m[0;34m:[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/collections.py[0m in [0;36m_set_facecolor[0;34m(self, c)[0m
[1;32m 664[0m [0;32mexcept[0m [0mAttributeError[0m[0;34m:[0m[0;34m[0m[0m
[1;32m 665[0m [0;32mpass[0m[0;34m[0m[0m
[0;32m--> 666[0;31m [0mself[0m[0;34m.[0m[0m_facecolors[0m [0;34m=[0m [0mmcolors[0m[0;34m.[0m[0mto_rgba_array[0m[0;34m([0m[0mc[0m[0;34m,[0m [0mself[0m[0;34m.[0m[0m_alpha[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 667[0m [0mself[0m[0;34m.[0m[0mstale[0m [0;34m=[0m [0;32mTrue[0m[0;34m[0m[0m
[1;32m 668[0m [0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/colors.py[0m in [0;36mto_rgba_array[0;34m(c, alpha)[0m
[1;32m 237[0m [0mresult[0m [0;34m=[0m [0mnp[0m[0;34m.[0m[0mempty[0m[0;34m([0m[0;34m([0m[0mlen[0m[0;34m([0m[0mc[0m[0;34m)[0m[0;34m,[0m [0;36m4[0m[0;34m)[0m[0;34m,[0m [0mfloat[0m[0;34m)[0m[0;34m[0m[0m
[1;32m 238[0m [0;32mfor[0m [0mi[0m[0;34m,[0m [0mcc[0m [0;32min[0m [0menumerate[0m[0;34m([0m[0mc[0m[0;34m)[0m[0;34m:[0m[0;34m[0m[0m
[0;32m--> 239[0;31m [0mresult[0m[0;34m[[0m[0mi[0m[0;34m][0m [0;34m=[0m [0mto_rgba[0m[0;34m([0m[0mcc[0m[0;34m,[0m [0malpha[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 240[0m [0;32mreturn[0m [0mresult[0m[0;34m[0m[0m
[1;32m 241[0m [0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/colors.py[0m in [0;36mto_rgba[0;34m(c, alpha)[0m
[1;32m 141[0m [0mrgba[0m [0;34m=[0m [0m_colors_full_map[0m[0;34m.[0m[0mcache[0m[0;34m[[0m[0mc[0m[0;34m,[0m [0malpha[0m[0;34m][0m[0;34m[0m[0m
[1;32m 142[0m [0;32mexcept[0m [0;34m([0m[0mKeyError[0m[0;34m,[0m [0mTypeError[0m[0;34m)[0m[0;34m:[0m [0;31m# Not in cache, or unhashable.[0m[0;34m[0m[0m
[0;32m--> 143[0;31m [0mrgba[0m [0;34m=[0m [0m_to_rgba_no_colorcycle[0m[0;34m([0m[0mc[0m[0;34m,[0m [0malpha[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 144[0m [0;32mtry[0m[0;34m:[0m[0;34m[0m[0m
[1;32m 145[0m [0m_colors_full_map[0m[0;34m.[0m[0mcache[0m[0;34m[[0m[0mc[0m[0;34m,[0m [0malpha[0m[0;34m][0m [0;34m=[0m [0mrgba[0m[0;34m[0m[0m
[0;32m~/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/colors.py[0m in [0;36m_to_rgba_no_colorcycle[0;34m(c, alpha)[0m
[1;32m 192[0m [0;31m# float)` and `np.array(...).astype(float)` all convert "0.5" to 0.5.[0m[0;34m[0m[0;34m[0m[0m
[1;32m 193[0m [0;31m# Test dimensionality to reject single floats.[0m[0;34m[0m[0;34m[0m[0m
[0;32m--> 194[0;31m [0;32mraise[0m [0mValueError[0m[0;34m([0m[0;34m"Invalid RGBA argument: {!r}"[0m[0;34m.[0m[0mformat[0m[0;34m([0m[0morig_c[0m[0;34m)[0m[0;34m)[0m[0;34m[0m[0m
[0m[1;32m 195[0m [0;31m# Return a tuple to prevent the cached value from being modified.[0m[0;34m[0m[0;34m[0m[0m
[1;32m 196[0m [0mc[0m [0;34m=[0m [0mtuple[0m[0;34m([0m[0mc[0m[0;34m.[0m[0mastype[0m[0;34m([0m[0mfloat[0m[0;34m)[0m[0;34m)[0m[0;34m[0m[0m
[0;31mValueError[0m: Invalid RGBA argument: 0.39416366815567017In [348]:
# save
pickle.dump('points267x266', open(points_file,'wb'))In [328]:
# save
pickle.dump(points, open(points_file,'wb'))In [78]:
# save
points = pickle.load(open(points_file,'rb'))In [346]:
x,y,z,v,n = np.array(points).T
# scale lossses to they look OK
z=(z-z.min())*10000000
logzmax = np.log(z.max())
logzmaxOut [346]:
19.090419312520552
In [349]:
# now reshape into square arrays
x = x.reshape((len(xs),len(ys)))
y = y.reshape((len(xs),len(ys)))
z = z.reshape((len(xs),len(ys)))
# dz = dz.reshape((len(xs),len(ys)))
v = v.reshape((len(xs),len(ys)))
# dzv = dzv.reshape((len(xs),len(ys)))
z.shapeOut [349]:
(267, 266)
In [345]:
ax = plt.gca()
cm=ax.contour(x, y, z, levels=np.logspace(0, logzmax//2, 55), norm=LogNorm(), cmap=plt.cm.jet, alpha=0.15)
plt.colorbar(cm)
# ax.plot(*minima_, 'r*', markersize=10)
# ax.plot(*problem.x0, 'r+', markersize=10)
plt.title('minst: a slice of the problem surface for 2 neurons: loss')
ax.set_xlabel('$x$')
ax.set_ylabel('$y$')
plt.show()/home/isisilon/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/contour.py:1518: UserWarning: Log scale: values of z <= 0 have been masked
warnings.warn('Log scale: values of z <= 0 have been masked')
In [351]:
# scale lossses to they look OK
v=(v-v.min())*10000000
logvmax = np.log(v.max())
logvmaxOut [351]:
15.494302237578435
In [352]:
ax = plt.gca()
cm=ax.contour(x, y, v, levels=np.logspace(0, logvmax//2, 55), norm=LogNorm(), cmap=plt.cm.jet, alpha=0.15)
plt.colorbar(cm)
# ax.plot(*minima_, 'r*', markersize=10)
# ax.plot(*problem.x0, 'r+', markersize=10)
plt.title('minst: a slice of the problem surface for 2 neurons - variance')
ax.set_xlabel('$x$')
ax.set_ylabel('$y$')
plt.show()/home/isisilon/.pyenv/versions/3.6.0/envs/jupyter3/lib/python3.6/site-packages/matplotlib/contour.py:1518: UserWarning: Log scale: values of z <= 0 have been masked
warnings.warn('Log scale: values of z <= 0 have been masked')
In [ ]:
model3 = torch.load(model_path)
# freeze all except last layer
mlayers = [
model3.conv1,
model3.conv2,
model3.fc1
]
for layer in mlayers:
for param in layer.parameters():
param.requires_grad = False
optimizer = optim.SGD(model3.parameters(), lr=1e-3)In [ ]:
epochs=100
points2=[]
for epoch in range(1, epochs + 1):
model.train()
for batch_idx, (data, target) in enumerate(train_loader):
if cuda:
data, target = data.cuda(), target.cuda()
data, target = Variable(data), Variable(target)
# reduce this to a binary problem
target = (target>5).type(torch.FloatTensor)
optimizer.zero_grad()
output = model3(data)
loss = F.binary_cross_entropy(output, target)
loss.backward()
optimizer.step()
x,y=model3.conv2.weights.data
z=loss.data
points2.append([x,y,z])
print('Train Epoch: {} [{}/{} ({:.0f}%)]\tLoss: {:.6f}'.format(
epoch, batch_idx * len(data), len(train_loader.dataset),
100. * batch_idx / len(train_loader), loss.data[0]))