(trainSplit = 0.25, valSplit = 0.15)
| 50 | * - testDateTuples, date tuples ([year, month, day]) for the test set. |
| 51 | */ |
| 52 | export 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 { |
no test coverage detected