import argparse import hashlib import json from pathlib import Path import pandas as pd from transformers import AutoTokenizer def render_prompt(url): return tokenizer.apply_chat_template([{'role': 'user', 'content': INSTRUCTION + str(url)}], tokenize=False, add_generation_prompt=True, enable_thinking=False) def encode_url_prompt(url): text = str(url) def encode(text): return tokenizer.encode(render_prompt(text), add_special_tokens=False) prompt_ids = encode(text) original_length = len(prompt_ids) if original_length <= MAX_PROMPT_TOKENS: return (prompt_ids, False, original_length) url_ids = tokenizer.encode(text, add_special_tokens=False) while len(prompt_ids) > MAX_PROMPT_TOKENS: excess = len(prompt_ids) - MAX_PROMPT_TOKENS keep = max(0, len(url_ids) - excess - 8) if keep >= len(url_ids): raise RuntimeError('Truncation did not make progress.') url_ids = url_ids[:keep] text = tokenizer.decode(url_ids, skip_special_tokens=False, clean_up_tokenization_spaces=False) prompt_ids = encode(text) if not url_ids and len(prompt_ids) > MAX_PROMPT_TOKENS: raise ValueError('Instructions alone exceed the token budget.') return (prompt_ids, True, original_length) def training_collator(examples): import torch longest = max((len(item['input_ids']) for item in examples)) batch = {'input_ids': [], 'attention_mask': [], 'labels': []} for item in examples: padding = longest - len(item['input_ids']) batch['input_ids'].append(item['input_ids'] + [tokenizer.pad_token_id] * padding) batch['attention_mask'].append(item['attention_mask'] + [0] * padding) batch['labels'].append(item['labels'] + [-100] * padding) return {key: torch.tensor(values, dtype=torch.long) for key, values in batch.items()} def read_config(directory, name): return json.loads( (Path(directory) / name).read_text(encoding="utf-8") ) def ordered_hash(frame): digest = hashlib.sha256() columns = ["URL", "label", "source", "public_domain", "answer"] for row in frame[columns].itertuples(index=False, name=None): record = [row[0], int(row[1]), row[2], row[3], row[4]] encoded = json.dumps( record, ensure_ascii=False, separators=(",", ":") ) digest.update(encoded.encode("utf-8")) digest.update(b"\n") return digest.hexdigest() def prepare(train_path, config_directory): global tokenizer, INSTRUCTION, MAX_SEQUENCE_TOKENS, MAX_PROMPT_TOKENS base = read_config(config_directory, "base_model.json") prompt = read_config(config_directory, "prompt_config.json") policy = read_config(config_directory, "input_policy.json") expected = read_config(config_directory, "expected_splits.json") frame = pd.read_parquet(train_path) assert len(frame) == expected["splits"]["train"]["rows"] assert ordered_hash(frame) == ( expected["splits"]["train"]["ordered_records_sha256"] ) assert frame["label"].isin([0, 1]).all() assert ( frame["answer"] == frame["label"].map({1: "A", 0: "B"}) ).all() tokenizer = AutoTokenizer.from_pretrained( base["model_id"], revision=base["revision"] ) if tokenizer.pad_token_id is None: tokenizer.pad_token = tokenizer.eos_token tokenizer.padding_side = "left" INSTRUCTION = prompt["instruction"] MAX_SEQUENCE_TOKENS = policy["max_sequence_tokens"] MAX_PROMPT_TOKENS = policy["max_prompt_tokens"] label_ids = prompt["label_token_ids"] for answer in ["A", "B"]: assert tokenizer.encode( answer, add_special_tokens=False ) == [label_ids[answer]] records = [] truncated_count = 0 for row in frame.itertuples(index=False): ids, truncated, _ = encode_url_prompt(row.URL) answer_id = label_ids[row.answer] record = { "input_ids": ids + [answer_id], "attention_mask": [1] * (len(ids) + 1), "labels": [-100] * len(ids) + [answer_id], } assert len(record["input_ids"]) <= MAX_SEQUENCE_TOKENS assert sum(value != -100 for value in record["labels"]) == 1 records.append(record) truncated_count += int(truncated) summary = { "training_examples": len(records), "truncated_training_urls": truncated_count, "longest_sequence": max(len(row["input_ids"]) for row in records), "ordered_records_sha256": ordered_hash(frame), } return records, summary def train_adapter(records, config_directory, output): import torch from datasets import Dataset from peft import ( LoraConfig, get_peft_model, prepare_model_for_kbit_training, ) from transformers import ( AutoModelForCausalLM, BitsAndBytesConfig, Trainer, TrainingArguments, set_seed, ) output = Path(output) if output.exists(): raise FileExistsError("Use a new output directory.") assert torch.cuda.is_available() assert torch.cuda.is_bf16_supported() base = read_config(config_directory, "base_model.json") settings = read_config(config_directory, "training_arguments.json") saved_lora = read_config(config_directory, "lora_config.json") quantization = BitsAndBytesConfig( load_in_4bit=True, bnb_4bit_quant_type="nf4", bnb_4bit_use_double_quant=True, bnb_4bit_compute_dtype=torch.bfloat16, ) model = AutoModelForCausalLM.from_pretrained( base["model_id"], revision=base["revision"], quantization_config=quantization, device_map={"": 0}, dtype=torch.bfloat16, attn_implementation="sdpa", ) model.config.use_cache = False set_seed(42) model = prepare_model_for_kbit_training( model, use_gradient_checkpointing=True, gradient_checkpointing_kwargs={"use_reentrant": False}, ) model = get_peft_model(model, LoraConfig( r=saved_lora["r"], lora_alpha=saved_lora["lora_alpha"], lora_dropout=saved_lora["lora_dropout"], target_modules=saved_lora["target_modules"], bias=saved_lora["bias"], task_type=saved_lora["task_type"], )) assert sum( p.numel() for p in model.parameters() if p.requires_grad ) == 2621440 set_seed(42) arguments = TrainingArguments( output_dir=str(output / "checkpoints"), **settings, ) trainer = Trainer( model=model, args=arguments, train_dataset=Dataset.from_list(records), data_collator=training_collator, ) result = trainer.train() destination = output / "adapter" trainer.save_model(str(destination)) tokenizer.save_pretrained(str(destination)) (output / "training_metrics.json").write_text( json.dumps(result.metrics, indent=2), encoding="utf-8", ) def main(): parser = argparse.ArgumentParser() parser.add_argument("--train", required=True) parser.add_argument("--config", required=True) parser.add_argument("--output", required=True) parser.add_argument("--prepare-only", action="store_true") args = parser.parse_args() output = Path(args.output) if output.exists(): raise FileExistsError("Use a new output directory.") records, summary = prepare(args.train, args.config) if args.prepare_only: output.mkdir(parents=True, exist_ok=False) else: train_adapter(records, args.config, output) (output / "training_data_summary.json").write_text( json.dumps(summary, indent=2), encoding="utf-8", ) print(json.dumps(summary, indent=2)) if __name__ == "__main__": main()