Skip to content

Commit 2afd115

Browse files
authored
Merge pull request #18 from atlasia-ma/fix/review-feedback-training
address review feedback
2 parents 3db3622 + 5cfd918 commit 2afd115

4 files changed

Lines changed: 14 additions & 11 deletions

File tree

pyproject.toml

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -11,14 +11,14 @@ dev = [
1111
"pytest>=9.1.1",
1212
]
1313
eval = [
14-
"sacrebleu>=2.4",
14+
"sacrebleu>=2.6.0",
1515
]
1616
train = [
17-
"datasets>=4.3.0",
18-
"trl>=0.24.0",
19-
"unsloth>=2026.6.9",
20-
"wandb>=0.19",
21-
"python-dotenv>=1.0",
17+
"datasets>=5.0.0",
18+
"trl>=1.7.1",
19+
"unsloth>=2026.7.2",
20+
"wandb>=0.28.0",
21+
"python-dotenv>=1.2.2",
2222
]
2323

2424
[project.scripts]

src/darija_translator/cli.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -59,11 +59,13 @@ def main():
5959
subparsers = parser.add_subparsers(required=True)
6060

6161
train_parser = subparsers.add_parser("train")
62-
train_parser.add_argument("--dataset", default="atlasia/darija_english")
62+
train_parser.add_argument("--dataset",
63+
default="atlasia/darija-english-combined")
6364
train_parser.set_defaults(func=run_train)
6465

6566
eval_parser = subparsers.add_parser("evaluate")
66-
eval_parser.add_argument("--dataset", default="atlasia/darija_english")
67+
eval_parser.add_argument("--dataset",
68+
default="atlasia/darija-english-combined")
6769
eval_parser.set_defaults(func=run_evaluate)
6870

6971
args = parser.parse_args()

src/darija_translator/evaluate.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -32,7 +32,7 @@ def generate_translations(model, tokenizer, sources: list[str],
3232
).to(model.device)
3333
input_len = inputs["input_ids"].shape[1]
3434
outputs = model.generate(**inputs,
35-
max_new_tokens=128,
35+
max_new_tokens=256,
3636
do_sample=False,
3737
use_cache=True)
3838
new_tokens = outputs[:, input_len:]

src/darija_translator/train.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,8 @@
88

99
def build_trainer(model, tokenizer, train_dataset, eval_dataset,
1010
config: TrainConfig):
11-
if config.report_to == "wandb" or (isinstance(config.report_to, list) and "wandb" in config.report_to):
11+
if config.report_to == "wandb" or (isinstance(config.report_to, list)
12+
and "wandb" in config.report_to):
1213
os.environ["WANDB_PROJECT"] = config.wandb_project
1314

1415
trainer = SFTTrainer(
@@ -33,7 +34,7 @@ def build_trainer(model, tokenizer, train_dataset, eval_dataset,
3334
lr_scheduler_type=config.lr_scheduler_type,
3435
seed=config.seed,
3536
report_to=config.report_to,
36-
group_by_length=config.group_by_length,
37+
# group_by_length=config.group_by_length,
3738
),
3839
)
3940
return train_on_responses_only(

0 commit comments

Comments
 (0)