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

Function createAndCompileModel

addition-rnn-webworker/worker.js:164–229  ·  view source on GitHub ↗
(
  layers, hiddenSize, rnnType, digits, vocabularySize)

Source from the content-addressed store, hash-verified

162}
163
164function 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 }) }));

Callers 1

constructorMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected