File size: 7,878 Bytes
ce827ed
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
"""

Full fine-tune (not LoRA) of h2oai/h2o-danube3-500m-chat for email triage.



VARIANT EXPERIMENT -- candidate replacement for cipher-nano. SmolLM2 (49K

vocab, 135M/360M) shrinks to the right disk size but plateaus at weak

category/importance accuracy after multiple tuning attempts (LoRA vs full-FT,

epoch sweeps, data reshaping) -- a base-pretraining-quality ceiling, not a

tuning problem (see DEPLOYMENT.md). Qwen2.5-0.5B has the opposite problem:

strong base quality but a 151,936-token vocabulary that floors its disk size

around 340-400MB regardless of quantization, so it can't shrink into nano's

target range either.



Danube3-500M is a plain LlamaForCausalLM with a 32,000-token vocabulary --

much smaller than Qwen/Gemma, comparable to SmolLM2 -- while coming from a

more conventional larger-scale pretraining recipe (h2oai's Danube series).

Worth testing whether it breaks the small-vocab-means-weak-base pattern.



Same data/format/eval as the other nano candidates -- only the base model

differs.



Outputs:

    outputs/danube3-500m-full/model/    - full fine-tuned HF model



Usage:

    python train/train_danube3_500m_full.py

    python train/train_danube3_500m_full.py --epochs 3 --output_dir ./my_run

"""

import argparse
import inspect
import re
from pathlib import Path


def parse_args():
    parser = argparse.ArgumentParser(description="Full fine-tune Danube3-500M for email triage")
    parser.add_argument("--model_name", default="h2oai/h2o-danube3-500m-chat", help="Base HF model")
    parser.add_argument("--train_file", default="train.jsonl", help="Training JSONL")
    parser.add_argument("--val_file", default="val.jsonl", help="Validation JSONL")
    parser.add_argument("--output_dir", default="outputs/danube3-500m-full", help="Root output directory")
    parser.add_argument("--max_seq_length", type=int, default=2048)
    parser.add_argument("--epochs", type=int, default=3)
    parser.add_argument("--lr", type=float, default=5e-5)
    parser.add_argument("--per_device_batch", type=int, default=2)
    parser.add_argument("--gradient_accumulation", type=int, default=4)
    parser.add_argument("--warmup_ratio", type=float, default=0.1)
    parser.add_argument("--seed", type=int, default=3407)
    parser.add_argument("--packing", action="store_true", default=False, help="Pack multiple short examples per sequence (default on)")
    parser.add_argument("--no-packing", dest="packing", action="store_false")
    return parser.parse_args()


