%% setup environment
addpath('/home/cicek/caffe/matlab');

try 
  caffe.set_mode_gpu();
  caffe.set_device(0);
catch e
  caffe.set_mode_cpu();
end

fid            = fopen('/where_your_train_prototxt_is/3dUnet_miccai2016_no_BN.prototxt');
train_prototxt = fread(fid);
fclose(fid);

% data: input data
data = ncread('test_blob.h5','data');

% dim: dimensions of data, input, and output
dim.bottom = [9 9 7];
dim.data = [248 244 64 3];
%3-3-3 / 3-3-3
scalefun   = @(x) deal([(((((((x(1)+4)*2)+4)*2)+4)*2)+4) (((((((x(2)+4)*2)+4)*2)+4)*2)+4) (((((((x(3)+4)*2)+4)*2)+4)*2)+4)], ...
                       [((((((x(1)*2)-4)*2)-4)*2)-4) ((((((x(2)*2)-4)*2)-4)*2)-4) ((((((x(3)*2)-4)*2)-4)*2)-4)]);
[dim.input dim.output] = scalefun(dim.bottom);
dim.tiles = ceil(dim.data(1:3)./dim.output);
dim.min_offset = -[(dim.input(1)-dim.output(1))/2 ...
               (dim.input(2)-dim.output(2))/2 ...
               (dim.input(3)-dim.output(3))/2];
dim.max_offset = dim.data(1:3)+dim.min_offset-dim.output;
dim.ratio = dim.input./dim.output;

% test file
test_file = strrep('/where_your_train_prototxt_is/3dUnet_miccai2016_no_BN.prototxt','-train.prototxt',sprintf('-test.prototxt'));
fid       = fopen(test_file,'w');
fprintf(fid, 'input: "data"\n');
fprintf(fid, 'input_shape: { dim: %g dim: %g dim: %g dim: %g dim: %g}\n', ...
       1, dim.data(4), dim.data(3), dim.data(2), dim.data(1));
fprintf(fid, 'input: "def"\n');
fprintf(fid, 'input_shape: { dim: %g dim: %g dim: %g dim: %g dim: %g}\n', ...
        1, dim.input(3), dim.input(2), dim.input(1), 3);
fprintf(fid, 'state: { phase: TEST }\n'); 
fwrite(fid, train_prototxt);
fclose(fid);

for p=1:1,% subject
  % net: network 
  d = dir(fullfile('kidney_seg_no_BN',sprintf('kidney_seg_no_BN_snapshot_iter_*.caffemodel')));
  for dd = 1:numel(d),
    %[~, nIdx] = max([d.datenum]);
    lastunderscore = strfind(d(dd).name,'_');
    lastunderscore = lastunderscore(end);
    iter = str2num(d(dd).name(lastunderscore+1:end-11));
    predfile = dir(sprintf(fullfile('kidney_seg_no_BN','pred1%02d_iter_%d.h5'),p,iter));
    if ~isempty(predfile) && (predfile.datenum>d(dd).datenum), continue; end
    fprintf(sprintf('Computing segmentation for subject %02d iteration %d ..',p,iter));
    snapshotd='kidney_seg_no_BN';
    cmd = ['net = caffe.Net(''' test_file ''',' ...
                                'fullfile(snapshotd,d(dd).name),''test'');'];
    evalc(cmd);


    % pred: prediction
    pred.TRAIN = NaN([dim.tiles.*dim.output 3 1],'single');

    for i=0:dim.tiles(1)-1,       % width
      for j=0:dim.tiles(2)-1,     % height
        for k=0:dim.tiles(3)-1,   % level
          [X Y Z ] = ndgrid((0:dim.input(1)-1)+dim.min_offset(1)+i*dim.output(1),  ...
                            (0:dim.input(2)-1)+dim.min_offset(2)+j*dim.output(2),  ...
                            (0:dim.input(3)-1)+dim.min_offset(3)+k*dim.output(3));
          def = single(permute( cat(4,Z,Y,X), [4 1 2 3]));

          % forward pass
          out = net.forward( {data(:,:,:,:,p); def});
          % size(out{1})
          % fill matrix of predictions
          pred.TRAIN((1:dim.output(1))+i*dim.output(1), ...
                     (1:dim.output(2))+j*dim.output(2), ...
                     (1:dim.output(3))+k*dim.output(3),:,1)  = out{1};
        end
      end
    end
    % crop to original image size
    pred.TRAIN = pred.TRAIN(1:dim.data(1),1:dim.data(2),1:dim.data(3),:,:);

    % compute hard predictions
    [~, pred.TRAIN] = max(permute(reshape(softmax(reshape(pred.TRAIN,[prod(dim.data(1:3)),3])')',[dim.data(1:3) 3]),[1 2 4 3]),[],3);
    pred.TRAIN = cast(pred.TRAIN,'uint8');
    
    % save predictions to disc
    h5outputfile = sprintf(fullfile('kidney_seg_no_BN','pred1%02d_iter_%d.h5'),p,iter);
    h5create(h5outputfile,'/decision',size(pred.TRAIN),'Datatype',class(pred.TRAIN));
    h5write(h5outputfile,'/decision',pred.TRAIN);
    h5writeatt(h5outputfile,'/decision','element_size_um', 2 .* [1.0217638 0.87890553 0.87890506]);
    fprintf('done\n');
    
    %% cleanup
    caffe.reset_all();

  end
end