Skip to content

Commit f5d2be6

Browse files
authored
Merge pull request #20 from atlasia-ma/feat/replace-packing-with-length-grouping-and-hf-checkpoints
feat(train): switch from packing to group_by_length and enable HF checkpointing
2 parents 595a665 + 4a1f63c commit f5d2be6

6 files changed

Lines changed: 244 additions & 41 deletions

File tree

.env.example

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1,2 @@
1-
WANDB_API_KEY=your-key-here
1+
WANDB_API_KEY=your-key-here
2+
HUB_TOKEN=your-token-here

README.md

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,10 @@ Training reports to Weights & Biases by default.
3737
cp .env.example .env
3838
# fill in WANDB_API_KEY from https://wandb.ai/authorize
3939

40+
## Hugging face checkpointing
41+
42+
huggingface-cli login
43+
4044
## Training (requires GPU)
4145

4246
uv run darija-translator train

pyproject.toml

Lines changed: 1 addition & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -14,9 +14,7 @@ eval = [
1414
"sacrebleu>=2.6.0",
1515
]
1616
train = [
17-
"datasets",
18-
"trl",
19-
"unsloth",
17+
"unsloth>=2026.7",
2018
"wandb",
2119
"python-dotenv",
2220
]

src/darija_translator/config.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -49,3 +49,4 @@ class TrainConfig:
4949
output_dir: str = "lora_model"
5050
report_to: str = "wandb"
5151
wandb_project: str = "darija-translator"
52+
hub_model_id: str = "atlasia/edge-device-darija-translator"

src/darija_translator/train.py

Lines changed: 128 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,31 @@
1+
from unsloth.chat_templates import train_on_responses_only
2+
13
import os
24

5+
from datasets import load_dataset
6+
7+
from darija_translator.model import attach_lora, load_model_and_tokenizer
8+
from darija_translator.data import split_dataset
39
from trl import SFTConfig, SFTTrainer
4-
from unsloth.chat_templates import train_on_responses_only
510

6-
from darija_translator.config import TrainConfig
11+
from darija_translator.config import DataConfig, ModelConfig, TrainConfig
12+
13+
14+
def prepare_data(dataset_name: str,
15+
data_config: DataConfig,
16+
tokenizer,
17+
remove_columns: bool = True) -> tuple:
18+
dataset = load_dataset(dataset_name, split="train[:10]")
19+
# dataset = dataset.filter(is_darija_script)
20+
# dataset = dataset.map(lambda b: to_conversations(b, data_config),
21+
# batched=True)
22+
# dataset = dataset.map(lambda b: format_conversations(b, tokenizer),
23+
# batched=True)
24+
# dataset = dataset.filter(lambda ex: is_within_length(ex, data_config))
25+
if remove_columns:
26+
dataset = dataset.remove_columns(
27+
[c for c in dataset.column_names if c != "text"])
28+
return split_dataset(dataset, data_config)
729

830

