| 170 | } |
| 171 | |
| 172 | function createAndCompileModel( |
| 173 | layers, hiddenSize, rnnType, digits, vocabularySize) { |
| 174 | const maxLen = digits + 1 + digits; |
| 175 | |
| 176 | const model = tf.sequential(); |
| 177 | switch (rnnType) { |
| 178 | case 'SimpleRNN': |
| 179 | model.add(tf.layers.simpleRNN({ |
| 180 | units: hiddenSize, |
| 181 | recurrentInitializer: 'glorotNormal', |
| 182 | inputShape: [maxLen, vocabularySize] |
| 183 | })); |
| 184 | break; |
| 185 | case 'GRU': |
| 186 | model.add(tf.layers.gru({ |
| 187 | units: hiddenSize, |
| 188 | recurrentInitializer: 'glorotNormal', |
| 189 | inputShape: [maxLen, vocabularySize] |
| 190 | })); |
| 191 | break; |
| 192 | case 'LSTM': |
| 193 | model.add(tf.layers.lstm({ |
| 194 | units: hiddenSize, |
| 195 | recurrentInitializer: 'glorotNormal', |
| 196 | inputShape: [maxLen, vocabularySize] |
| 197 | })); |
| 198 | break; |
| 199 | default: |
| 200 | throw new Error(`Unsupported RNN type: '${rnnType}'`); |
| 201 | } |
| 202 | model.add(tf.layers.repeatVector({n: digits + 1})); |
| 203 | switch (rnnType) { |
| 204 | case 'SimpleRNN': |
| 205 | model.add(tf.layers.simpleRNN({ |
| 206 | units: hiddenSize, |
| 207 | recurrentInitializer: 'glorotNormal', |
| 208 | returnSequences: true |
| 209 | })); |
| 210 | break; |
| 211 | case 'GRU': |
| 212 | model.add(tf.layers.gru({ |
| 213 | units: hiddenSize, |
| 214 | recurrentInitializer: 'glorotNormal', |
| 215 | returnSequences: true |
| 216 | })); |
| 217 | break; |
| 218 | case 'LSTM': |
| 219 | model.add(tf.layers.lstm({ |
| 220 | units: hiddenSize, |
| 221 | recurrentInitializer: 'glorotNormal', |
| 222 | returnSequences: true |
| 223 | })); |
| 224 | break; |
| 225 | default: |
| 226 | throw new Error(`Unsupported RNN type: '${rnnType}'`); |
| 227 | } |
| 228 | model.add(tf.layers.timeDistributed( |
| 229 | {layer: tf.layers.dense({units: vocabularySize})})); |