bengoldberg0's picture
Publish checked experiment pipeline and verification scope
d28e9c6 verified
Raw History Blame Contribute Delete
7.85 kB
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()