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

Function createAndCompileModel

addition-rnn/index.js:172–237  ·  view source on GitHub ↗
(
    layers, hiddenSize, rnnType, digits, vocabularySize)

Source from the content-addressed store, hash-verified

170}
171
172function 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})}));

Callers 1

constructorMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected