MCPcopy Create free account
hub / github.com/tensorflow/tfjs-examples / trainModel

Function trainModel

iris/index.js:40–85  ·  view source on GitHub ↗

* 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)

Source from the content-addressed store, hash-verified

38 * @returns The trained `tf.Model` instance.
39 */
40async 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.

Callers 1

irisFunction · 0.70

Calls 1

Tested by

no test coverage detected