| 125 | } |
| 126 | |
| 127 | function parseArguments() { |
| 128 | const argParser = new argparse.ArgumentParser({ |
| 129 | description: |
| 130 | 'Train an attention-based date-conversion model in TensorFlow.js' |
| 131 | }); |
| 132 | argParser.addArgument('--gpu', { |
| 133 | action: 'storeTrue', |
| 134 | help: 'Use tfjs-node-gpu to train the model. Requires CUDA/CuDNN.' |
| 135 | }); |
| 136 | argParser.addArgument('--epochs', { |
| 137 | type: 'int', |
| 138 | defaultValue: 2, |
| 139 | help: 'Number of epochs to train the model for' |
| 140 | }); |
| 141 | argParser.addArgument('--batchSize', { |
| 142 | type: 'int', |
| 143 | defaultValue: 128, |
| 144 | help: 'Batch size to be used during model training' |
| 145 | }); |
| 146 | argParser.addArgument('--trainSplit ', { |
| 147 | type: 'float', |
| 148 | defaultValue: 0.25, |
| 149 | help: 'Fraction of all possible dates to use for training. Must be ' + |
| 150 | '> 0 and < 1. Its sum with valSplit must be <1.' |
| 151 | }); |
| 152 | argParser.addArgument('--valSplit', { |
| 153 | type: 'float', |
| 154 | defaultValue: 0.15, |
| 155 | help: 'Fraction of all possible dates to use for training. Must be ' + |
| 156 | '> 0 and < 1. Its sum with trainSplit must be <1.' |
| 157 | }); |
| 158 | argParser.addArgument('--savePath', { |
| 159 | type: 'string', |
| 160 | defaultValue: './dist/model', |
| 161 | }); |
| 162 | argParser.addArgument('--logDir', { |
| 163 | type: 'string', |
| 164 | help: 'Optional tensorboard log directory, to which the loss and ' + |
| 165 | 'accuracy will be logged during model training.' |
| 166 | }); |
| 167 | argParser.addArgument('--logUpdateFreq', { |
| 168 | type: 'string', |
| 169 | defaultValue: 'batch', |
| 170 | optionStrings: ['batch', 'epoch'], |
| 171 | help: 'Frequency at which the loss and accuracy will be logged to ' + |
| 172 | 'tensorboard.' |
| 173 | }); |
| 174 | return argParser.parseArgs(); |
| 175 | } |
| 176 | |
| 177 | async function run() { |
| 178 | const args = parseArguments(); |