git-svn-id: http://moon:8086/svn/matlab/trunk@91 801c6759-fa7c-4059-a304-17956f83a07c
76 lines
1.6 KiB
Plaintext
76 lines
1.6 KiB
Plaintext
function rbm_eval(nh, batch, numEpochs)
|
|
|
|
[nv batchsize] = size(batch);
|
|
|
|
numGibbs = 1;
|
|
muW = 0.1;
|
|
muBV = 0.1;
|
|
muBH = 0.1;
|
|
momentum = 0.5;
|
|
momentum_final = 0.9;
|
|
weightDecay = 0.000;
|
|
|
|
w = 0.1*randn(nv, nh);
|
|
bv = zeros(nv, 1);
|
|
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(nv, 1);
|
|
dBH = zeros(nh, 1);
|
|
|
|
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, 2);
|
|
diffBH = sum(h')';
|
|
diffErr = v;
|
|
% Negative phase
|
|
for n=1:numGibbs
|
|
% h is sampled
|
|
h = h > rand(batchsize, nh);
|
|
|
|
% reconstruct
|
|
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, 2);
|
|
diffBH = diffBH - sum(h');
|
|
diffErr = diffErr - v;
|
|
|
|
err = sum(sum((diffErr).^2));
|
|
errsum = errsum + err;
|
|
|
|
if k > 5,
|
|
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
|
|
v = v
|
|
err = err
|
|
errsum = errsum |