931
def build_trainer(model, tokenizer, train_dataset, eval_dataset,
@@ -12,31 +34,41 @@ def build_trainer(model, tokenizer, train_dataset, eval_dataset,
1234
and "wandb" in config.report_to):
1335
os.environ["WANDB_PROJECT"] = config.wandb_project
1436

15-
trainer = SFTTrainer(
16-
model=model,
17-
tokenizer=tokenizer,
18-
train_dataset=train_dataset,
19-
eval_dataset=eval_dataset,
20-
args=SFTConfig(
21-
dataset_text_field="text",
22-
per_device_train_batch_size=config.per_device_train_batch_size,
23-
gradient_accumulation_steps=config.gradient_accumulation_steps,
24-
packing=config.packing,
25-
max_seq_length=config.max_seq_length,
26-
warmup_ratio=config.warmup_ratio,
27-
num_train_epochs=config.num_train_epochs,
28-
per_device_eval_batch_size=config.per_device_eval_batch_size,
29-
eval_strategy="epoch",
30-
learning_rate=config.learning_rate,
31-
logging_steps=config.logging_steps,
32-
optim=config.optim,
33-
weight_decay=config.weight_decay,
34-
lr_scheduler_type=config.lr_scheduler_type,
35-
seed=config.seed,
36-
report_to=config.report_to,
37-
# group_by_length=config.group_by_length,
38-
),
37+
sft_args = SFTConfig(
38+
dataset_text_field="text",
39+
per_device_train_batch_size=config.per_device_train_batch_size,
40+
gradient_accumulation_steps=config.gradient_accumulation_steps,
41+
packing=False,
42+
max_seq_length=config.max_seq_length,
43+
warmup_ratio=config.warmup_ratio,
44+
num_train_epochs=config.num_train_epochs,
45+
per_device_eval_batch_size=config.per_device_eval_batch_size,
46+
eval_strategy="epoch",
47+
learning_rate=config.learning_rate,
48+
logging_steps=config.logging_steps,
49+
optim=config.optim,
50+
weight_decay=config.weight_decay,
51+
lr_scheduler_type=config.lr_scheduler_type,
52+
seed=config.seed,
53+
report_to=config.report_to,
54+
output_dir=config.output_dir,
55+
save_strategy="steps", # Save checkpoints at step intervals
56+
save_steps=400,
57+
save_total_limit=3,
58+
push_to_hub=True, # Enable auto-uploading to HF
59+
hub_model_id=config.hub_model_id,
60+
hub_strategy="checkpoint",
61+
62+
# group_by_length=config.group_by_length,
3963
)
64+
sft_args.group_by_length = True
65+
66+
trainer = SFTTrainer(model=model,
67+
tokenizer=tokenizer,
68+
train_dataset=train_dataset,
69+
eval_dataset=eval_dataset,
70+
args=sft_args)
71+
4072
return train_on_responses_only(
4173
trainer,
4274
instruction_part="<|im_start|>user\n",
@@ -47,3 +79,73 @@ def build_trainer(model, tokenizer, train_dataset, eval_dataset,
4779
def save_model(model, tokenizer, config: TrainConfig):
4880
model.save_pretrained(config.output_dir)
4981
tokenizer.save_pretrained(config.output_dir)
82+
83+
84+
if __name__ == "__main__":
85+
data_config, model_config, train_config = DataConfig(), ModelConfig(
86+
), TrainConfig()
87+
model, tokenizer = load_model_and_tokenizer(model_config)
88+
model = attach_lora(model, model_config)
89+
train_dataset, eval_dataset = prepare_data(
90+
"atlasia/english-to-darija-arabic-script-formatted", data_config,
91+
tokenizer)
92+
trainer = build_trainer(model, tokenizer, train_dataset, eval_dataset,
93+
train_config)
94+
for i in range(min(10, len(trainer.train_dataset))):
95+
row = trainer.train_dataset[i]
96+
print(row)
97+
input_ids = row["input_ids"]
98+
labels = row["labels"]
99+
print(f"\n--- row {i} ---")
100+
print(f"input_ids: {input_ids}")
101+
print(f"labels: {labels}")
102+
print(f"decoded input: {tokenizer.decode(input_ids)}")
103+
print(
104+
f"decoded labels: {tokenizer.decode([lab for lab in labels if lab != -100])}"
105+
)
106+
107+
# # 1. Run your training
108+
# trainer.train()
109+
110+
# # 2. Push the highly optimized LoRA adapter to your repo
111+
# model.push_to_hub("your-username/darija-translator", token=True)
112+
# tokenizer.push_to_hub("your-username/darija-translator", token=True)
113+
114+
# # Merges LoRA weights back into the base structure and pushes the whole thing
115+
# model.push_to_hub_merged(
116+
# "your-username/darija-translator-merged",
117+
# tokenizer,
118+
# save_method="merged_16bit"
119+
# )
120+
# from unsloth import FastLanguageModel
121+
122+
# max_seq_length = 2048
123+
# dtype = None # None for auto detection. Float16 for Tesla T4/V100, Bfloat16 for Ampere+
124+
# load_in_4bit = True # Use True if you want to keep VRAM footprint low
125+
126+
# # 1. Load the base model and tokenizer
127+
# model, tokenizer = FastLanguageModel.from_pretrained(
128+
# model_name = "unsloth/llama-3-8b-Instruct", # Use whatever base model you started with
129+
# max_seq_length = max_seq_length,
130+
# dtype = dtype,
131+
# load_in_4bit = load_in_4bit,
132+
# )
133+
134+
# # 2. Layer your tiny checkpoint adapter right on top from the Hugging Face Hub
135+
# # You can reference specific checkpoint folders directly using the 'subfolder' argument!
136+
# model = FastLanguageModel.for_inference(model)
137+
# model.load_adapter(
138+
# "your-username/darija-translator",
139+
# subfolder="checkpoint-400" # Change this to checkpoint-800, checkpoint-1200, etc.
140+
# )
141+
142+
# # 3. Test your translator!
143+
# inputs = tokenizer(
144+
# [
145+
# "<|im_start|>system\nYou are a professional English to Darija translator.<|im_end|>\n<|im_start|>user\nHow are you doing today?<|im_end|>\n<|im_start|>assistant\n"
146+
# ],
147+
# return_tensors = "pt"
148+
# ).to("cuda")
149+
150+
# outputs = model.generate(**inputs, max_new_tokens = 64, use_cache = True)
151+
# print(tokenizer.batch_decode(outputs, skip_special_tokens=True)[0])

uv.lock

Lines changed: 108 additions & 11 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

0 commit comments

Comments
 (0)