MCPcopy Create free account
hub / github.com/BorealisAI/scaleformer / create_stack

Method create_stack

models/NHitsMS.py:461–518  ·  view source on GitHub ↗
(self, stack_types, n_blocks,
                     n_time_in, n_time_out,
                     n_x, n_x_hidden, n_s, n_s_hidden,
                     n_layers, n_mlp_units,
                     n_pool_kernel_size, n_freq_downsample, pooling_mode, interpolation_mode,
                     batch_normalization, dropout_prob_theta,
                     activation, shared_weights, initialization)

Source from the content-addressed store, hash-verified

459 self.blocks = t.nn.ModuleList(self.blocks)
460
461 def create_stack(self, stack_types, n_blocks,
462 n_time_in, n_time_out,
463 n_x, n_x_hidden, n_s, n_s_hidden,
464 n_layers, n_mlp_units,
465 n_pool_kernel_size, n_freq_downsample, pooling_mode, interpolation_mode,
466 batch_normalization, dropout_prob_theta,
467 activation, shared_weights, initialization):
468 block_list = []
469 for i in range(len(stack_types)):
470 assert stack_types[i] in ['identity', 'exogenous', 'exogenous_tcn', 'exogenous_wavenet'], 'f Invalid stack type {stack_types[i]}'
471 for block_id in range(n_blocks[i]):
472 # Batch norm only on first block
473 if (len(block_list)==0) and (batch_normalization):
474 batch_normalization_block = True
475 else:
476 batch_normalization_block = False
477 # Shared weights
478 if shared_weights and block_id>0:
479 nbeats_block = block_list[-1]
480 else:
481 if stack_types[i] == 'identity':
482 n_theta = (n_time_in + max(n_time_out//n_freq_downsample[i], 1) )
483 basis = _IdentityBasis(backcast_size=n_time_in,
484 forecast_size=n_time_out,
485 interpolation_mode=interpolation_mode)
486
487 elif stack_types[i] == 'exogenous':
488 n_theta = 2 * n_x
489 basis = _ExogenousBasisInterpretable()
490
491 elif stack_types[i] == 'exogenous_tcn':
492 n_theta = 2 * n_x_hidden
493 basis = _ExogenousBasisTCN(n_x_hidden, n_x)
494
495 elif stack_types[i] == 'exogenous_wavenet':
496 n_theta = 2 * n_x_hidden
497 basis = _ExogenousBasisWavenet(n_x_hidden, n_x)
498
499 nbeats_block = _NHITSBlock(n_time_in=n_time_in,
500 n_time_out=n_time_out,
501 n_x=n_x,
502 n_s=n_s,
503 n_s_hidden=n_s_hidden,
504 n_theta=n_theta,
505 n_mlp_units=n_mlp_units[i],
506 n_pool_kernel_size=n_pool_kernel_size[i],
507 pooling_mode=pooling_mode,
508 basis=basis,
509 n_layers=n_layers[i],
510 batch_normalization=batch_normalization_block,
511 dropout_prob=dropout_prob_theta,
512 activation=activation)
513
514 # Select type of evaluation and apply it to all layers of block
515 init_function = partial(_init_weights, initialization=initialization)
516 nbeats_block.layers.apply(init_function)
517 block_list.append(nbeats_block)
518 return block_list

Callers 1

__init__Method · 0.95

Calls 5

_IdentityBasisClass · 0.70
_ExogenousBasisTCNClass · 0.70
_NHITSBlockClass · 0.70

Tested by

no test coverage detected