()
| 52 | } |
| 53 | |
| 54 | async function main() { |
| 55 | const args = parseArgs(); |
| 56 | if (args.gpu) { |
| 57 | tf = require('@tensorflow/tfjs-node-gpu'); |
| 58 | } else { |
| 59 | tf = require('@tensorflow/tfjs-node'); |
| 60 | } |
| 61 | |
| 62 | let dataset; |
| 63 | if (args.dataset === 'fashion-mnist') { |
| 64 | dataset = new FashionMnistDataset(); |
| 65 | } else if (args.dataset === 'mnist') { |
| 66 | dataset = new MnistDataset(); |
| 67 | } else { |
| 68 | throw new Error(`Unrecognized dataset name: ${args.dataset}`); |
| 69 | } |
| 70 | await dataset.loadData(); |
| 71 | const {images: testImages, labels: testLabels} = dataset.getTestData(); |
| 72 | |
| 73 | console.log(`Loading model from ${args.modelSavePath}...`); |
| 74 | const model = await tf.loadLayersModel(`file://${args.modelSavePath}`); |
| 75 | compileModel(model); |
| 76 | |
| 77 | console.log(`Performing evaluation...`); |
| 78 | const t0 = tf.util.now(); |
| 79 | const evalOutput = model.evaluate(testImages, testLabels); |
| 80 | const t1 = tf.util.now(); |
| 81 | console.log(`\nEvaluation took ${(t1 - t0).toFixed(2)} ms.`); |
| 82 | console.log( |
| 83 | `\nEvaluation result:\n` + |
| 84 | ` Loss = ${evalOutput[0].dataSync()[0].toFixed(6)}; `+ |
| 85 | `Accuracy = ${evalOutput[1].dataSync()[0].toFixed(6)}`); |
| 86 | } |
| 87 | |
| 88 | if (require.main === module) { |
| 89 | main(); |
no test coverage detected