From e4cb1ffa556cfb1bb21162ea07c59289308db2b9 Mon Sep 17 00:00:00 2001 From: 51616 Date: Fri, 27 Jun 2025 17:43:37 +0900 Subject: [PATCH] add todo for self-gen logits --- README.md | 4 ++-- data/self_generate_qa.py | 4 ++-- src/ctx_to_lora/data/processing.py | 1 + src/ctx_to_lora/trainer.py | 2 +- 4 files changed, 6 insertions(+), 5 deletions(-) diff --git a/README.md b/README.md index 504752c..756516f 100644 --- a/README.md +++ b/README.md @@ -105,8 +105,8 @@ ctx_model_max_len = 2**13 # load via a custom function because of custom chat_template tokenizer = get_tokenizer(base_model_name) ctx_tokenizer = get_tokenizer(base_model_name) -ds = def get_tokenized_dataset( - ds_name +ds = get_tokenized_dataset( + ds_name, split="train", base_model_max_len=base_model_max_len, tokenizer=tokenizer, diff --git a/data/self_generate_qa.py b/data/self_generate_qa.py index 05a0bf0..5fa93a4 100644 --- a/data/self_generate_qa.py +++ b/data/self_generate_qa.py @@ -1,6 +1,5 @@ import argparse import os -import random from glob import glob import pandas as pd @@ -160,11 +159,12 @@ def self_generate( messages = create_messages(ctxs, questions, args.vllm_model, SYSTEM_TEMPLATE) print(f"Generating from {len(messages)} contexts") + # TODO (distillation): make vllm outputs logits here too completions = llm.chat( messages, sampling_params=SamplingParams( max_tokens=2048, - temperature=1.0, + temperature=1.0, # TODO: lower the temp ), ) diff --git a/src/ctx_to_lora/data/processing.py b/src/ctx_to_lora/data/processing.py index 7128139..64d9794 100644 --- a/src/ctx_to_lora/data/processing.py +++ b/src/ctx_to_lora/data/processing.py @@ -592,6 +592,7 @@ def get_tokenized_dataset( set_format: str | None = None, streaming: bool = False, ) -> dict[str, Any]: + # TODO (distillation, Tan): make this works with pre-computed logits assert not use_kl_loss, "KL loss is deprecated" if add_repeat_prompt: assert repeat_prob > 0, f"add_repeat_prompt is set but repeat_prob = 0" diff --git a/src/ctx_to_lora/trainer.py b/src/ctx_to_lora/trainer.py index e25534e..65fb425 100644 --- a/src/ctx_to_lora/trainer.py +++ b/src/ctx_to_lora/trainer.py @@ -10,7 +10,7 @@ from ctx_to_lora.modeling.hypernet import ModulatedPretrainedModel logger = logging.getLogger() -# TODO (distillation): implement +# TODO (distillation, Shin): implement class DistillationTrainer(Trainer): ...