git-svn-id: http://moon:8086/svn/matlab/trunk@91 801c6759-fa7c-4059-a304-17956f83a07c
79 lines
1.7 KiB
Matlab
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 |