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