Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 6 additions & 5 deletions src/fairseq2/recipes/wav2vec2/asr/_train.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
#
Expand All @@ -8,15 +8,13 @@

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
Expand Down Expand Up @@ -60,17 +58,20 @@
)
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


@dataclass(kw_only=True)
class Wav2Vec2AsrTrainConfig:
is_june_17_codebase: bool = True
model: ModelSection = field(
default_factory=lambda: ModelSection(family="wav2vec2_asr", arch="base_10h")
)
Expand Down Expand Up @@ -301,7 +302,7 @@

# Make sure that the final projection layer is instantiated along with
# the pretrained parameters if it was on the meta device.
if gangs.dp.rank == 0:

Check failure on line 305 in src/fairseq2/recipes/wav2vec2/asr/_train.py

View workflow job for this annotation

GitHub Actions / Lint Python / Lint

Variable "tp" is not valid as a type
to_device(module, gangs.root.device)

try:
Expand Down
Loading