From e6294f08dae2d1054178992706cac648872eec4d Mon Sep 17 00:00:00 2001 From: 51616 Date: Wed, 14 May 2025 07:59:32 +0000 Subject: [PATCH] fix templating webui bug --- README.md | 7 +++++++ webui/app.py | 28 ++++++++++++---------------- 2 files changed, 19 insertions(+), 16 deletions(-) diff --git a/README.md b/README.md index 47cbbb3..72abfe5 100644 --- a/README.md +++ b/README.md @@ -50,6 +50,13 @@ run python intx_sft.py configs/...yaml ... --from_pretrained_checkpoint=train_ou ### Evaluation LongBench ```bash +# generative + run python src/ctx_to_lora/eval.py --checkpoint_path train_outputs/runs/May08_13-56-31_slurm0-a3nodeset-5_59383_906acb28/checkpoint-105000/pytorch_model.bin + +# benchmark cd LongBench/LongBench + +run python pred_ctx_to_lora.py --checkpoint_path ../../train_outputs/runs/Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782/pytorch_model.bin + run python eval_ctx_to_lora.py --model_name Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782 --checkpoint_path ../../train_outputs/runs/Mar16_12-38-01_slurm0-a3nodeset-12_54818_32426662/checkpoint-136782/pytorch_model.bin ``` \ No newline at end of file diff --git a/webui/app.py b/webui/app.py index 93f2071..7a9a0b5 100644 --- a/webui/app.py +++ b/webui/app.py @@ -8,6 +8,7 @@ import yaml from flask import Flask, abort, jsonify, render_template, request from transformers import pipeline, AutoTokenizer +from ctx_to_lora.data_utils import tokenize_ctx_text from ctx_to_lora.model_loading import get_tokenizer app = Flask(__name__) @@ -354,7 +355,7 @@ def load_checkpoint(): train=False, use_flash_attn=True, ) - modulated_model = modulated_model.to(device) + modulated_model = modulated_model.to(device).to(torch.bfloat16) modulated_model.eval() result = { @@ -393,23 +394,18 @@ def process_multiple_contexts(contexts, ctx_tokenizer): print(f"Processing {len(contexts)} non-empty contexts") - # Tokenize each context - all_ctx_ids = [] - all_ctx_attn_mask = [] - - for context in contexts: - # Tokenize the single context - inputs = ctx_tokenizer(context, return_tensors="pt") - - all_ctx_ids.append(inputs["input_ids"][0]) - all_ctx_attn_mask.append(inputs["attention_mask"][0]) - - # Pad to the same length + tokenized_contexts = tokenize_ctx_text({"context": contexts}, ctx_tokenizer) + ctx_ids = tokenized_contexts["ctx_ids"] + ctx_attn_mask = tokenized_contexts["ctx_attn_mask"] ctx_ids = torch.nn.utils.rnn.pad_sequence( - all_ctx_ids, batch_first=True, padding_value=ctx_tokenizer.pad_token_id + torch.tensor(ctx_ids, dtype=torch.long, device=device), + batch_first=True, + padding_value=0, ) ctx_attn_mask = torch.nn.utils.rnn.pad_sequence( - all_ctx_attn_mask, batch_first=True, padding_value=0 + torch.tensor(ctx_attn_mask, dtype=torch.long, device=device), + batch_first=True, + padding_value=0, ) return {"ctx_ids": ctx_ids, "ctx_attn_mask": ctx_attn_mask} @@ -483,7 +479,7 @@ def chat(): # Tokenize the chat history model_inputs = base_tokenizer.apply_chat_template( - chat_history, return_tensors="pt" + chat_history, return_tensors="pt", add_generation_prompt=True ).to(device) # Generate response with context-modulated model