Skip to content

Commit d1c98a8

Browse files
committed
Lint
1 parent 4c95061 commit d1c98a8

6 files changed

Lines changed: 41 additions & 13 deletions

File tree

src/fairseq2/models/opt/__init__.py

Lines changed: 3 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -6,16 +6,10 @@
66

77
from __future__ import annotations
88

9-
from fairseq2.models.opt.interop import (
10-
convert_opt_state_dict as convert_opt_state_dict,
11-
)
129
from fairseq2.models.opt.config import OPT_MODEL_FAMILY as OPT_MODEL_FAMILY
1310
from fairseq2.models.opt.config import OPTConfig as OPTConfig
14-
from fairseq2.models.opt.config import (
15-
register_opt_configs as register_opt_configs,
16-
)
11+
from fairseq2.models.opt.config import register_opt_configs as register_opt_configs
1712
from fairseq2.models.opt.factory import OPTFactory as OPTFactory
18-
from fairseq2.models.opt.factory import (
19-
create_opt_model as create_opt_model,
20-
)
13+
from fairseq2.models.opt.factory import create_opt_model as create_opt_model
2114
from fairseq2.models.opt.hub import get_opt_model_hub as get_opt_model_hub
15+
from fairseq2.models.opt.interop import convert_opt_state_dict as convert_opt_state_dict

src/fairseq2/models/opt/factory.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,6 @@
3333
Linear,
3434
PositionEncoder,
3535
Projection,
36-
RMSNorm,
3736
StandardEmbedding,
3837
StandardLayerNorm,
3938
)
@@ -156,7 +155,9 @@ def create_self_attention(self) -> MultiheadAttention:
156155
def create_ffn(self) -> FeedForwardNetwork:
157156
config = self._config
158157

159-
return StandardFeedForwardNetwork(config.model_dim, config.ffn_inner_dim, bias=True)
158+
return StandardFeedForwardNetwork(
159+
config.model_dim, config.ffn_inner_dim, bias=True
160+
)
160161

161162
def create_layer_norm(self) -> LayerNorm:
162163
config = self._config

src/fairseq2/models/opt/hub.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,6 @@
1111

1212
# isort: split
1313

14-
from fairseq2.models.opt.config import OPTConfig, OPT_MODEL_FAMILY
14+
from fairseq2.models.opt.config import OPT_MODEL_FAMILY, OPTConfig
1515

1616
get_opt_model_hub = ModelHubAccessor(OPT_MODEL_FAMILY, TransformerLM, OPTConfig)

src/fairseq2/models/opt/interop.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,6 @@
1212

1313
from fairseq2.models.opt.config import OPTConfig
1414

15-
1615
_OPT_HG_KEY_MAP = {
1716
# fmt: off
1817
r"^model\.decoder\.embed_tokens\.": r"decoder_frontend.embed.",

tests/unit/models/opt/__init__.py

Whitespace-only changes.
Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,34 @@
1+
# Copyright (c) Meta Platforms, Inc. and affiliates.
2+
# All rights reserved.
3+
#
4+
# This source code is licensed under the BSD-style license found in the
5+
# LICENSE file in the root directory of this source tree.
6+
7+
from __future__ import annotations
8+
9+
import torch
10+
11+
from fairseq2.models.opt.config import OPTConfig
12+
from fairseq2.models.opt.factory import OPTFactory
13+
from fairseq2.nn import BatchLayout
14+
15+
16+
class TestOptFactory:
17+
def test_opt_factory(self) -> None:
18+
"""Sanity check for factory + forward on a small model."""
19+
config = OPTConfig() # by default opt-125m
20+
config.num_layers = 2
21+
config.vocab_size = 258
22+
config.model_dim = 24
23+
config.ffn_inner_dim = 48
24+
25+
factory = OPTFactory(config)
26+
27+
model = factory.create_model()
28+
29+
model.eval()
30+
31+
_ = model.forward(
32+
seqs=torch.randint(0, config.vocab_size, (2, 10)),
33+
seqs_layout=BatchLayout((2, 10), [10, 10]),
34+
)

0 commit comments

Comments
 (0)