File size: 5,433 Bytes
54fc58d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
89f2d87
54fc58d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
89f2d87
 
 
54fc58d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
89f2d87
54fc58d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
"""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()