* Train a `tf.Model` to recognize Iris flower type. * * @param xTrain Training feature data, a `tf.Tensor` of shape * [numTrainExamples, 4]. The second dimension include the features * petal length, petalwidth, sepal length and sepal width. * @param yTrain One-hot training labels, a `tf.Ten
(xTrain, yTrain, xTest, yTest)
| 38 | * @returns The trained `tf.Model` instance. |
| 39 | */ |
| 40 | async function trainModel(xTrain, yTrain, xTest, yTest) { |
| 41 | ui.status('Training model... Please wait.'); |
| 42 | |
| 43 | const params = ui.loadTrainParametersFromUI(); |
| 44 | |
| 45 | // Define the topology of the model: two dense layers. |
| 46 | const model = tf.sequential(); |
| 47 | model.add(tf.layers.dense( |
| 48 | {units: 10, activation: 'sigmoid', inputShape: [xTrain.shape[1]]})); |
| 49 | model.add(tf.layers.dense({units: 3, activation: 'softmax'})); |
| 50 | model.summary(); |
| 51 | |
| 52 | const optimizer = tf.train.adam(params.learningRate); |
| 53 | model.compile({ |
| 54 | optimizer: optimizer, |
| 55 | loss: 'categoricalCrossentropy', |
| 56 | metrics: ['accuracy'], |
| 57 | }); |
| 58 | |
| 59 | const trainLogs = []; |
| 60 | const lossContainer = document.getElementById('lossCanvas'); |
| 61 | const accContainer = document.getElementById('accuracyCanvas'); |
| 62 | const beginMs = performance.now(); |
| 63 | // Call `model.fit` to train the model. |
| 64 | const history = await model.fit(xTrain, yTrain, { |
| 65 | epochs: params.epochs, |
| 66 | validationData: [xTest, yTest], |
| 67 | callbacks: { |
| 68 | onEpochEnd: async (epoch, logs) => { |
| 69 | // Plot the loss and accuracy values at the end of every training epoch. |
| 70 | const secPerEpoch = |
| 71 | (performance.now() - beginMs) / (1000 * (epoch + 1)); |
| 72 | ui.status(`Training model... Approximately ${ |
| 73 | secPerEpoch.toFixed(4)} seconds per epoch`) |
| 74 | trainLogs.push(logs); |
| 75 | tfvis.show.history(lossContainer, trainLogs, ['loss', 'val_loss']) |
| 76 | tfvis.show.history(accContainer, trainLogs, ['acc', 'val_acc']) |
| 77 | calculateAndDrawConfusionMatrix(model, xTest, yTest); |
| 78 | }, |
| 79 | } |
| 80 | }); |
| 81 | const secPerEpoch = (performance.now() - beginMs) / (1000 * params.epochs); |
| 82 | ui.status( |
| 83 | `Model training complete: ${secPerEpoch.toFixed(4)} seconds per epoch`); |
| 84 | return model; |
| 85 | } |
| 86 | |
| 87 | /** |
| 88 | * Run inference on manually-input Iris flower data. |
no test coverage detected