From 64e419f625f215f0e8a0e2e591f73356052b0e44 Mon Sep 17 00:00:00 2001 From: Artyom Kozhevnikov Date: Tue, 5 Aug 2025 03:02:05 -0700 Subject: [PATCH 1/3] adding bleu metric for s2tt --- src/fairseq2/models/wav2vec2/asr/_config.py | 2 +- src/fairseq2/nn/utils/module.py | 15 +++++++++++++++ src/fairseq2/recipes/asr/_common.py | 11 ++++++++++- 3 files changed, 26 insertions(+), 2 deletions(-) diff --git a/src/fairseq2/models/wav2vec2/asr/_config.py b/src/fairseq2/models/wav2vec2/asr/_config.py index 7d45e6b80..6305b594e 100644 --- a/src/fairseq2/models/wav2vec2/asr/_config.py +++ b/src/fairseq2/models/wav2vec2/asr/_config.py @@ -160,7 +160,7 @@ def bib61_300m() -> Wav2Vec2AsrConfig: @wav2vec2_asr_arch("300m_bib1143") def bib1143_300m() -> Wav2Vec2AsrConfig: config = bib61_300m() - config.vocab_info.size = 3335 + config.vocab_info.size = 3292 return config @wav2vec2_asr_arch("1b_bib61") diff --git a/src/fairseq2/nn/utils/module.py b/src/fairseq2/nn/utils/module.py index a3ba5f0ce..c6e6985ac 100644 --- a/src/fairseq2/nn/utils/module.py +++ b/src/fairseq2/nn/utils/module.py @@ -464,6 +464,21 @@ def load_state_dict( ``state_dict`` does not contain any keys corresponding to descendants that are set to ``None`` via :meth:`Module.register_module()`. """ + # Key mapping + need_mapping = False + sample_key = list(state_dict.keys())[0] + if ( + sample_key.startswith("module.") + and not sample_key in module.state_dict().keys() + ): + mapped_key = sample_key[7:] + if mapped_key in module.state_dict().keys(): + need_mapping = True + + if need_mapping: + key_mapping = lambda key: key[7:] if key.startswith("module.") else key + state_dict = {key_mapping(key): value for key, value in state_dict.items()} + module.load_state_dict(state_dict, strict=strict) unexpected_keys = [] diff --git a/src/fairseq2/recipes/asr/_common.py b/src/fairseq2/recipes/asr/_common.py index b023b714b..9c12c0900 100644 --- a/src/fairseq2/recipes/asr/_common.py +++ b/src/fairseq2/recipes/asr/_common.py @@ -16,7 +16,7 @@ from fairseq2.data.text.tokenizers import TextTokenDecoder, TextTokenizer from fairseq2.gang import Gang from fairseq2.metrics import Mean -from fairseq2.metrics.text import WerMetric +from fairseq2.metrics.text import BleuMetric, WerMetric from fairseq2.models.asr import AsrModel, AsrModelOutput from fairseq2.models.seq2seq import Seq2SeqBatch from fairseq2.models.sequence import SequenceBatch @@ -121,6 +121,8 @@ def __call__( refs, ref_seqs, ref_padding_mask, hyps, hyp_seqs, hyp_padding_mask ) + metric_bag.bleu.update(refs, hyps) + try: # Dump references. stream = self._ref_output_stream @@ -148,6 +150,7 @@ def __call__( class AsrMetricBag(BaseMetricBag): ctc_loss: Mean wer: WerMetric + bleu: BleuMetric def __init__(self, gang: Gang, train: bool = True) -> None: super().__init__(gang, train=train) @@ -158,6 +161,12 @@ def __init__(self, gang: Gang, train: bool = True) -> None: self.register_metric("wer", WerMetric(device=self.device), persistent=False) + self.register_metric( + "bleu", + BleuMetric(tokenizer="flores200", device=self.device), + persistent=False, + ) + @torch.inference_mode() def update_ctc_loss(self, batch: Seq2SeqBatch, loss: Tensor) -> None: n = batch.batch_size From a4cb0be1f356bf270d69891376a3d6e65e68dfba Mon Sep 17 00:00:00 2001 From: Artyom Kozhevnikov Date: Mon, 11 Aug 2025 10:44:52 -0700 Subject: [PATCH 2/3] rm custom stuff --- src/fairseq2/models/wav2vec2/asr/_config.py | 308 ++------------------ src/fairseq2/nn/utils/module.py | 101 +------ 2 files changed, 31 insertions(+), 378 deletions(-) diff --git a/src/fairseq2/models/wav2vec2/asr/_config.py b/src/fairseq2/models/wav2vec2/asr/_config.py index 6305b594e..c65739c1a 100644 --- a/src/fairseq2/models/wav2vec2/asr/_config.py +++ b/src/fairseq2/models/wav2vec2/asr/_config.py @@ -10,8 +10,7 @@ from typing import Final from fairseq2.context import RuntimeContext -from fairseq2.data import VocabularyInfo -from fairseq2.models.wav2vec2 import Wav2Vec2EncoderConfig +from fairseq2.models.wav2vec2 import Wav2Vec2Config, Wav2Vec2EncoderConfig WAV2VEC2_ASR_MODEL_FAMILY: Final = "wav2vec2_asr" @@ -26,7 +25,7 @@ class Wav2Vec2AsrConfig: encoder_config: Wav2Vec2EncoderConfig = field( default_factory=lambda: Wav2Vec2EncoderConfig( - feature_gradient_scale=1.0, + feature_grad_scale=1.0, dropout_p=0.0, attn_dropout_p=0.0, ffn_inner_dropout_p=0.1, @@ -34,19 +33,13 @@ class Wav2Vec2AsrConfig: ) """The configuration of the encoder.""" - vocab_info: VocabularyInfo = field( - default_factory=lambda: VocabularyInfo( - size=32, unk_idx=3, bos_idx=0, eos_idx=2, pad_idx=1 - ) - ) - """The vocabulary information.""" + target_vocab_size: int = 32 + """The size of the target vocabulary.""" final_dropout_p: float = 0.0 """The dropout probability on the output of the encoder.""" # Mask - mask_codebase: str = "fairseq2" - use_masking: bool = True """If ``True``, masks features as regularization.""" @@ -73,15 +66,14 @@ class Wav2Vec2AsrConfig: def register_wav2vec2_asr_configs(context: RuntimeContext) -> None: registry = context.get_config_registry(Wav2Vec2AsrConfig) - wav2vec2_asr_arch = registry.decorator - w2v2_encoder_registry = context.get_config_registry(Wav2Vec2EncoderConfig) + arch = registry.decorator - @wav2vec2_asr_arch("base_10h") + @arch("base_10h") def base_10h() -> Wav2Vec2AsrConfig: return Wav2Vec2AsrConfig() - @wav2vec2_asr_arch("base_100h") + @arch("base_100h") def base_100h() -> Wav2Vec2AsrConfig: config = base_10h() @@ -89,12 +81,16 @@ def base_100h() -> Wav2Vec2AsrConfig: return config - @wav2vec2_asr_arch("large_10h") + w2v2_registry = context.get_config_registry(Wav2Vec2Config) + + @arch("large_10h") def large_10h() -> Wav2Vec2AsrConfig: config = base_10h() - config.encoder_config = w2v2_encoder_registry.get("large") - config.encoder_config.feature_gradient_scale = 1.0 + w2v2_config = w2v2_registry.get("large") + + config.encoder_config = w2v2_config.encoder_config + config.encoder_config.feature_grad_scale = 1.0 config.encoder_config.dropout_p = 0.0 config.encoder_config.attn_dropout_p = 0.0 config.encoder_config.ffn_inner_dropout_p = 0.1 @@ -105,7 +101,7 @@ def large_10h() -> Wav2Vec2AsrConfig: return config - @wav2vec2_asr_arch("large_100h") + @arch("large_100h") def large_100h() -> Wav2Vec2AsrConfig: config = large_10h() @@ -114,12 +110,14 @@ def large_100h() -> Wav2Vec2AsrConfig: return config - @wav2vec2_asr_arch("large_lv60k_10h") + @arch("large_lv60k_10h") def large_lv60k_10h() -> Wav2Vec2AsrConfig: config = base_10h() - config.encoder_config = w2v2_encoder_registry.get("large_lv60k") - config.encoder_config.feature_gradient_scale = 1.0 + w2v2_config = w2v2_registry.get("large_lv60k") + + config.encoder_config = w2v2_config.encoder_config + config.encoder_config.feature_grad_scale = 1.0 config.encoder_config.dropout_p = 0.0 config.encoder_config.attn_dropout_p = 0.0 config.encoder_config.ffn_inner_dropout_p = 0.1 @@ -130,7 +128,7 @@ def large_lv60k_10h() -> Wav2Vec2AsrConfig: return config - @wav2vec2_asr_arch("large_lv60k_100h") + @arch("large_lv60k_100h") def large_lv60k_100h() -> Wav2Vec2AsrConfig: config = large_lv60k_10h() @@ -138,267 +136,3 @@ def large_lv60k_100h() -> Wav2Vec2AsrConfig: config.max_spatial_mask_prob = 0.55 return config - - @wav2vec2_asr_arch("300m_bib61") - def bib61_300m() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("large_lv60k") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 2475 - - return config - - @wav2vec2_asr_arch("300m_bib1143") - def bib1143_300m() -> Wav2Vec2AsrConfig: - config = bib61_300m() - config.vocab_info.size = 3292 - return config - - @wav2vec2_asr_arch("1b_bib61") - def bib61_1b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("1b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 2475 - - return config - - @wav2vec2_asr_arch("1b_llama_bib61") - def llama_bib61_1b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("1b_llama") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 2475 - - return config - - @wav2vec2_asr_arch("2b_bib61") - def bib61_2b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("2b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 2475 - - return config - - @wav2vec2_asr_arch("3b_bib61") - def bib61_3b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("3b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 2475 - - return config - - @wav2vec2_asr_arch("5b_bib61") - def bib61_5b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("5b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 2475 - - return config - - @wav2vec2_asr_arch("7b_bib61") - def bib61_7b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("7b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 2475 - - return config - - @wav2vec2_asr_arch("3.25b_bib61") - def higher_bib61_3b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("3.25b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 2475 - - return config - - @wav2vec2_asr_arch("5b_front51") - def front51_5b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("5b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 222 - - return config - - @wav2vec2_asr_arch("7b_front51") - def front51_7b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("7b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 222 - - return config - - @wav2vec2_asr_arch("1b_bib1143") - def bib1143_1b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("1b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 3335 - - return config - - @wav2vec2_asr_arch("3b_bib1143") - def bib1143_3b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("3b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 3335 - - return config - - @wav2vec2_asr_arch("5b_bib1143") - def bib1143_5b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("5b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 3335 # following bibfront1194's vocab size - - return config - - @wav2vec2_asr_arch("7b_bib1143") - def bib1143_7b() -> Wav2Vec2AsrConfig: - config = base_10h() - - config.encoder_config = w2v2_encoder_registry.get("7b") - config.encoder_config.feature_gradient_scale = 1.0 - config.encoder_config.dropout_p = 0.0 - config.encoder_config.attn_dropout_p = 0.0 - config.encoder_config.ffn_inner_dropout_p = 0.1 - config.encoder_config.layer_drop_p = 0.1 - - config.use_masking = False - config.max_temporal_mask_prob = 0.0 - config.max_spatial_mask_prob = 0.0 - config.vocab_info.size = 3335 - - return config - - @wav2vec2_asr_arch("7b_v3_tokenizer") - def v3_tokenizer_7b() -> Wav2Vec2AsrConfig: - config = bib1143_7b() - config.vocab_info.size = 9656 - return config diff --git a/src/fairseq2/nn/utils/module.py b/src/fairseq2/nn/utils/module.py index c6e6985ac..0a69a00bf 100644 --- a/src/fairseq2/nn/utils/module.py +++ b/src/fairseq2/nn/utils/module.py @@ -7,8 +7,7 @@ from __future__ import annotations import re -from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence -from dataclasses import dataclass +from collections.abc import Callable, Iterator, Mapping, Sequence from itertools import chain from typing import Protocol, runtime_checkable @@ -17,9 +16,9 @@ from torch.nn import Module, Parameter from torch.nn.utils import remove_weight_norm # type: ignore[attr-defined] +from fairseq2.device import CPU, Device from fairseq2.gang import Gang from fairseq2.logging import log -from fairseq2.typing import CPU, Device @runtime_checkable @@ -238,7 +237,7 @@ def apply_to_parameters( """ if no_memo: memo = None - elif memo is None and recurse: + elif memo is None: memo = {} if recurse: @@ -275,7 +274,8 @@ def call_fn( setattr(module, param_name, new_param) - if (grad := param.grad) is not None: + grad = param.grad + if grad is not None: with torch.no_grad(): new_grad = call_fn(grad, requires_grad=grad.requires_grad) @@ -298,7 +298,7 @@ def freeze_parameters(module: Module | None, value: bool = True) -> None: def select_parameters( module: Module, names: Sequence[str], *, exclude: bool = False -) -> Iterable[tuple[str, Parameter]]: +) -> Iterator[tuple[str, Parameter]]: """Select the parameters of ``module`` and its descendant modules whose names match ``names``. @@ -310,7 +310,7 @@ def select_parameters( If ``True``, return the parameters that do not match ``names``. :returns: - An iterable of name-parameter tuples. + An iterator of name-parameter tuples. """ for name, param in module.named_parameters(): matched = any(name == pattern or re.match(pattern, name) for pattern in names) @@ -460,25 +460,10 @@ def load_state_dict( """Copy parameters and buffers from ``state_dict`` into ``module`` and its descendant modules. - This implementation internally calls :meth:`Module.load_state_dict()`, and also enforces that - ``state_dict`` does not contain any keys corresponding to descendants that are set to ``None`` - via :meth:`Module.register_module()`. + This implementation internally calls :meth:`Module.load_state_dict()`, and + also enforces that ``state_dict`` does not contain any keys corresponding to + descendants that are set to ``None`` via :meth:`Module.register_module()`. """ - # Key mapping - need_mapping = False - sample_key = list(state_dict.keys())[0] - if ( - sample_key.startswith("module.") - and not sample_key in module.state_dict().keys() - ): - mapped_key = sample_key[7:] - if mapped_key in module.state_dict().keys(): - need_mapping = True - - if need_mapping: - key_mapping = lambda key: key[7:] if key.startswith("module.") else key - state_dict = {key_mapping(key): value for key, value in state_dict.items()} - module.load_state_dict(state_dict, strict=strict) unexpected_keys = [] @@ -548,69 +533,3 @@ def _get_named_modules( if post_order: yield prefix, module - - -@dataclass(kw_only=True) -class ModuleSizeInfo: - """Holds the size information of a module.""" - - param_size: int = 0 - """The total size of all parameters.""" - - param_size_bytes: int = 0 - """The total size of all parameters, in bytes.""" - - trainable_param_size: int = 0 - """The total size of all trainable parameters.""" - - trainable_param_size_bytes: int = 0 - """The total size of all trainable parameters, in bytes.""" - - buffer_size: int = 0 - """The total size of all buffers.""" - - buffer_size_bytes: int = 0 - """The total size of all buffers, in bytes.""" - - total_size: int = 0 - """The total size of the module.""" - - total_size_bytes: int = 0 - """The total size of the module, in bytes.""" - - -def get_module_size(module: Module) -> ModuleSizeInfo: - """Return the size information of ``module`` and its descendant modules.""" - info = ModuleSizeInfo() - - for param in module.parameters(): - if param is None: - continue - - size = param.numel() - size_bytes = size * param.element_size() - - info.param_size += size - info.param_size_bytes += size_bytes - - if param.requires_grad: - info.trainable_param_size += size - info.trainable_param_size_bytes += size_bytes - - info.total_size += size - info.total_size_bytes += size_bytes - - for buffer in module.buffers(): - if buffer is None: - continue - - size = buffer.numel() - size_bytes = size * buffer.element_size() - - info.buffer_size += size - info.buffer_size_bytes += size_bytes - - info.total_size += size - info.total_size_bytes += size_bytes - - return info From ad6e3d4307ad12dddec3cdc9a85ecc87de21badb Mon Sep 17 00:00:00 2001 From: Artyom Kozhevnikov Date: Mon, 11 Aug 2025 10:49:02 -0700 Subject: [PATCH 3/3] rm custom stuff --- src/fairseq2/models/wav2vec2/asr/_config.py | 308 ++++++++++++++++++-- src/fairseq2/nn/utils/module.py | 86 +++++- 2 files changed, 363 insertions(+), 31 deletions(-) diff --git a/src/fairseq2/models/wav2vec2/asr/_config.py b/src/fairseq2/models/wav2vec2/asr/_config.py index c65739c1a..7d45e6b80 100644 --- a/src/fairseq2/models/wav2vec2/asr/_config.py +++ b/src/fairseq2/models/wav2vec2/asr/_config.py @@ -10,7 +10,8 @@ from typing import Final from fairseq2.context import RuntimeContext -from fairseq2.models.wav2vec2 import Wav2Vec2Config, Wav2Vec2EncoderConfig +from fairseq2.data import VocabularyInfo +from fairseq2.models.wav2vec2 import Wav2Vec2EncoderConfig WAV2VEC2_ASR_MODEL_FAMILY: Final = "wav2vec2_asr" @@ -25,7 +26,7 @@ class Wav2Vec2AsrConfig: encoder_config: Wav2Vec2EncoderConfig = field( default_factory=lambda: Wav2Vec2EncoderConfig( - feature_grad_scale=1.0, + feature_gradient_scale=1.0, dropout_p=0.0, attn_dropout_p=0.0, ffn_inner_dropout_p=0.1, @@ -33,13 +34,19 @@ class Wav2Vec2AsrConfig: ) """The configuration of the encoder.""" - target_vocab_size: int = 32 - """The size of the target vocabulary.""" + vocab_info: VocabularyInfo = field( + default_factory=lambda: VocabularyInfo( + size=32, unk_idx=3, bos_idx=0, eos_idx=2, pad_idx=1 + ) + ) + """The vocabulary information.""" final_dropout_p: float = 0.0 """The dropout probability on the output of the encoder.""" # Mask + mask_codebase: str = "fairseq2" + use_masking: bool = True """If ``True``, masks features as regularization.""" @@ -66,14 +73,15 @@ class Wav2Vec2AsrConfig: def register_wav2vec2_asr_configs(context: RuntimeContext) -> None: registry = context.get_config_registry(Wav2Vec2AsrConfig) + wav2vec2_asr_arch = registry.decorator - arch = registry.decorator + w2v2_encoder_registry = context.get_config_registry(Wav2Vec2EncoderConfig) - @arch("base_10h") + @wav2vec2_asr_arch("base_10h") def base_10h() -> Wav2Vec2AsrConfig: return Wav2Vec2AsrConfig() - @arch("base_100h") + @wav2vec2_asr_arch("base_100h") def base_100h() -> Wav2Vec2AsrConfig: config = base_10h() @@ -81,16 +89,12 @@ def base_100h() -> Wav2Vec2AsrConfig: return config - w2v2_registry = context.get_config_registry(Wav2Vec2Config) - - @arch("large_10h") + @wav2vec2_asr_arch("large_10h") def large_10h() -> Wav2Vec2AsrConfig: config = base_10h() - w2v2_config = w2v2_registry.get("large") - - config.encoder_config = w2v2_config.encoder_config - config.encoder_config.feature_grad_scale = 1.0 + config.encoder_config = w2v2_encoder_registry.get("large") + config.encoder_config.feature_gradient_scale = 1.0 config.encoder_config.dropout_p = 0.0 config.encoder_config.attn_dropout_p = 0.0 config.encoder_config.ffn_inner_dropout_p = 0.1 @@ -101,7 +105,7 @@ def large_10h() -> Wav2Vec2AsrConfig: return config - @arch("large_100h") + @wav2vec2_asr_arch("large_100h") def large_100h() -> Wav2Vec2AsrConfig: config = large_10h() @@ -110,14 +114,12 @@ def large_100h() -> Wav2Vec2AsrConfig: return config - @arch("large_lv60k_10h") + @wav2vec2_asr_arch("large_lv60k_10h") def large_lv60k_10h() -> Wav2Vec2AsrConfig: config = base_10h() - w2v2_config = w2v2_registry.get("large_lv60k") - - config.encoder_config = w2v2_config.encoder_config - config.encoder_config.feature_grad_scale = 1.0 + config.encoder_config = w2v2_encoder_registry.get("large_lv60k") + config.encoder_config.feature_gradient_scale = 1.0 config.encoder_config.dropout_p = 0.0 config.encoder_config.attn_dropout_p = 0.0 config.encoder_config.ffn_inner_dropout_p = 0.1 @@ -128,7 +130,7 @@ def large_lv60k_10h() -> Wav2Vec2AsrConfig: return config - @arch("large_lv60k_100h") + @wav2vec2_asr_arch("large_lv60k_100h") def large_lv60k_100h() -> Wav2Vec2AsrConfig: config = large_lv60k_10h() @@ -136,3 +138,267 @@ def large_lv60k_100h() -> Wav2Vec2AsrConfig: config.max_spatial_mask_prob = 0.55 return config + + @wav2vec2_asr_arch("300m_bib61") + def bib61_300m() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("large_lv60k") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 2475 + + return config + + @wav2vec2_asr_arch("300m_bib1143") + def bib1143_300m() -> Wav2Vec2AsrConfig: + config = bib61_300m() + config.vocab_info.size = 3335 + return config + + @wav2vec2_asr_arch("1b_bib61") + def bib61_1b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("1b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 2475 + + return config + + @wav2vec2_asr_arch("1b_llama_bib61") + def llama_bib61_1b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("1b_llama") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 2475 + + return config + + @wav2vec2_asr_arch("2b_bib61") + def bib61_2b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("2b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 2475 + + return config + + @wav2vec2_asr_arch("3b_bib61") + def bib61_3b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("3b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 2475 + + return config + + @wav2vec2_asr_arch("5b_bib61") + def bib61_5b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("5b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 2475 + + return config + + @wav2vec2_asr_arch("7b_bib61") + def bib61_7b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("7b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 2475 + + return config + + @wav2vec2_asr_arch("3.25b_bib61") + def higher_bib61_3b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("3.25b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 2475 + + return config + + @wav2vec2_asr_arch("5b_front51") + def front51_5b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("5b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 222 + + return config + + @wav2vec2_asr_arch("7b_front51") + def front51_7b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("7b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 222 + + return config + + @wav2vec2_asr_arch("1b_bib1143") + def bib1143_1b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("1b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 3335 + + return config + + @wav2vec2_asr_arch("3b_bib1143") + def bib1143_3b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("3b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 3335 + + return config + + @wav2vec2_asr_arch("5b_bib1143") + def bib1143_5b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("5b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 3335 # following bibfront1194's vocab size + + return config + + @wav2vec2_asr_arch("7b_bib1143") + def bib1143_7b() -> Wav2Vec2AsrConfig: + config = base_10h() + + config.encoder_config = w2v2_encoder_registry.get("7b") + config.encoder_config.feature_gradient_scale = 1.0 + config.encoder_config.dropout_p = 0.0 + config.encoder_config.attn_dropout_p = 0.0 + config.encoder_config.ffn_inner_dropout_p = 0.1 + config.encoder_config.layer_drop_p = 0.1 + + config.use_masking = False + config.max_temporal_mask_prob = 0.0 + config.max_spatial_mask_prob = 0.0 + config.vocab_info.size = 3335 + + return config + + @wav2vec2_asr_arch("7b_v3_tokenizer") + def v3_tokenizer_7b() -> Wav2Vec2AsrConfig: + config = bib1143_7b() + config.vocab_info.size = 9656 + return config diff --git a/src/fairseq2/nn/utils/module.py b/src/fairseq2/nn/utils/module.py index 0a69a00bf..a3ba5f0ce 100644 --- a/src/fairseq2/nn/utils/module.py +++ b/src/fairseq2/nn/utils/module.py @@ -7,7 +7,8 @@ from __future__ import annotations import re -from collections.abc import Callable, Iterator, Mapping, Sequence +from collections.abc import Callable, Iterable, Iterator, Mapping, Sequence +from dataclasses import dataclass from itertools import chain from typing import Protocol, runtime_checkable @@ -16,9 +17,9 @@ from torch.nn import Module, Parameter from torch.nn.utils import remove_weight_norm # type: ignore[attr-defined] -from fairseq2.device import CPU, Device from fairseq2.gang import Gang from fairseq2.logging import log +from fairseq2.typing import CPU, Device @runtime_checkable @@ -237,7 +238,7 @@ def apply_to_parameters( """ if no_memo: memo = None - elif memo is None: + elif memo is None and recurse: memo = {} if recurse: @@ -274,8 +275,7 @@ def call_fn( setattr(module, param_name, new_param) - grad = param.grad - if grad is not None: + if (grad := param.grad) is not None: with torch.no_grad(): new_grad = call_fn(grad, requires_grad=grad.requires_grad) @@ -298,7 +298,7 @@ def freeze_parameters(module: Module | None, value: bool = True) -> None: def select_parameters( module: Module, names: Sequence[str], *, exclude: bool = False -) -> Iterator[tuple[str, Parameter]]: +) -> Iterable[tuple[str, Parameter]]: """Select the parameters of ``module`` and its descendant modules whose names match ``names``. @@ -310,7 +310,7 @@ def select_parameters( If ``True``, return the parameters that do not match ``names``. :returns: - An iterator of name-parameter tuples. + An iterable of name-parameter tuples. """ for name, param in module.named_parameters(): matched = any(name == pattern or re.match(pattern, name) for pattern in names) @@ -460,9 +460,9 @@ def load_state_dict( """Copy parameters and buffers from ``state_dict`` into ``module`` and its descendant modules. - This implementation internally calls :meth:`Module.load_state_dict()`, and - also enforces that ``state_dict`` does not contain any keys corresponding to - descendants that are set to ``None`` via :meth:`Module.register_module()`. + This implementation internally calls :meth:`Module.load_state_dict()`, and also enforces that + ``state_dict`` does not contain any keys corresponding to descendants that are set to ``None`` + via :meth:`Module.register_module()`. """ module.load_state_dict(state_dict, strict=strict) @@ -533,3 +533,69 @@ def _get_named_modules( if post_order: yield prefix, module + + +@dataclass(kw_only=True) +class ModuleSizeInfo: + """Holds the size information of a module.""" + + param_size: int = 0 + """The total size of all parameters.""" + + param_size_bytes: int = 0 + """The total size of all parameters, in bytes.""" + + trainable_param_size: int = 0 + """The total size of all trainable parameters.""" + + trainable_param_size_bytes: int = 0 + """The total size of all trainable parameters, in bytes.""" + + buffer_size: int = 0 + """The total size of all buffers.""" + + buffer_size_bytes: int = 0 + """The total size of all buffers, in bytes.""" + + total_size: int = 0 + """The total size of the module.""" + + total_size_bytes: int = 0 + """The total size of the module, in bytes.""" + + +def get_module_size(module: Module) -> ModuleSizeInfo: + """Return the size information of ``module`` and its descendant modules.""" + info = ModuleSizeInfo() + + for param in module.parameters(): + if param is None: + continue + + size = param.numel() + size_bytes = size * param.element_size() + + info.param_size += size + info.param_size_bytes += size_bytes + + if param.requires_grad: + info.trainable_param_size += size + info.trainable_param_size_bytes += size_bytes + + info.total_size += size + info.total_size_bytes += size_bytes + + for buffer in module.buffers(): + if buffer is None: + continue + + size = buffer.numel() + size_bytes = size * buffer.element_size() + + info.buffer_size += size + info.buffer_size_bytes += size_bytes + + info.total_size += size + info.total_size_bytes += size_bytes + + return info