Skip to content

Commit 1422a8e

Browse files
authored
Merge pull request #21 from atlasia-ma/fix/suppress-warnings-dataloader-workers
fix: suppress transformers warnings, add dataloader workers config
2 parents f5d2be6 + fda51e1 commit 1422a8e

4 files changed

Lines changed: 15 additions & 3 deletions

File tree

src/darija_translator/__init__.py

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,7 @@
1+
import warnings
2+
3+
warnings.filterwarnings("ignore",
4+
message=".*is_flash_linear_attention_available.*")
5+
import logging
6+
7+
logging.getLogger("transformers").setLevel(logging.ERROR)

src/darija_translator/cli.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,9 @@
1313
from darija_translator.evaluate import compute_translation_metrics, generate_translations
1414
from darija_translator.model import attach_lora, load_model_and_tokenizer
1515
from darija_translator.train import build_trainer, save_model
16-
# from dotenv import load_dotenv
16+
from dotenv import load_dotenv
17+
18+
load_dotenv()
1719

1820

1921
def prepare_data(dataset_name: str,

src/darija_translator/config.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -32,7 +32,8 @@ class ModelConfig:
3232

3333
@dataclass(frozen=True)
3434
class TrainConfig:
35-
per_device_train_batch_size: int = 16
35+
dataloader_num_workers: int = 4
36+
per_device_train_batch_size: int = 64
3637
gradient_accumulation_steps: int = 2
3738
per_device_eval_batch_size: int = 8
3839
num_train_epochs: int = 3

src/darija_translator/train.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@ def prepare_data(dataset_name: str,
1515
data_config: DataConfig,
1616
tokenizer,
1717
remove_columns: bool = True) -> tuple:
18-
dataset = load_dataset(dataset_name, split="train[:10]")
18+
dataset = load_dataset(dataset_name, split="train")
1919
# dataset = dataset.filter(is_darija_script)
2020
# dataset = dataset.map(lambda b: to_conversations(b, data_config),
2121
# batched=True)
@@ -36,6 +36,7 @@ def build_trainer(model, tokenizer, train_dataset, eval_dataset,
3636

3737
sft_args = SFTConfig(
3838
dataset_text_field="text",
39+
dataloader_num_workers=config.dataloader_num_workers,
3940
per_device_train_batch_size=config.per_device_train_batch_size,
4041
gradient_accumulation_steps=config.gradient_accumulation_steps,
4142
packing=False,
@@ -52,6 +53,7 @@ def build_trainer(model, tokenizer, train_dataset, eval_dataset,
5253
seed=config.seed,
5354
report_to=config.report_to,
5455
output_dir=config.output_dir,
56+
padding_free=False,
5557
save_strategy="steps", # Save checkpoints at step intervals
5658
save_steps=400,
5759
save_total_limit=3,

0 commit comments

Comments
 (0)