gonzalolinares's picture
Use GRPO adapter (SFT+DPO+GRPO)
89f2d87 verified
Raw
History Blame Contribute Delete
5.43 kB
"""Gradio chat for the C++ compiler-tuned model."""
from __future__ import annotations
import os
import re
import torch
import gradio as gr
from peft import PeftModel
from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer
from threading import Thread
BASE_MODEL = os.environ.get("BASE_MODEL", "Qwen/Qwen2.5-1.5B-Instruct")
SFT_ADAPTER = os.environ.get("SFT_ADAPTER", "gonzalolinares/qwen25-1.5b-cpp-sft")
DPO_ADAPTER = os.environ.get("DPO_ADAPTER", "gonzalolinares/qwen25-1.5b-cpp-dpo")
GRPO_ADAPTER = os.environ.get("GRPO_ADAPTER", "gonzalolinares/qwen25-1.5b-cpp-grpo")
SYSTEM_PROMPT = (
"Eres un asistente que solo programa en C++ moderno (C++20). "
"Responde siempre con un único bloque de código ```cpp``` completo y compilable primero, "
"y después una breve explicación en español o inglés según el idioma del usuario."
)
EXAMPLES = [
"Escribe un programa C++ que imprima hola en una línea.",
"Crea un std::vector con {1,2,3} e imprime su tamaño con size().",
"Ordena el vector {3,1,2} con std::sort e imprime los valores separados por espacio.",
"Usa std::make_unique<int>(42) e imprime el valor.",
"Write a C++ program that prints the sum of 10 and 5.",
]
print("Loading tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(BASE_MODEL)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.float16 if device == "cuda" else torch.float32
print(f"Loading base model on {device}...")
model = AutoModelForCausalLM.from_pretrained(
BASE_MODEL,
torch_dtype=dtype,
device_map=device if device == "cuda" else None,
low_cpu_mem_usage=True,
)
print("Merging SFT adapter...")
model = PeftModel.from_pretrained(model, SFT_ADAPTER)
model = model.merge_and_unload()
print("Loading DPO adapter...")
model = PeftModel.from_pretrained(model, DPO_ADAPTER)
model = model.merge_and_unload()
print("Loading GRPO adapter...")
model = PeftModel.from_pretrained(model, GRPO_ADAPTER)
if device == "cpu":
model = model.to(device)
model.eval()
print("Model ready.")
def extract_cpp(text: str) -> str:
m = re.search(r"```(?:cpp|c\+\+)?\s*([\s\S]*?)```", text, re.IGNORECASE)
return m.group(1).strip() if m else ""
def build_messages(history: list[list[str | None]], user_message: str) -> list[dict]:
messages = [{"role": "system", "content": SYSTEM_PROMPT}]
for user_msg, assistant_msg in history:
if user_msg:
messages.append({"role": "user", "content": user_msg})
if assistant_msg:
messages.append({"role": "assistant", "content": assistant_msg})
messages.append({"role": "user", "content": user_message})
return messages
def stream_reply(history: list, max_tokens: int, temperature: float):
user_message = history[-1][0]
messages = build_messages(history[:-1], user_message)
text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
inputs = tokenizer(text, return_tensors="pt").to(model.device)
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)
gen_kwargs = dict(
**inputs,
streamer=streamer,
max_new_tokens=int(max_tokens),
do_sample=temperature > 0.01,
temperature=max(float(temperature), 0.01),
top_p=0.9,
pad_token_id=tokenizer.eos_token_id,
)
thread = Thread(target=model.generate, kwargs=gen_kwargs)
thread.start()
partial = ""
for chunk in streamer:
partial += chunk
history[-1][1] = partial
yield history, extract_cpp(partial)
thread.join()
with gr.Blocks(title="C++ Compiler Chat", theme=gr.themes.Soft()) as demo:
gr.Markdown(
"""
# ⚙️ C++ Compiler Chat
Modelo fine-tuned para **C++20** (`gonzalolinares/qwen25-1.5b-cpp-grpo` — SFT + DPO + GRPO con `g++`).
Pide un programa en lenguaje natural; la respuesta empieza con ` ```cpp `.
"""
)
with gr.Row():
max_tokens = gr.Slider(64, 1024, value=512, step=64, label="Max tokens")
temperature = gr.Slider(0.0, 1.0, value=0.1, step=0.05, label="Temperature")
chatbot = gr.Chatbot(height=420, label="Chat", type="tuples")
msg = gr.Textbox(
placeholder="Ej: Escribe un programa que imprima los números del 1 al 5...",
label="Tu mensaje",
lines=2,
)
code_preview = gr.Code(language="cpp", label="Código extraído", lines=14)
with gr.Row():
send = gr.Button("Enviar", variant="primary")
clear = gr.Button("Limpiar")
gr.Examples(examples=[[e] for e in EXAMPLES], inputs=msg, label="Ejemplos")
def add_message(user_message, history):
if not user_message.strip():
return "", history
return "", history + [[user_message, None]]
def respond(history, max_tok, temp):
yield from stream_reply(history, max_tok, temp)
msg.submit(add_message, [msg, chatbot], [msg, chatbot], queue=False).then(
respond, [chatbot, max_tokens, temperature], [chatbot, code_preview]
)
send.click(add_message, [msg, chatbot], [msg, chatbot], queue=False).then(
respond, [chatbot, max_tokens, temperature], [chatbot, code_preview]
)
clear.click(lambda: ([], ""), None, [chatbot, code_preview])
if __name__ == "__main__":
demo.queue(max_size=8).launch()