var sample_training_instance = function() { // find an unloaded batch var bi = Math.floor(Math.random()*loaded_train_batches.length); var b = loaded_train_batches[bi]; var k = Math.floor(Math.random()*num_samples_per_batch); // sample within the batch var n = b*num_samples_per_batch+k; // load more batches over time if(step_num%(2 * num_samples_per_batch)===0 && step_num>0) { for(var i=0;i 0 ? dwq : 0.0; } draw_activations_COLOR(activations_div, dd, scale); for(var q=0;q3) { // actual weights filters_div.appendChild(document.createTextNode('Weights:')); filters_div.appendChild(document.createElement('br')); for(var j=0;j' + classes_txt[preds[k].k] + '' } probsdiv.innerHTML = t; probsdiv.className = 'probsdiv'; div.appendChild(probsdiv); // add it into DOM $(div).prependTo($("#testset_vis")).hide().fadeIn('slow').slideDown('slow'); if($(".probsdiv").length>200) { $("#testset_vis > .probsdiv").last().remove(); // pop to keep upper bound of shown items } } testAccWindow.add(num_correct/num_total); $("#testset_acc").text('test accuracy based on last 200 test images: ' + testAccWindow.get_average()); } var testImage = function(img) { var x = convnetjs.img_to_vol(img); var out_p = net.forward(x); var vis_elt = document.getElementById("visnet"); visualize_activations(net, vis_elt); var preds =[] for(var k=0;k' + classes_txt[preds[k].k] + '' } probsdiv.innerHTML = t; probsdiv.className = 'probsdiv'; div.appendChild(probsdiv); // add it into DOM $(div).prependTo($("#testset_vis")).hide().fadeIn('slow').slideDown('slow'); if($(".probsdiv").length>200) { $("#testset_vis > .probsdiv").last().remove(); // pop to keep upper bound of shown items } } var lossGraph = new cnnvis.Graph(); var xLossWindow = new cnnutil.Window(100); var wLossWindow = new cnnutil.Window(100); var trainAccWindow = new cnnutil.Window(100); var valAccWindow = new cnnutil.Window(100); var testAccWindow = new cnnutil.Window(50, 1); var step_num = 0; var step = function(sample) { var x = sample.x; var y = sample.label; if(sample.isval) { // use x to build our estimate of validation error net.forward(x); var yhat = net.getPrediction(); var val_acc = yhat === y ? 1.0 : 0.0; valAccWindow.add(val_acc); return; // get out } // train on it with network var stats = trainer.train(x, y); var lossx = stats.cost_loss; var lossw = stats.l2_decay_loss; // keep track of stats such as the average training error and loss var yhat = net.getPrediction(); var train_acc = yhat === y ? 1.0 : 0.0; xLossWindow.add(lossx); wLossWindow.add(lossw); trainAccWindow.add(train_acc); // visualize training status var train_elt = document.getElementById("trainstats"); train_elt.innerHTML = ''; var t = 'Forward time per example: ' + stats.fwd_time + 'ms'; train_elt.appendChild(document.createTextNode(t)); train_elt.appendChild(document.createElement('br')); var t = 'Backprop time per example: ' + stats.bwd_time + 'ms'; train_elt.appendChild(document.createTextNode(t)); train_elt.appendChild(document.createElement('br')); var t = 'Classification loss: ' + f2t(xLossWindow.get_average()); train_elt.appendChild(document.createTextNode(t)); train_elt.appendChild(document.createElement('br')); var t = 'L2 Weight decay loss: ' + f2t(wLossWindow.get_average()); train_elt.appendChild(document.createTextNode(t)); train_elt.appendChild(document.createElement('br')); var t = 'Training accuracy: ' + f2t(trainAccWindow.get_average()); train_elt.appendChild(document.createTextNode(t)); train_elt.appendChild(document.createElement('br')); var t = 'Validation accuracy: ' + f2t(valAccWindow.get_average()); train_elt.appendChild(document.createTextNode(t)); train_elt.appendChild(document.createElement('br')); var t = 'Examples seen: ' + step_num; train_elt.appendChild(document.createTextNode(t)); train_elt.appendChild(document.createElement('br')); // visualize activations if(step_num % 100 === 0) { var vis_elt = document.getElementById("visnet"); visualize_activations(net, vis_elt); } // log progress to graph, (full loss) if(step_num % 200 === 0) { var xa = xLossWindow.get_average(); var xw = wLossWindow.get_average(); if(xa >= 0 && xw >= 0) { // if they are -1 it means not enough data was accumulated yet for estimates lossGraph.add(step_num, xa + xw); lossGraph.drawSelf(document.getElementById("lossgraph")); } } // run prediction on test set if((step_num % 100 === 0 && step_num > 0) || step_num===100) { test_predict(); } step_num++; } // user settings var change_lr = function() { trainer.learning_rate = parseFloat(document.getElementById("lr_input").value); update_net_param_display(); } var change_momentum = function() { trainer.momentum = parseFloat(document.getElementById("momentum_input").value); update_net_param_display(); } var change_batch_size = function() { trainer.batch_size = parseFloat(document.getElementById("batch_size_input").value); update_net_param_display(); } var change_decay = function() { trainer.l2_decay = parseFloat(document.getElementById("decay_input").value); update_net_param_display(); } var update_net_param_display = function() { document.getElementById('lr_input').value = trainer.learning_rate; document.getElementById('momentum_input').value = trainer.momentum; document.getElementById('batch_size_input').value = trainer.batch_size; document.getElementById('decay_input').value = trainer.l2_decay; } var toggle_pause = function() { paused = !paused; var btn = document.getElementById('buttontp'); if(paused) { btn.value = 'resume' } else { btn.value = 'pause'; } } var dump_json = function() { document.getElementById("dumpjson").value = JSON.stringify(this.net.toJSON()); } var clear_graph = function() { lossGraph = new cnnvis.Graph(); // reinit graph too } var reset_all = function() { // reinit trainer trainer = new convnetjs.SGDTrainer(net, {learning_rate:trainer.learning_rate, momentum:trainer.momentum, batch_size:trainer.batch_size, l2_decay:trainer.l2_decay}); update_net_param_display(); // reinit windows that keep track of val/train accuracies xLossWindow.reset(); wLossWindow.reset(); trainAccWindow.reset(); valAccWindow.reset(); testAccWindow.reset(); lossGraph = new cnnvis.Graph(); // reinit graph too step_num = 0; } var load_from_json = function() { var jsonString = document.getElementById("dumpjson").value; var json = JSON.parse(jsonString); net = new convnetjs.Net(); net.fromJSON(json); reset_all(); } var load_pretrained = function() { $.getJSON(dataset_name + "_snapshot.json", function(json){ net = new convnetjs.Net(); net.fromJSON(json); trainer.learning_rate = 0.0001; trainer.momentum = 0.9; trainer.batch_size = 2; trainer.l2_decay = 0.00001; reset_all(); }); } var change_net = function() { eval($("#newnet").val()); reset_all(); }