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

Function main

quantization/train_housing.js:69–105  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

67}
68
69async function main() {
70 const args = parseArgs();
71 if (args.gpu) {
72 tf = require('@tensorflow/tfjs-node-gpu');
73 } else {
74 tf = require('@tensorflow/tfjs-node');
75 }
76
77 const {count, featureMeans, featureStddevs, labelMean, labelStddev} =
78 await getDatasetStats();
79 const {trainXs, trainYs, valXs, valYs, evalXs, evalYs} =
80 await getNormalizedDatasets(
81 count, featureMeans, featureStddevs, labelMean, labelStddev,
82 args.validationSplit, args.evaluationSplit);
83
84 const model = createModel();
85 model.summary();
86
87 await model.fit(trainXs, trainYs, {
88 epochs: args.epochs,
89 batchSize: args.batchSize,
90 validationData: [valXs, valYs]
91 });
92
93 const evalOutput = model.evaluate(evalXs, evalYs);
94 console.log(
95 `\nEvaluation result:\n` +
96 ` Loss = ${evalOutput.dataSync()[0].toFixed(6)}`);
97
98 if (args.modelSavePath != null) {
99 if (!fs.existsSync(path.dirname(args.modelSavePath))) {
100 shelljs.mkdir('-p', path.dirname(args.modelSavePath));
101 }
102 await model.save(`file://${args.modelSavePath}`);
103 console.log(`Saved model to path: ${args.modelSavePath}`);
104 }
105}
106
107if (require.main === module) {
108 main();

Callers 1

train_housing.jsFile · 0.70

Calls 4

getDatasetStatsFunction · 0.90
getNormalizedDatasetsFunction · 0.90
createModelFunction · 0.90
parseArgsFunction · 0.70

Tested by

no test coverage detected