Function
create_block
(
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 | |
| 468 | def 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 | |
| 512 | def _init_weights( |
Tested by
no test coverage detected