Files
matlab/RBM/jbatch.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

77 lines
1.9 KiB
Matlab

function jbatch(numSamplesPerBatch, numBatches)
numLabels = 10;
maxNumSamples = -1;
for d=1:numLabels,
Dall{d} = load(['test' num2str(d-1) '.ascii'],'-ascii');
fprintf('%5d Digits of class %d\n',size(Dall{d},1),d-1);
maxNumSamples = max(maxNumSamples, size(Dall{d},1));
% jconv(Dall{d}/256, d-1);
% jimage(Dall{d}/256, d-1, 28, 28);
end;
already_used = zeros(numLabels, maxNumSamples);
for k=1:numBatches,
s = ['mnist_batch_' num2str(k) '.trainingStates.dat']
fid = fopen(s, 'w');
fprintf(fid, '%d\n', numSamplesPerBatch);
for m=1:numSamplesPerBatch,
label = floor(numLabels*rand()) + 1;
D = Dall{label};
[numTrain, numPixel] = size(D);
while(1)
sample = floor(numTrain*rand()) + 1;
if already_used(label, sample) == 0
break;
end
end
already_used(label, sample) = 1;
jwrite(fid, D(sample, :)/256);
end
fclose(fid);
end
for k=1:numBatches,
s = ['mnist_retina_batch_' num2str(k) '.trainingStates.dat']
fid = fopen(s, 'w');
fprintf(fid, '%d\n', numSamplesPerBatch);
for m=1:numSamplesPerBatch,
label = floor(numLabels*rand()) + 1;
D = Dall{label};
[numTrain, numPixel] = size(D);
while(1)
sample = floor(numTrain*rand()) + 1;
if already_used(label, sample) == 0
break;
end
end
already_used(label, sample) = 1;
jwrite_retina(fid, D(sample, :)/256, 28, 28);
end
fclose(fid);
end
s = ['mnist_batch_0_9.trainingStates.dat']
fid = fopen(s, 'w');
fprintf(fid, '%d\n', numLabels*numSamplesPerBatch);
already_used = zeros(10, maxNumSamples);
for label=1:numLabels,
for m=1:numSamplesPerBatch,
D = Dall{label};
[numTrain, numPixel] = size(D);
while(1)
sample = floor(numTrain*rand()) + 1;
if already_used(label, sample) == 0
break;
end
end
already_used(label, sample) = 1;
jwrite(fid, D(sample, :)/256);
end
end
fclose(fid);