MCPcopy Create free account
hub / github.com/CompVis/zigma / create_block

Function create_block

model_zigma.py:468–509  ·  view source on GitHub ↗
(
    d_model,
    ssm_cfg=None,
    has_text=False,
    norm_epsilon=1e-5,
    drop_path=0.0,
    rms_norm=False,
    residual_in_fp32=False,
    fused_add_norm=False,
    skip=False,
    layer_idx=None,
    device=None,
    dtype=None,
    scan_type="none",
    **block_kwargs,
)

Source from the content-addressed store, hash-verified

466
467
468def create_block(
469 d_model,
470 ssm_cfg=None,
471 has_text=False,
472 norm_epsilon=1e-5,
473 drop_path=0.0,
474 rms_norm=False,
475 residual_in_fp32=False,
476 fused_add_norm=False,
477 skip=False,
478 layer_idx=None,
479 device=None,
480 dtype=None,
481 scan_type="none",
482 **block_kwargs,
483):
484 if ssm_cfg is None:
485 ssm_cfg = {}
486 factory_kwargs = {"device": device, "dtype": dtype}
487 mixer_cls = partial(
488 Mamba,
489 layer_idx=layer_idx,
490 scan_type=scan_type,
491 **ssm_cfg,
492 **block_kwargs,
493 **factory_kwargs,
494 )
495 norm_cls = partial(
496 nn.LayerNorm if not rms_norm else RMSNorm, eps=norm_epsilon, **factory_kwargs
497 )
498 block = Block(
499 d_model,
500 mixer_cls,
501 has_text=has_text,
502 norm_cls=norm_cls,
503 drop_path=drop_path,
504 fused_add_norm=fused_add_norm,
505 residual_in_fp32=residual_in_fp32,
506 skip=skip,
507 )
508 block.layer_idx = layer_idx
509 return block
510
511
512def _init_weights(

Callers 1

__init__Method · 0.70

Calls 1

BlockClass · 0.70

Tested by

no test coverage detected