def main(args):
    from datasets import disable_caching, load_dataset
    from trl import SFTConfig, SFTTrainer
    from unsloth import FastLanguageModel, is_bfloat16_supported

    disable_caching()

    out_root = Path(args.output_dir)
    model_dir = out_root / "model"
    out_root.mkdir(parents=True, exist_ok=True)

    print(f"Loading {args.model_name} for FULL fine-tune (no LoRA, no quantization) ...")
    model, tokenizer = FastLanguageModel.from_pretrained(
        model_name=args.model_name,
        max_seq_length=args.max_seq_length,
        dtype=None,
        load_in_4bit=False,
        full_finetuning=True,
    )

    print(f"Loading datasets: {args.train_file}, {args.val_file}")
    train_ds = load_dataset("json", data_files=args.train_file, split="train")
    val_ds = load_dataset("json", data_files=args.val_file, split="train")

    def format_chat(example):
        # Danube3-500m-chat uses its own native format -- <|prompt|>...eos
        # for user turns, <|answer|>...eos for assistant turns, strictly
        # alternating, no system role (confirmed against the tokenizer's own
        # chat_template and vocab: ChatML tokens aren't even present).
        # Fold the system prompt into the first user turn's content.
        msgs = example["messages"]
        system_content = ""
        if msgs and msgs[0]["role"] == "system":
            system_content = msgs[0]["content"] + "\n\n"
            msgs = msgs[1:]

        parts = []
        first_user = True
        for msg in msgs:
            content = msg["content"]
            if msg["role"] == "user" and first_user:
                content = system_content + content
                first_user = False
            if msg["role"] == "user":
                parts.append(f"<|prompt|>{content.strip()}{tokenizer.eos_token}")
            else:
                parts.append(f"<|answer|>{content.strip()}{tokenizer.eos_token}")
        text = "".join(parts)
        return {"text": text}

    train_ds = train_ds.map(format_chat, remove_columns=train_ds.column_names)
    val_ds = val_ds.map(format_chat, remove_columns=val_ds.column_names)

    print(f"Train examples: {len(train_ds)}  Validation examples: {len(val_ds)}")

    config_params = inspect.signature(SFTConfig).parameters
    training_kwargs = dict(
        output_dir=str(model_dir),
        num_train_epochs=args.epochs,
        per_device_train_batch_size=args.per_device_batch,
        per_device_eval_batch_size=args.per_device_batch,
        gradient_accumulation_steps=args.gradient_accumulation,
        learning_rate=args.lr,
        warmup_ratio=args.warmup_ratio,
        lr_scheduler_type="cosine",
        optim="adamw_8bit",
        eval_steps=100,
        save_strategy="steps",
        save_steps=100,
        logging_steps=10,
        seed=args.seed,
        fp16=not is_bfloat16_supported(),
        bf16=is_bfloat16_supported(),
        load_best_model_at_end=True,
        metric_for_best_model="eval_loss",
        greater_is_better=False,
        report_to="none",
        dataset_text_field="text",
        packing=args.packing,
    )

    if "eval_strategy" in config_params:
        training_kwargs["eval_strategy"] = "steps"
    else:
        training_kwargs["evaluation_strategy"] = "steps"
    if "max_length" in config_params:
        training_kwargs["max_length"] = args.max_seq_length
    else:
        training_kwargs["max_seq_length"] = args.max_seq_length
    training_args = SFTConfig(**training_kwargs)

    trainer_kwargs = dict(
        model=model,
        train_dataset=train_ds,
        eval_dataset=val_ds,
        args=training_args,
    )
    trainer_params = inspect.signature(SFTTrainer).parameters
    if "processing_class" in trainer_params:
        trainer_kwargs["processing_class"] = tokenizer
    else:
        trainer_kwargs["tokenizer"] = tokenizer

    _orig_convert_tokens_to_ids = tokenizer.convert_tokens_to_ids
    _sentinel_re = re.compile(r"^<([A-Z]+)_TOKEN>$")

    def _convert_tokens_to_ids_patched(token):
        match = _sentinel_re.match(token) if isinstance(token, str) else None
        if match:
            real_id = getattr(tokenizer, f"{match.group(1).lower()}_token_id", None)
            if real_id is not None:
                return real_id
        return _orig_convert_tokens_to_ids(token)

    _orig_prepare_dataset = SFTTrainer._prepare_dataset

    def _prepare_dataset_patched(self, dataset, processing_class, ds_args, *rest, **kw):
        ds_args.dataset_num_proc = None
        return _orig_prepare_dataset(self, dataset, processing_class, ds_args, *rest, **kw)

    SFTTrainer._prepare_dataset = _prepare_dataset_patched

    tokenizer.convert_tokens_to_ids = _convert_tokens_to_ids_patched
    try:
        trainer = SFTTrainer(**trainer_kwargs)
    finally:
        tokenizer.convert_tokens_to_ids = _orig_convert_tokens_to_ids
        SFTTrainer._prepare_dataset = _orig_prepare_dataset

    print("Starting training...")
    trainer.train()

    print(f"Saving full fine-tuned model to {model_dir}")
    model.save_pretrained(model_dir)
    tokenizer.save_pretrained(model_dir)

    print("Done.")


if __name__ == "__main__":
    args = parse_args()
    main(args)