[RBM]
- added git-svn-id: http://moon:8086/svn/matlab/trunk@93 801c6759-fa7c-4059-a304-17956f83a07c
This commit is contained in:
@@ -0,0 +1,64 @@
|
||||
function mnist_gen_pca(retained_variance_target)
|
||||
|
||||
close all;
|
||||
|
||||
addpath(genpath('UFLDL/common'))
|
||||
fprintf('Load mnist raw data\n');
|
||||
|
||||
x = loadMNISTImages('UFLDL/common/train-images-idx3-ubyte');
|
||||
|
||||
rbmWrite(x, 'mnist.trainingStates.dat')
|
||||
figure('name','Raw images');
|
||||
randsel = randi(size(x,2),64,1); % A random selection of samples for visualization
|
||||
display_network(x(:,randsel));
|
||||
|
||||
fprintf('Zero mean raw data\n');
|
||||
avg = mean(x, 1); % Compute the mean pixel intensity value separately for each patch.
|
||||
x = x - repmat(avg, size(x, 1), 1);
|
||||
|
||||
fprintf('Do the PCA whitening\n');
|
||||
[xHat, k, xZCAWhite] = pca(x, retained_variance_target);
|
||||
|
||||
figure('name',['PCA processed images ',sprintf('(%d / %d dimensions)', k, size(x, 1)),'']);
|
||||
display_network(xHat(:,randsel));
|
||||
|
||||
fprintf('Normalize the whitened data\n');
|
||||
minZ = min(xZCAWhite(:))
|
||||
maxZ = max(xZCAWhite(:))
|
||||
xZCAWhite_norm = (xZCAWhite - minZ)/(maxZ-minZ);
|
||||
|
||||
minZ = min(xHat(:))
|
||||
maxZ = max(xHat(:))
|
||||
xHat_norm = (xHat - minZ)/(maxZ-minZ);
|
||||
|
||||
figure('name','ZCA whitened images');
|
||||
display_network(xZCAWhite_norm(:,randsel));
|
||||
|
||||
figure('name','xHat');
|
||||
display_network(xHat_norm(:,randsel));
|
||||
|
||||
fprintf('Save data\n');
|
||||
rbmWrite(xZCAWhite_norm, 'mnist_zca.trainingStates.dat')
|
||||
rbmWrite(xHat_norm, 'mnist_xhat.trainingStates.dat')
|
||||
|
||||
|
||||
function rbmWrite(data, name)
|
||||
[numPixel, numTrain] = size(data)
|
||||
nx = sqrt(numPixel)
|
||||
ny = nx
|
||||
|
||||
data = reshape(data, nx, ny, numTrain);
|
||||
|
||||
fid = fopen(name, 'w');
|
||||
fprintf(fid, '%d\n', numTrain);
|
||||
fprintf(fid, '%d\n', numPixel);
|
||||
|
||||
for m=1:numTrain,
|
||||
d = data(:,:,m)';
|
||||
d = d(:);
|
||||
for n=1:numPixel,
|
||||
fprintf(fid, '%f\n', d(n));
|
||||
end
|
||||
end
|
||||
|
||||
fclose(fid);
|
||||
Reference in New Issue
Block a user