Files
matlab/RBM/rbm_eval.m
T
jens 7b34529b24 imported RBM
git-svn-id: http://moon:8086/svn/matlab/trunk@91 801c6759-fa7c-4059-a304-17956f83a07c
2016-07-12 11:24:12 +00:00

79 lines
1.7 KiB
Matlab

function rbm_eval(nh, batch, numEpochs)
[batchsize nv] = size(batch);
numGibbs = 1;
muW = 0.1;
muBV = 0.1;
muBH = 0.1;
momentum = 0.5;
momentum_final = 0.9;
weightDecay = 0.0002;
w = 0.1*randn(nv, nh);
bv = zeros(1, nv);
bh = zeros(1, nh);
v = zeros(batchsize, nv);
h = zeros(batchsize, nh);
diffW = zeros(nv, nh);
diffBV = zeros(1, nv);
diffBH = zeros(1, nh);
diffErr = zeros(nv, batchsize);
dW = zeros(nv, nh);
dBV = zeros(1, nv);
dBH = zeros(1, nh);
for k=1:numEpochs
errsum = 0;
% Positive phase
v = batch;
h = 1./(1 + exp(-(v * w + repmat(bh, batchsize, 1))));
% Update difference
diffW = v' * h;
diffBV = sum(v);
diffBH = sum(h);
diffErr = v;
% Negative phase
for n=1:numGibbs
% h is sampled
h = h > rand(batchsize, nh);
% Reconstruct visible
v = 1./(1 + exp(-(h * w' + repmat(bv, batchsize, 1))));
% Get new h probabilities from reconstruction
h = 1./(1 + exp(-(v * w + repmat(bh, batchsize, 1))));
end
% Update difference
diffW = diffW - v' * h;
diffBV = diffBV - sum(v);
diffBH = diffBH - sum(h);
diffErr = diffErr - v;
err(k) = sum(sum((diffErr).^2));
errsum = errsum + err(k);
if k > fix(numEpochs/20),
momentum = momentum_final;
end;
% Update parameter gradients
dW = momentum*dW + muW*(diffW/batchsize - weightDecay*w);
dBV = momentum*dBV + muBV*diffBV/batchsize;
dBH = momentum*dBH + muBH*diffBH/batchsize;
% Update parameters
w = w + dW;
bv = bv + dBV;
bh = bh + dBH;
end
close all;
plot(1:numEpochs, err); grid;
v = v
errsum = errsum