From 865725e45e6763ea7c7104b252866efed067cf5c Mon Sep 17 00:00:00 2001 From: gilkeren Date: Fri, 22 Aug 2025 13:06:41 -0700 Subject: [PATCH 1/4] debug commit --- .../assets/cards/datasets/librilight.yaml | 12 +++++ src/fairseq2/cli/_main.py | 50 ++++++++++--------- src/fairseq2/nn/utils/module.py | 21 ++++++-- src/fairseq2/recipes/_trainer.py | 39 ++++++++------- src/fairseq2/recipes/_validator.py | 32 ++++++++---- src/fairseq2/recipes/asr/_common.py | 32 ++++++++---- src/fairseq2/recipes/wav2vec2/asr/_train.py | 20 +++++--- 7 files changed, 138 insertions(+), 68 deletions(-) diff --git a/src/fairseq2/assets/cards/datasets/librilight.yaml b/src/fairseq2/assets/cards/datasets/librilight.yaml index 2f8ba52a7..415990fb9 100644 --- a/src/fairseq2/assets/cards/datasets/librilight.yaml +++ b/src/fairseq2/assets/cards/datasets/librilight.yaml @@ -7,3 +7,15 @@ name: librilight_asr_10h dataset_family: generic_asr data: /checkpoint/mms/shared/assets/debug/librilight + +--- + +name: librilight_asr_10h_unlabeled +dataset_family: generic_speech +data: /checkpoint/mms/shared/assets/debug/librilight + +--- + +name: librilight_asr_10h_inference +dataset_family: speech_inference +data: /checkpoint/mms/shared/assets/debug/librilight diff --git a/src/fairseq2/cli/_main.py b/src/fairseq2/cli/_main.py index c03940a0c..0e3492d18 100644 --- a/src/fairseq2/cli/_main.py +++ b/src/fairseq2/cli/_main.py @@ -8,47 +8,47 @@ import os import sys -from signal import SIG_DFL, SIGINT, raise_signal, signal +from signal import raise_signal, SIG_DFL, SIGINT, signal import torch -from torch.cuda import OutOfMemoryError from fairseq2 import setup_fairseq2 -from fairseq2.cli.utils.rich import create_rich_progress_reporter -from fairseq2.error import ContractError, InternalError -from fairseq2.extensions import ExtensionError -from fairseq2.logging import LoggingSetupError, log -from fairseq2.setup import SetupError -from fairseq2.utils.env import InvalidEnvironmentVariableError, get_rank # isort: split from fairseq2.cli._logging import setup_logging from fairseq2.cli._setup import setup_cli +from fairseq2.cli.utils.rich import create_rich_progress_reporter +from fairseq2.error import ContractError, InternalError +from fairseq2.extensions import ExtensionError +from fairseq2.logging import log, LoggingSetupError +from fairseq2.setup import SetupError +from fairseq2.utils.env import get_rank, InvalidEnvironmentVariableError +from torch.cuda import OutOfMemoryError def main() -> None: """Runs the command line fairseq2 program.""" exit_code = 1 - try: - exit_code = _run() - except KeyboardInterrupt: - log.info("Command canceled!") + # try: + exit_code = _run() + # except KeyboardInterrupt: + # log.info("Command canceled!") - signal(SIGINT, SIG_DFL) + # signal(SIGINT, SIG_DFL) - raise_signal(SIGINT) - except OutOfMemoryError: - s = torch.cuda.memory_summary() + # raise_signal(SIGINT) + # except OutOfMemoryError: + # s = torch.cuda.memory_summary() - log.exception("CUDA out of memory. See logged memory stats.\n{}", s) - except InternalError: - log.exception("Command failed with an unexpected internal error. Please file a bug report.") # fmt: skip - except ContractError: - log.exception("Command failed with an unexpected internal error caused by an extension. Please file a bug report to the corresponding extension author.") # fmt: skip - except Exception: - log.exception("Command failed with an unexpected error. See the logged stack trace for details.") # fmt: skip + # log.exception("CUDA out of memory. See logged memory stats.\n{}", s) + # except InternalError: + # log.exception("Command failed with an unexpected internal error. Please file a bug report.") # fmt: skip + # except ContractError: + # log.exception("Command failed with an unexpected internal error caused by an extension. Please file a bug report to the corresponding extension author.") # fmt: skip + # except Exception: + # log.exception("Command failed with an unexpected error. See the logged stack trace for details.") # fmt: skip sys.exit(exit_code) @@ -84,3 +84,7 @@ def _run() -> int: return 1 return cli.run(context) + + +if __name__ == "__main__": + main() diff --git a/src/fairseq2/nn/utils/module.py b/src/fairseq2/nn/utils/module.py index a3ba5f0ce..962da59aa 100644 --- a/src/fairseq2/nn/utils/module.py +++ b/src/fairseq2/nn/utils/module.py @@ -13,13 +13,13 @@ from typing import Protocol, runtime_checkable import torch -from torch import Tensor -from torch.nn import Module, Parameter -from torch.nn.utils import remove_weight_norm # type: ignore[attr-defined] from fairseq2.gang import Gang from fairseq2.logging import log from fairseq2.typing import CPU, Device +from torch import Tensor +from torch.nn import Module, Parameter +from torch.nn.utils import remove_weight_norm # type: ignore[attr-defined] @runtime_checkable @@ -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/_trainer.py b/src/fairseq2/recipes/_trainer.py index b7cff3ed1..bdee11c67 100644 --- a/src/fairseq2/recipes/_trainer.py +++ b/src/fairseq2/recipes/_trainer.py @@ -9,16 +9,10 @@ from abc import ABC, abstractmethod from contextlib import nullcontext from enum import Enum -from typing import Final, Generic, Mapping, TypeVar, final +from typing import Final, final, Generic, Mapping, TypeVar import torch import torch.distributed -from rich.pretty import pretty_repr -from torch import Tensor -from torch.cuda import OutOfMemoryError -from torch.optim import Optimizer -from torch.profiler import record_function -from typing_extensions import override from fairseq2.checkpoint import ( CheckpointError, @@ -30,21 +24,14 @@ from fairseq2.datasets import DataReader, DataReadError from fairseq2.device import SupportsDeviceTransfer from fairseq2.error import InternalError, InvalidOperationError -from fairseq2.gang import GangError, Gangs, broadcast_flag +from fairseq2.gang import broadcast_flag, GangError, Gangs from fairseq2.logging import log from fairseq2.metrics import Mean, MetricBag, MetricBagError, MetricDescriptor from fairseq2.metrics.recorders import MetricRecorder, MetricRecordError from fairseq2.nn.utils.gradient import check_gradient_norms, normalize_gradients from fairseq2.optim import DynamicLossScaler -from fairseq2.optim.lr_scheduler import LRScheduler, get_effective_lr +from fairseq2.optim.lr_scheduler import get_effective_lr, LRScheduler from fairseq2.profilers import Profiler -from fairseq2.typing import CPU, ContextManager, DataType -from fairseq2.utils.device_stat import DeviceStatTracker -from fairseq2.utils.gc import GarbageCollector -from fairseq2.utils.progress import ProgressReporter, ProgressTask -from fairseq2.utils.rng import RngBag -from fairseq2.utils.state import Stateful -from fairseq2.utils.stopwatch import Stopwatch # isort: split @@ -59,6 +46,19 @@ from fairseq2.recipes._model import Model from fairseq2.recipes._recipe import Recipe, RecipeStopException from fairseq2.recipes._validator import Validator +from fairseq2.typing import ContextManager, CPU, DataType +from fairseq2.utils.device_stat import DeviceStatTracker +from fairseq2.utils.gc import GarbageCollector +from fairseq2.utils.progress import ProgressReporter, ProgressTask +from fairseq2.utils.rng import RngBag +from fairseq2.utils.state import Stateful +from fairseq2.utils.stopwatch import Stopwatch +from rich.pretty import pretty_repr +from torch import Tensor +from torch.cuda import OutOfMemoryError +from torch.optim import Optimizer +from torch.profiler import record_function +from typing_extensions import override BatchT_contra = TypeVar( "BatchT_contra", bound=SupportsDeviceTransfer, contravariant=True @@ -718,7 +718,12 @@ def _do_run_step(self, progress_task: ProgressTask) -> _TrainerState: batch = batches.pop() try: - batch.to(gangs.root.device) + try: + batch.to(gangs.root.device) + except Exception as e: + log.info(f"{gangs.root.device=}") + log.info(f"{batch=}") + raise e with self._maybe_no_sync(batch_nr, num_batches): with record_function(f"step_{step_nr}_{batch_nr}_forward"): diff --git a/src/fairseq2/recipes/_validator.py b/src/fairseq2/recipes/_validator.py index 6fb712abb..6fbc8e44d 100644 --- a/src/fairseq2/recipes/_validator.py +++ b/src/fairseq2/recipes/_validator.py @@ -6,15 +6,14 @@ from __future__ import annotations +import socket + from abc import ABC, abstractmethod from collections.abc import Sequence from contextlib import nullcontext -from typing import Generic, TypeVar, final +from typing import final, Generic, TypeVar import torch -from torch import Tensor -from torch.profiler import record_function -from typing_extensions import override from fairseq2.checkpoint import CheckpointError, CheckpointManager, CheckpointSaveError from fairseq2.datasets import DataReader, DataReadError @@ -25,17 +24,20 @@ from fairseq2.metrics import MetricBagError, MetricDescriptor from fairseq2.metrics.recorders import MetricRecorder, MetricRecordError from fairseq2.profilers import Profiler -from fairseq2.typing import CPU, ContextManager, DataType -from fairseq2.utils.device_stat import DeviceStatTracker -from fairseq2.utils.progress import ProgressReporter, ProgressTask -from fairseq2.utils.rng import RngBag -from fairseq2.utils.stopwatch import Stopwatch # isort: split from fairseq2.recipes._error import RecipeError, UnitError from fairseq2.recipes._evaluator import EvalUnit from fairseq2.recipes._metrics import extend_batch_metrics +from fairseq2.typing import ContextManager, CPU, DataType +from fairseq2.utils.device_stat import DeviceStatTracker +from fairseq2.utils.progress import ProgressReporter, ProgressTask +from fairseq2.utils.rng import RngBag +from fairseq2.utils.stopwatch import Stopwatch +from torch import Tensor +from torch.profiler import record_function +from typing_extensions import override class Validator(ABC): @@ -243,7 +245,14 @@ def _run_unit( f"The {s} unit has failed. See the nested exception for details." ) from ex + machine_name = socket.gethostname() + if machine_name.startswith("devvm"): + _max_num_valid_steps = 5 + else: + _max_num_valid_steps = 50000000000 + c = 0 while not eod: + log.info(f"s1: Running validation step {c}.") try: self._checkpoint_manager.maybe_complete_async_checkpoint() except CheckpointSaveError as ex: @@ -252,10 +261,13 @@ def _run_unit( ) from ex batches = self._read_next_batches(unit, data_reader) - if batches is None: + log.info(f"s2: Read batches step {c}.") + if batches is None or c == _max_num_valid_steps: eod = True else: self._run_step(unit, batches, progress_task) + log.info(f"s7: Done with step {c}.") + c += 1 with self._compute_watch: with record_function("finalize"): diff --git a/src/fairseq2/recipes/asr/_common.py b/src/fairseq2/recipes/asr/_common.py index b79d7c8a9..d501d68c3 100644 --- a/src/fairseq2/recipes/asr/_common.py +++ b/src/fairseq2/recipes/asr/_common.py @@ -7,20 +7,23 @@ from __future__ import annotations import math -from typing import Any, Dict, TextIO, final +import re +from typing import Any, Dict, final, TextIO import torch -from torch import Tensor -from typing_extensions import override from fairseq2.data.text.tokenizers import TextTokenDecoder, TextTokenizer from fairseq2.gang import Gang + +from fairseq2.logging import log from fairseq2.metrics import Mean 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 from fairseq2.recipes import BaseMetricBag, Model, UnitError +from torch import Tensor +from typing_extensions import override @final @@ -41,8 +44,10 @@ def __init__(self, model: Model, scorer: AsrScorer | None = None) -> None: def __call__( self, batch: Seq2SeqBatch, metric_bag: AsrMetricBag ) -> tuple[Tensor, int]: + log.info(f"s3: calling forward") output = self._forward(batch) + log.info(f"s4: calling loss") loss, extra_metrics = output.compute_loss( batch.target_seqs, batch.target_padding_mask ) @@ -53,8 +58,10 @@ def __call__( metric_bag.update_extra_metrics(batch, extra_metrics) + log.info(f"s5: calling scorer") if self._scorer is not None: self._scorer(batch, output, metric_bag) + log.info(f"s6: done scorer") return loss, batch.batch_size @@ -117,11 +124,18 @@ def __call__( refs = [self._text_decoder(s) for s in ref_seqs] hyps = [self._text_decoder(s) for s in hyp_seqs] + for i, (r, h) in enumerate(zip(refs, hyps)): + # if torch.rand([]) < 0.01 or bool(re.search(r"[\u0590-\u05FF]", r)): + if "lang" in batch.example: + log.info(f"Lang: {batch.example['lang'][i]}") + log.info(f"Reference: {r}") + log.info(f"Hypothesis: {h}") + metric_bag.wer.update( refs, ref_seqs, ref_padding_mask, hyps, hyp_seqs, hyp_padding_mask ) - metric_bag.bleu.update(refs, hyps) + # metric_bag.bleu.update(refs, hyps) try: # Dump references. @@ -161,11 +175,11 @@ 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="13a", device=self.device), - persistent=False, - ) + # self.register_metric( + # "bleu", + # BleuMetric(tokenizer="13a", device=self.device), + # persistent=False, + # ) @torch.inference_mode() def update_ctc_loss(self, batch: Seq2SeqBatch, loss: Tensor) -> None: diff --git a/src/fairseq2/recipes/wav2vec2/asr/_train.py b/src/fairseq2/recipes/wav2vec2/asr/_train.py index 3595d9114..5823e3e98 100644 --- a/src/fairseq2/recipes/wav2vec2/asr/_train.py +++ b/src/fairseq2/recipes/wav2vec2/asr/_train.py @@ -7,18 +7,18 @@ from __future__ import annotations import re +import socket from copy import deepcopy + from dataclasses import dataclass, field from pathlib import Path -from typing import Literal, cast, final +from typing import cast, final, Literal import torch -from torch import Tensor -from typing_extensions import override from fairseq2.context import RuntimeContext from fairseq2.datasets import LengthBatching, SyncMode -from fairseq2.datasets.asr import GENERIC_ASR_DATASET_FAMILY, AsrDataset +from fairseq2.datasets.asr import AsrDataset, GENERIC_ASR_DATASET_FAMILY from fairseq2.datasets.speech import ManifestDatasetInterface, SpeechReadOptions from fairseq2.gang import Gang, GangError from fairseq2.logging import log @@ -62,13 +62,15 @@ ) from fairseq2.recipes.utils.log import log_model from fairseq2.recipes.wav2vec2.batch_weighted_datareader import ( - MIXTURE_DATASET_FAMILY, BatchMixtureDataset, + MIXTURE_DATASET_FAMILY, ) from fairseq2.typing import CPU from fairseq2.utils.rng import manual_seed from fairseq2.utils.structured import structure from fairseq2.utils.validation import validate +from torch import Tensor +from typing_extensions import override def _strict_name(s: str) -> str: @@ -134,6 +136,7 @@ class Wav2Vec2AsrTrainConfig: validate_after_n_steps=10_000, validate_every_n_steps=1_000, publish_metrics_every_n_steps=200, + keep_last_n_checkpoints=1, ) ) @@ -310,7 +313,12 @@ def load_wav2vec2_asr_trainer( # If we start the training with an empty ASR model, use the weights of a # pretrained wav2vec 2.0 model. - if model.is_empty_initialized and config.pretrained_encoder.name: + machine_name = socket.gethostname() + if ( + model.is_empty_initialized + and config.pretrained_encoder.name + and not machine_name.startswith("devvm") + ): tp = AsrModel if config.pretrained_encoder_is_ctc else Wav2Vec2Model pt_model = load_reference_model( tp, From 8d4d97e8eefaa8da1f085a3af0709db846b19790 Mon Sep 17 00:00:00 2001 From: gilkeren Date: Tue, 2 Sep 2025 20:05:26 -0700 Subject: [PATCH 2/4] freqs hack --- src/fairseq2/nn/utils/module.py | 30 ++++++++++++++++++++++++++++-- 1 file changed, 28 insertions(+), 2 deletions(-) diff --git a/src/fairseq2/nn/utils/module.py b/src/fairseq2/nn/utils/module.py index 962da59aa..abce3c362 100644 --- a/src/fairseq2/nn/utils/module.py +++ b/src/fairseq2/nn/utils/module.py @@ -210,7 +210,12 @@ def collect_tensors(m: Module) -> None: # Do not memoize. No need anyways, and would also break the sync between the # traversed tensors and the iterator. - apply_to_parameters(target_module, lambda _: next(it), no_memo=True) + # apply_to_parameters( + # target_module, lambda _: next(it), no_memo=True, skip_freqs=True + # ) + apply_to_parameters( + target_module, lambda _: next(it), no_memo=True, skip_freqs=False + ) def apply_to_parameters( @@ -220,6 +225,7 @@ def apply_to_parameters( recurse: bool = True, memo: dict[Tensor, Tensor] | None = None, no_memo: bool = False, + skip_freqs: bool = False, ) -> None: """Apply ``fn`` to the parameters and buffers of ``module``. @@ -245,7 +251,12 @@ def apply_to_parameters( for child in module.children(): if child is not None: apply_to_parameters( - child, fn, recurse=recurse, memo=memo, no_memo=no_memo + child, + fn, + recurse=recurse, + memo=memo, + no_memo=no_memo, + skip_freqs=skip_freqs, ) def call_fn( @@ -273,6 +284,12 @@ def call_fn( with torch.no_grad(): new_param = call_fn(param, is_param=True, requires_grad=param.requires_grad) + if param.shape != new_param.shape: + log.warning( + f"The shape of {param_name} changed from {param.shape} to {new_param.shape}." + ) + continue + setattr(module, param_name, new_param) if (grad := param.grad) is not None: @@ -285,6 +302,15 @@ def call_fn( if buffer is None: continue + if skip_freqs and buffer_name == "freqs": + log.warning( + f"The `freqs` buffer of `module` was not updated :{buffer_name}." + ) + target = call_fn(buffer) + log.info(f"{target=} {buffer=}") + setattr(module, buffer_name, buffer) + continue + setattr(module, buffer_name, call_fn(buffer)) From 49449e33fa0842a9663f3431336448769a632d96 Mon Sep 17 00:00:00 2001 From: gilkeren Date: Mon, 15 Sep 2025 13:41:23 -0700 Subject: [PATCH 3/4] Matt's neural LM fusion --- src/fairseq2/models/llama/_config.py | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) diff --git a/src/fairseq2/models/llama/_config.py b/src/fairseq2/models/llama/_config.py index 6d73665f9..118249e3e 100644 --- a/src/fairseq2/models/llama/_config.py +++ b/src/fairseq2/models/llama/_config.py @@ -224,6 +224,25 @@ def llama3_1_8b() -> LLaMAConfig: return config + @arch("llama3_1_8b_v4_tokenizer") + def llama3_1_8b_v4_tokenizer() -> LLaMAConfig: + config = llama3_1_8b() + config.vocab_size = 9812 + config.pad_idx = 1 + config.model_dim = 2048 + config.tie_embeddings = True # remapped from tied_embeddings + config.ffn_inner_dim = 2048 * 4 + config.ffn_inner_dim_multiplier = 1.5 + config.ffn_inner_dim_to_multiple = ( + 256 # remapped from ffn_inner_dim_multiple_of + ) + config.num_attn_heads = 32 + config.num_key_value_heads = 8 + config.num_layers = 16 + config.use_scaled_rope = True + config.rope_scaling.factor = 32.0 # renamed from rope_scale.factor + return config + @arch("llama3_1_70b") def llama3_1_70b() -> LLaMAConfig: config = llama3_70b() From be3b40b8ab456779995a7058d387fd2c0d444bf2 Mon Sep 17 00:00:00 2001 From: gilkeren Date: Tue, 21 Oct 2025 11:10:41 -0700 Subject: [PATCH 4/4] v7 stuff --- src/fairseq2/models/wav2vec2/asr/_config.py | 48 +++++++++++++++++++++ 1 file changed, 48 insertions(+) diff --git a/src/fairseq2/models/wav2vec2/asr/_config.py b/src/fairseq2/models/wav2vec2/asr/_config.py index b76b7cd30..86e189450 100644 --- a/src/fairseq2/models/wav2vec2/asr/_config.py +++ b/src/fairseq2/models/wav2vec2/asr/_config.py @@ -500,3 +500,51 @@ def v4_tokenizer_updated_300m() -> Wav2Vec2AsrConfig: pad_idx=1, ) return config + + @wav2vec2_asr_arch("7b_v7_tokenizer") + def v7_tokenizer_7b() -> Wav2Vec2AsrConfig: + config = bib1143_7b() + config.vocab_info = VocabularyInfo( + size=9818, + unk_idx=3, + bos_idx=0, + eos_idx=2, + pad_idx=1, + ) + return config + + @wav2vec2_asr_arch("300m_v7_tokenizer") + def v7_tokenizer_300m() -> Wav2Vec2AsrConfig: + config = bib1143_300m() + config.vocab_info = VocabularyInfo( + size=9818, + unk_idx=3, + bos_idx=0, + eos_idx=2, + pad_idx=1, + ) + return config + + @wav2vec2_asr_arch("1b_v7_tokenizer") + def v7_tokenizer_1b() -> Wav2Vec2AsrConfig: + config = bib1143_1b() + config.vocab_info = VocabularyInfo( + size=9818, + unk_idx=3, + bos_idx=0, + eos_idx=2, + pad_idx=1, + ) + return config + + @wav2vec2_asr_arch("3b_v7_tokenizer") + def v7_tokenizer_3b() -> Wav2Vec2AsrConfig: + config = bib1143_3b() + config.vocab_info = VocabularyInfo( + size=9818, + unk_idx=3, + bos_idx=0, + eos_idx=2, + pad_idx=1, + ) + return config