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

Function generateDataForTraining

date-conversion-attention/train.js:52–125  ·  view source on GitHub ↗
(trainSplit = 0.25, valSplit = 0.15)

Source from the content-addressed store, hash-verified

50 * - testDateTuples, date tuples ([year, month, day]) for the test set.
51 */
52export function generateDataForTraining(trainSplit = 0.25, valSplit = 0.15) {
53 tf.util.assert(
54 trainSplit > 0 && valSplit > 0 && trainSplit + valSplit <= 1,
55 `Invalid trainSplit (${trainSplit}) and valSplit (${valSplit})`);
56
57 const dateTuples = [];
58 const MIN_YEAR = 1950;
59 const MAX_YEAR = 2050;
60 for (let date = new Date(MIN_YEAR,0,1);
61 date.getFullYear() < MAX_YEAR;
62 date.setDate(date.getDate() + 1)) {
63 dateTuples.push([date.getFullYear(), date.getMonth() + 1, date.getDate()]);
64 }
65 tf.util.shuffle(dateTuples);
66
67 const numTrain = Math.floor(dateTuples.length * trainSplit);
68 const numVal = Math.floor(dateTuples.length * valSplit);
69 console.log(`Number of dates used for training: ${numTrain}`);
70 console.log(`Number of dates used for validation: ${numVal}`);
71 console.log(
72 `Number of dates used for testing: ` +
73 `${dateTuples.length - numTrain - numVal}`);
74
75 function dateTuplesToTensor(dateTuples) {
76 return tf.tidy(() => {
77 const inputs =
78 dateFormat.INPUT_FNS.map(fn => dateTuples.map(tuple => fn(tuple)));
79 const inputStrings = [];
80 inputs.forEach(inputs => inputStrings.push(...inputs));
81 const encoderInput =
82 dateFormat.encodeInputDateStrings(inputStrings);
83 const trainTargetStrings = dateTuples.map(
84 tuple => dateFormat.dateTupleToYYYYDashMMDashDD(tuple));
85 let decoderInput =
86 dateFormat.encodeOutputDateStrings(trainTargetStrings)
87 .asType('float32');
88 // One-step time shift: The decoder input is shifted to the left by
89 // one time step with respect to the encoder input. This accounts for
90 // the step-by-step decoding that happens during inference time.
91 decoderInput = tf.concat([
92 tf.ones([decoderInput.shape[0], 1]).mul(dateFormat.START_CODE),
93 decoderInput.slice(
94 [0, 0], [decoderInput.shape[0], decoderInput.shape[1] - 1])
95 ], 1).tile([dateFormat.INPUT_FNS.length, 1]);
96 const decoderOutput = tf.oneHot(
97 dateFormat.encodeOutputDateStrings(trainTargetStrings),
98 dateFormat.OUTPUT_VOCAB.length).tile(
99 [dateFormat.INPUT_FNS.length, 1, 1]);
100 return {encoderInput, decoderInput, decoderOutput};
101 });
102 }
103
104 const {
105 encoderInput: trainEncoderInput,
106 decoderInput: trainDecoderInput,
107 decoderOutput: trainDecoderOutput
108 } = dateTuplesToTensor(dateTuples.slice(0, numTrain));
109 const {

Callers 2

train_test.jsFile · 0.90
runFunction · 0.85

Calls 1

dateTuplesToTensorFunction · 0.85

Tested by

no test coverage detected