Team Ai
Apppublic

tecuhtli/assistant-t5-qa-data-processing

sourceHugging Facemitupdated 10mo agoView on Hugging Face
0likes
app.py904 linesDownload Raw Back to root
1#***************************************************************************2# Mori (tech-only) — Streamlit App sin sidebar ni social, con RAG opcional3#***************************************************************************4import os, sys, warnings, json, joblib, random, re, unicodedata, uuid, torch, csv5import numpy as np6os.environ["TF_ENABLE_ONEDNN_OPTS"] = "0"7import streamlit as st8import datetime as dt9from pathlib import Path10import torch11import numpy as np12from transformers import AutoTokenizer, AutoModelForSeq2SeqLM, AutoModelForSequenceClassification13from huggingface_hub import hf_hub_download14from sentence_transformers import SentenceTransformer  # RAG embeddings15 16# =========================17# Configuración general18# =========================19HF_TOKEN = os.environ.get("HF_TOKEN")  # Token privado (colócalo en Secrets o variable de entorno)20 21#***************************************************************************22# Sidebar controls for generation params23#***************************************************************************24 25def sidebar_params():26    27    with st.sidebar:28        st.title("🎮 Adjustments (T5-Base)")29 30        ss = st.session_state31        # Defaults (solo 1ª vez)32 33        # Estado inicial: ocultar ajustes avanzados34        ss = st.session_state35        if "show_llm_controls" not in ss:36            ss.show_llm_controls = False37 38        39        ss.setdefault("persona", "Normal")40        ss.setdefault("mode", "beam")  # 'beam' | 'sampling'41        ss.setdefault("max_new", 128)42        ss.setdefault("min_tok", 16)43        ss.setdefault("no_repeat", 3)44        ss.setdefault("num_beams", 4)45        ss.setdefault("length_penalty", 1.0)46        ss.setdefault("temperature", 0.7)47        ss.setdefault("top_p", 0.9)48        ss.setdefault("repetition_penalty", 1.0)49        ss.setdefault("show_llm_controls", True)  # Toggle principal50 51        # ----------------------------52        # Personalidad (presets)53        # ----------------------------54        st.header("💡 Predefined Personalities")55        c1, c2 = st.columns(2)56 57        with c1:58            if st.button("Normal 🧐", use_container_width=True):59                ss.update({60                    "persona": "Normal",61                    "mode": "beam",62                    "num_beams": 1,63                    "max_new": 92,64                    "min_tok": 32,65                    "no_repeat": 3,66                    "length_penalty": .3,67                    "temperature": 0.4,68                    "top_p": 0.9,69                    "repetition_penalty": .4,70                })71                st.rerun()72 73        with c2:74            if st.button("Enthusiastic 😃", use_container_width=True):75                ss.update({76                    "persona": "Enthusiastic",   # <- corregido77                    "mode": "sampling",78                    "max_new": 192,79                    "min_tok": 48,80                    "no_repeat": 3,81                    "temperature": .8,82                    "top_p": 0.95,83                    "repetition_penalty": 1.0,84                })85                st.rerun()86 87        st.caption(f"Selected Personality: **{ss.persona}**")88 89        # ----------------------------90        # Botón para mostrar/ocultar parámetros91        # ----------------------------92        if st.button(("🔼 Hide" if ss.show_llm_controls else "🔽 Show") + " Advanced Settings"):93            ss.show_llm_controls = not ss.show_llm_controls94            st.rerun()95 96        # ----------------------------97        # Controles del modelo (sliders, estrategia, etc.)98        # ----------------------------99        if ss.show_llm_controls:100            st.header("⚙️ Manual Adjustments")101            st.subheader("📝 Text Generation")102            picked = st.radio(103                "Strategy",104                ["Beam search (stable)", "Sampling (creative)"],105                index=0 if ss.mode == "beam" else 1,106                help="https://huggingface.co/docs/transformers/generation_strategies"107            )108            ss.mode = "beam" if picked.startswith("Beam") else "sampling"109 110            st.subheader("🔧 LLM text generation parameters")111            ss.max_new = st.slider(112                "max_new_tokens", 16, 256, int(ss.max_new), step=8,113                help="https://huggingface.co/docs/transformers/main_classes/text_generation"114            )115            ss.min_tok = st.slider(116                "min_tokens", 0, int(ss.max_new), int(ss.min_tok),117                help="https://huggingface.co/docs/transformers/main_classes/text_generation"118            )119            ss.no_repeat = st.slider(120                "no_repeat_ngram_size", 0, 6, int(ss.no_repeat),121                help="https://huggingface.co/docs/transformers/main_classes/text_generation"122            )123 124            # Subcontroles según modo125            if ss.mode == "beam":126                ss.num_beams = st.slider(127                    "num_beams", 2, 8, int(ss.num_beams),128                    help="https://huggingface.co/docs/transformers/main_classes/text_generation"129                )130                ss.length_penalty = st.slider(131                    "length_penalty", 0.0, 2.0, float(ss.length_penalty),132                    step=0.1, help="https://huggingface.co/docs/transformers/main_classes/text_generation"133                )134            else:135                ss.temperature = st.slider(136                    "temperature", 0.1, 1.5, float(ss.temperature),137                    step=0.05, help="https://huggingface.co/docs/transformers/main_classes/text_generation"138                )139                ss.top_p = st.slider(140                    "top_p", 0.5, 1.0, float(ss.top_p),141                    step=0.01, help="https://huggingface.co/docs/transformers/main_classes/text_generation"142                )143 144       145        if "last_prompt" in st.session_state and st.session_state["last_prompt"]:146            with st.expander("Show generated prompt"):147                st.text_area(148                    "Prompt actual:",149                    st.session_state["last_prompt"],150                    height=200,151                    disabled=True152                )153        else:154            st.caption("👉 No prompt is available yet.")        155 156        # ----------------------------157        # Construir diccionario de parámetros158        # ----------------------------159        params = {160            "persona": ss.persona,161            "mode": ss.mode,162            "max_new_tokens": int(ss.max_new),163            "min_tokens": int(ss.min_tok),164            "no_repeat_ngram_size": int(ss.no_repeat),165            "repetition_penalty": float(ss.repetition_penalty),166        }167        if ss.mode == "beam":168            params.update({169                "num_beams": int(ss.num_beams),170                "length_penalty": float(ss.length_penalty),171            })172        else:173            params.update({174                "temperature": float(ss.temperature),175                "top_p": float(ss.top_p),176            })177 178        return params179 180 181#***************************************************************************182# Functions183#***************************************************************************184 185 186def truncate_sentences(text: str, max_sentences: int = 4) -> str:187    _SENT_SPLIT = re.compile(r'(?<=[\.\!\?…])\s+')188    s = text.strip()189    if not s: return s190    parts = _SENT_SPLIT.split(s)191    cut = " ".join(parts[:max_sentences]).strip()192    if cut and cut[-1] not in ".!?…": cut += "."193    return cut194 195 196def _load_json_safe(path: Path, fallback: dict) -> dict:197    try:198        with open(path, "r", encoding="utf-8") as f:199            return json.load(f)200    except Exception:201        return fallback202 203# Function to clean the question field204def limpiar_input():205    st.session_state["entrada"] = ""206 207# ✅ Corrige la ruta correctamente desde Scripts hacia Models208def get_model_path(folder_name):209    return Path("Models") / folder_name210 211# Function to save user interaction212def saving_interaction(question, response, context, user_id):213    '''214    inputs:215    question --> User input question216    response --> Assistant response to the user question217    context --> Context related to the user input, found by the trained classifier218    user_id --> ID for the current user (Unique ID per session)219    '''220    timestamp = dt.datetime.now().isoformat()221    stats_dir = Path("Statistics")222    stats_dir.mkdir(parents=True, exist_ok=True)223 224    archivo_csv = stats_dir / "conversaciones_log.csv"225    existe_csv = archivo_csv.exists()226 227    with open(archivo_csv, mode="a", encoding="utf-8", newline="") as f_csv:228        writer = csv.writer(f_csv)229        if not existe_csv:230            writer.writerow(["timestamp", "user_id", "contexto", "pregunta", "respuesta"])231        writer.writerow([timestamp, user_id, context, question, response])232 233    archivo_jsonl = stats_dir / "conversaciones_log.jsonl"234    with open(archivo_jsonl, mode="a", encoding="utf-8") as f_jsonl:235        registro = {236            "timestamp": timestamp,237            "user_id": user_id,238            "context": context,239            "pregunta": question,240            "respuesta": response}241        f_jsonl.write(json.dumps(registro, ensure_ascii=False) + "\n")242 243# Function to load models within the huggingface repositories space244@st.cache_resource245def load_model(path_str):246    path = Path(path_str).resolve()247    tokenizer = AutoTokenizer.from_pretrained(path, local_files_only=True)248    model = AutoModelForSeq2SeqLM.from_pretrained(path, local_files_only=True)249    return model, tokenizer250 251#-------------------------------------------------------------------------252# Function to correct Spanish sentences' punctuation and missing characters253#-------------------------------------------------------------------------254 255def polish_spanish(s: str) -> str:256    s = unicodedata.normalize("NFC", s).strip()257    s = re.sub(r'\s*[\[\(]\s*Assistant\s+(?:Social|T[eé]nico|T[eé]cnico)\s*[\]\)]\s*', '', s, flags=re.I)258    fixes = [259        (r'(?i)(^|\W)T\s+puedes(?P<p>[^\w]|$)', r'\1Tú puedes\g<p>'),260        (r'(?i)(^|\W)T\s+(ya|eres|estas|estás|tienes|puedes)\b', r'\1Tú \2'),261        (r'(?i)\bclaro que s(?:i|í)?\b(?P<p>[,.\!?…])?', r'Claro que sí\g<p>'),262        (r'(?i)(^|\s)si,', r'\1Sí,'),263        (r'(?i)(\beso\s+)s(\s+est[áa]\b)', r'\1sí\2'),264        (r'(?i)(^|[\s,;:])s(\s+es\b)', r'\1sí\2'),265        (r'(?i)\btiles\b', 'útiles'),266        (r'(?i)\butiles\b', 'útiles'),267        (r'(?i)\butil\b', 'útil'),268        (r'(?i)\baqui\b', 'aquí'),269        (r'(?i)\baqu\b(?=\s+estoy\b)', 'aquí'),270        (r'(?i)\balgn\b', 'algún'),271        (r'(?i)\balgun\b', 'algún'),272        (r'(?i)\bAnimo\b', 'Ánimo'),273        (r'(?i)\bcario\b', 'cariño'),274        (r'(?i)\baprendisaje\b', 'aprendizaje'),275        (r'(?i)\bmanana\b', 'mañana'),276        (r'(?i)\bmaana\b', 'mañana'),277        (r'(?i)\benergia\b', 'energía'),278        (r'(?i)\benerga\b', 'energía'),279        (r'(?i)\bextrano\b', 'extraño'),280        (r'(?i)\bextrana\b', 'extraña'),281        (r'(?i)\bextranar\b', 'extrañar'),282        (r'(?i)\bextranarte\b', 'extrañarte'),283        (r'(?i)\bextranas\b', 'extrañas'),284        (r'(?i)\bextranos\b', 'extraños'),285        (r'(?i)\baqu\b', 'aquí'),286        (r'(?i)\baqui\b', 'aquí'),287        (r'(?i)\bestare\b', 'estaré'),288        (r'(?i)\bclarn\b', 'clarín'),289        (r'(?i)\bclarin\b', 'clarín'),290        (r'(?i)\bclar[íi]n\s+cornetas\b', 'clarín cornetas'),291        (r'(?i)(^|\s)s([,.;:!?])', r'\1Sí\2'),292        (r'(?i)\bfutbol\b', 'fútbol'),293        (r'(?i)(^|\s)as(\s+se\b)', r'\1Así\2'),294        (r'(?i)(^|\s)s(\s+orientarte\b)', r'\1sí\2'),295        (r'(?i)\bbuen dia\b', 'buen día'),296        (r'(?i)\bgran dia\b', 'gran día'),297        (r'(?i)\bdias\b', 'días'),298        (r'(?i)\bdia\b', 'día'),299        (r'(?i)\bgran da\b', 'gran día'),300        (r'(?i)\bacompa?a(r|rte|do|da|dos|das)?\b', r'acompaña\1'),301        (r'(?i)(^|\s)as([,.;:!?]|\s|$)', r'\1así\2'),302        (r'(?i)(^|\s)S lo se\b', r'\1Sí lo sé'),303        (r'(?i)(^|\s)S lo sé\b', r'\1Sí lo sé'),304        (r'(?i)\bcudese\b', 'cuídese'),305        (r'(?i)\bpequeo\b', 'pequeño'),306        (r'(?i)\bpequea\b', 'pequeña'),307        (r'(?i)\bpequeos\b', 'pequeños'),308        (r'(?i)\bpequeas\b', 'pequeñas'),309        (r'(?i)\bunico\b', 'único'),310        (r'(?i)\bunica\b', 'única'),311        (r'(?i)\bunicos\b', 'únicos'),312        (r'(?i)\bunicas\b', 'únicas'),313        (r'(?i)\bnico\b', 'único'),314        (r'(?i)\bnica\b', 'única'),315        (r'(?i)\bnicos\b', 'únicos'),316        (r'(?i)\bnicas\b', 'únicas'),317        (r'(?i)\bestadstico\b', 'estadístico'),318        (r'(?i)\bestadstica\b', 'estadística'),319        (r'(?i)\bestadsticos\b', 'estadísticos'),320        (r'(?i)\bestadsticas\b', 'estadísticas'),321        (r'(?i)\bcudate\b', 'cuídate'),322        (r'(?i)\bcuidate\b', 'cuídate'),323        (r'(?i)\bcuidese\b', 'cuídese'),324        (r'(?i)\bcudese\b', 'cuídese'),325        (r'(?i)\bcuidense\b', 'cuídense'),326        (r'(?i)\bcudense\b', 'cuídense'),327        (r'(?i)\bgracias por confiar en m\b', 'gracias por confiar en mí'),328        (r'(?i)\bcada dia\b', 'cada día'),329        (r'(?i)\bcada da\b', 'cada día'),330        (r'(?i)\bsegun\b', 'según'),331        (r'(?i)\bcaracteristica(s)?\b', r'característica\1'),332        (r'(?i)\bcaracterstica(s)?\b', r'característica\1'),333        (r'(?i)\b([a-záéíóúñ]+)cion\b', r'\1ción'),334        (r'(?i)\bdeterminacio\b', 'determinación'),335    ]336    for pat, rep in fixes:337        s = re.sub(pat, rep, s)338 339    s = re.sub(r'(?i)^eso es todo!(?P<r>(\s|$).*)', r'¡Eso es todo!\g<r>', s)340 341    def add_opening_q(m):342        cuerpo = m.group('qbody')343        if '¿' in cuerpo:344            return m.group(0)345        return f"{m.group('pre')}¿{cuerpo}"346    s = re.sub(r'(?P<pre>(^|[\.!\…]\s+))(?P<qbody>[^?]*\?)', add_opening_q, s)347 348    def _open_exclam(m):349        palabra = m.group('w')350        resto   = m.group('r') or ''351        return f'¡{palabra}!{resto}'352    s = re.sub(r'(?i)^(?P<w>(hola|gracias|genial|perfecto|claro|por supuesto|con gusto|listo|vaya|wow|tu puedes|tú puedes|clarín|clarin|clarín cornetas))!(?P<r>(\s|$).*)',_open_exclam, s)353 354    s = re.sub(r'\s+', ' ', s).strip()355    if s and s[-1] not in ".!?…":356        s += "."357    return s358 359#-------------------------------------------------------------------------360# Function to remove repeated input in the Model answer361#-------------------------------------------------------------------------362 363def anti_echo(response: str, user_text: str) -> str:364    rn = normalize_for_route(response)365    un = normalize_for_route(user_text)366    def _clean_leading(s: str) -> str:367        s = re.sub(r'^\s*[,;:\-–—]\s*', '', s)368        s = re.sub(r'^\s+', '', s)369        return s370    if len(un) >= 4 and rn.startswith(un):371        cut = re.sub(r'^\s*[^,;:\.\!\?]{0,120}[,;:\-]\s*', '', response).lstrip()372        if cut and cut != response:373            return _clean_leading(cut)374        return _clean_leading(response[len(user_text):])375    return response376 377#-------------------------------------------------------------------------378# Normalization helpers379#-------------------------------------------------------------------------380 381def normalize_for_route(s: str) -> str:382    s = unicodedata.normalize("NFKD", s)383    s = "".join(ch for ch in s if not unicodedata.combining(ch))384    s = re.sub(r"[^\w\s-]", " ", s, flags=re.UNICODE)385    s = re.sub(r"\s+", " ", s).strip().lower()386    return s387 388_Q_STARTERS = {389    "como","que","quien","quienes","cuando","donde","por que","para que",390    "cual","cuales","cuanto","cuantos","cuanta","cuantas"391}392_EXC_TRIGGERS = {"motiva","motivame","animate","animame","animo","ayudame","ayudame porfa", "clarin", "clarín", "clarinete", "clarin cornetas"}393SPECIAL_NOPUNCT = {"kiubo", "quiubo", "que chido", "qué chido", "que buena onda"}394_Q_VERB_STARTERS = {"eres","estas","estás","puedes","sabes","tienes","quieres","conoces",395        "crees","piensas","dirias","dirías","podrias","podrías","podras","podrás"}396 397#-------------------------------------------------------------------------398# Punctuation helpers399#-------------------------------------------------------------------------400 401def needs_question_marks(norm: str) -> bool:402    if "?" in norm: return False403    for w in _Q_STARTERS:404        if norm.startswith(w + " ") or norm == w:405            return True406    return False407 408def needs_exclam(norm: str) -> bool:409    if "!" in norm: return False410    return any(t in norm for t in _EXC_TRIGGERS)411 412#-------------------------------------------------------------------------413# Greetings detection414#-------------------------------------------------------------------------415 416def is_slang_greeting(norm: str) -> bool:417    SHORT = {418        "que pex", "que onda", "ke pex", "k pex", "q onda",419        "kiubo", "quiubo", "quiubole", "quiúbole", "kionda", "q onda", "k onda",420        "que rollo", "ke onda", "que show", "que tranza"421    }422    if norm in SHORT: return True423    if re.match(r"^(q|k|ke|que)\s+(pex|onda|rollo|show|tranza)\b", norm): return True424    if re.match(r"^(kiubo|quiubo|quiubole|quiúbole|quiubol[e]?)\b", norm): return True425    return False426 427#-------------------------------------------------------------------------428# Capitalization & autopunct429#-------------------------------------------------------------------------430 431def capitalize_spanish(s: str) -> str:432    s = s.strip()433    i = 0434    while i < len(s) and not s[i].isalpha():435        i += 1436    if i < len(s):437        s = s[:i] + s[i].upper() + s[i+1:]438    return s439 440def smart_autopunct(user_text: str) -> str:441    s = user_text.strip()442    if len(s) > 20:443        return capitalize_spanish(s)444    norm = normalize_for_route(s)445    if norm in SPECIAL_NOPUNCT:446        s = re.sub(r'[¿?!¡]+', '', s).strip()447        return capitalize_spanish(s)448    if norm.startswith("y si "):449        s = f"¿{s}?"450        return capitalize_spanish(s)451    if "?" in s and "¿" not in s:452        s = "¿" + s453        return capitalize_spanish(s)454    if "!" in s and "¡" not in s:455        s = "¡" + s456        return capitalize_spanish(s)457    if is_slang_greeting(norm):458        s = f"¡{s}!"459        return capitalize_spanish(s)460    if needs_question_marks(norm):461        s = f"¿{s}?"462        return capitalize_spanish(s)463    toks = norm.split()464    if toks and toks[0] in _Q_VERB_STARTERS:465        s = f"¿{s}?"466        return capitalize_spanish(s)467    if re.match(r"^(me\s+ayudas?|me\s+puedes|podrias?|podras?)\b", norm):468        s = f"¿{s}?"469        return capitalize_spanish(s)470    if needs_exclam(norm):471        s = f"¡{s}!"472        return capitalize_spanish(s)473    return capitalize_spanish(s)474 475 476#-------------------------------------------------------------------------477# Seeds & helpers478#-------------------------------------------------------------------------479 480def set_seeds(seed: int = 42):481    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)482    if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)483    torch.backends.cudnn.deterministic = True484    torch.backends.cudnn.benchmark = False485 486# --- Personalidades (solo estilo en prompt; parámetros ya vienen del sidebar) ---487 488def persona_style_prompt(persona: str, domain: str) -> str:489    """Instrucción breve de estilo según personalidad y dominio (technical/social)."""490    if persona == "Enthusiastic":491        return (492            "Responde de forma creativa, usa al menos 232 palabras. ")493    if persona == "Normal":  # ya no se usa, pero por compatibilidad494        return ""495    return ""  # Assistant response496 497#-------------------------------------------------------------------------498# Classifier499#-------------------------------------------------------------------------500 501def classify_context(question, label_classes, model, tokenizer, device):502    model = model.to(device)503    inputs = tokenizer(question, return_tensors="pt", padding=True, truncation=True, max_length=128)504    inputs = {k: v.to(device) for k, v in inputs.items()}505    with torch.no_grad():506        outputs = model(**inputs)507        logits = outputs.logits508    pred_intent = torch.argmax(logits, dim=1).item()509    predicted_label = label_classes[pred_intent]510    return predicted_label511 512#-------------------------------------------------------------------------513# Chatbot response for technical contexts using a Hugging Face model514#-------------------------------------------------------------------------515 516def technical_asnwer(question, context, model, tokenizer, device, gen_params=None):517    model = model.to(device).eval()518    persona_name = (gen_params or {}).get("persona", st.session_state.get("persona", "Normal"))519    style = persona_style_prompt(persona_name, "technical")520    521    # Promp Engineering para ayudar al asistente a encontrar la mejor respuesta522    input_text = f"{style}Context: {context} [SEP] Question: {question}."523        524    st.session_state["last_prompt"] = input_text  # o prompt525    st.session_state["just_generated"] = True526    #st.rerun()527    enc = tokenizer(input_text, return_tensors="pt", padding=True, truncation=True, max_length=256).to(device)528 529    bad_words = ["["]530    bad_ids = [tokenizer(bw, add_special_tokens=False).input_ids for bw in bad_words]531 532    # --- construir kwargs de generación, SIN tocar nada por personalidad ---533    max_new   = int((gen_params).get("max_new_tokens"))534    min_new   = int((gen_params).get("min_tokens"))          # <- ahora SIEMPRE min_new_tokens535    no_repeat = int((gen_params).get("no_repeat_ngram_size"))536    rep_pen   = float((gen_params).get("repetition_penalty"))537    mode      = (gen_params or {}).get("mode", "beam")538 539    if mode == "sampling":540        temperature = float((gen_params or {}).get("temperature", 0.7))541        top_p       = float((gen_params or {}).get("top_p", 0.9))542        kwargs = dict(543            do_sample=True,544            num_beams=1,545            temperature=max(0.1, temperature),546            top_p=min(1.0, max(0.5, top_p)),547            max_new_tokens=max_new,548            min_new_tokens=max(0, min_new),   # 👈 consistente549            no_repeat_ngram_size=no_repeat,550            repetition_penalty=max(1.0, rep_pen),551            bad_words_ids=bad_ids,552            eos_token_id=tokenizer.eos_token_id,553            pad_token_id=tokenizer.pad_token_id,554        )555    else:556        num_beams      = max(2, int((gen_params or {}).get("num_beams", 4)))557        length_penalty = float((gen_params or {}).get("length_penalty", 1.0))558        kwargs = dict(559            do_sample=False,560            num_beams=num_beams,561            length_penalty=length_penalty,562            max_new_tokens=max_new,563            min_new_tokens=max(0, min_new),   # 👈 también aquí (no min_length)564            no_repeat_ngram_size=no_repeat,565            repetition_penalty=max(1.0, rep_pen),566            bad_words_ids=bad_ids,567            eos_token_id=tokenizer.eos_token_id,568            pad_token_id=tokenizer.pad_token_id,569        )570 571    out_ids = model.generate(572        input_ids=enc["input_ids"], attention_mask=enc["attention_mask"], **kwargs573    )574    text = tokenizer.decode(out_ids[0], skip_special_tokens=True)575 576    if persona_name == "Normal":577        text = truncate_sentences(text, max_sentences=1)578 579    st.session_state["last_response"] = text580    #st.rerun()581 582    583    return polish_spanish(text)584 585#-------------------------------------------------------------------------586# Chatbot response for social contexts using a Hugging Face model587#-------------------------------------------------------------------------588 589def social_asnwer(question, model, tokenizer, device, gen_params=None, block_web=True):590    591    model = model.to(device).eval()592    persona_name = (gen_params or {}).get("persona", st.session_state.get("persona", "Normal"))593    prompt_type  = st.session_state.get("prompt_type", "Zero-shot")594    prompt = question595 596    st.session_state["last_prompt"] = prompt  # o prompt597    st.session_state["just_generated"] = True598 599 600    enc = tokenizer(prompt, return_tensors="pt", padding=True, truncation=True, max_length=192).to(device)601 602    bad_words = ["[", "Thanks", "thank you"]603    if block_web:604        bad_words += ["website", "http", "www", ".com"]605    bad_ids = [tokenizer(bw, add_special_tokens=False).input_ids for bw in bad_words]606 607    608    max_new = int((gen_params).get("max_new_tokens"))609    min_tokens = int((gen_params).get("min_tokens"))610    min_length = int(enc["input_ids"].shape[1]) + max(0, min_tokens)611    no_repeat = int((gen_params).get("no_repeat_ngram_size"))612    rep_pen = float((gen_params).get("repetition_penalty"))613    mode = (gen_params or {}).get("mode", "beam")614 615    if mode == "sampling":616        temperature = float((gen_params or {}).get("temperature", 0.7))617        top_p = float((gen_params or {}).get("top_p", 0.9))618        kwargs = dict(619            do_sample=True, num_beams=1,620            temperature=max(0.1, temperature),621            top_p=min(1.0, max(0.5, top_p)),622            max_new_tokens=max_new,623            #min_length=min_length,624            min_new_tokens=max(0, min_tokens),625            no_repeat_ngram_size=no_repeat,626            repetition_penalty=max(1.0, rep_pen),627            bad_words_ids=bad_ids,628            eos_token_id=tokenizer.eos_token_id,629            pad_token_id=tokenizer.pad_token_id,           630        )631    else:632        num_beams = max(2, int((gen_params or {}).get("num_beams", 4)))633        length_penalty = float((gen_params or {}).get("length_penalty", 1.0))634        kwargs = dict(635            do_sample=False, num_beams=num_beams, length_penalty=length_penalty,636            max_new_tokens=max_new,637            #min_length=min_length,638            min_new_tokens=max(0, min_tokens),   # <- usar min_new_tokens639            no_repeat_ngram_size=no_repeat,640            repetition_penalty=max(1.0, rep_pen),641            bad_words_ids=bad_ids,642            eos_token_id=tokenizer.eos_token_id,643            pad_token_id=tokenizer.pad_token_id,644            645        )646 647    out_ids = model.generate(648        input_ids=enc["input_ids"], attention_mask=enc["attention_mask"], **kwargs649    )650    text = tokenizer.decode(out_ids[0], skip_special_tokens=True)651    if persona_name == "Normal":652        text = truncate_sentences(text, max_sentences=2)653    #text = anti_echo(text, question)654    text = polish_spanish(text)655    text = capitalize_spanish(text)656 657    st.session_state["last_response"] = text658    #st.rerun()659 660    661    return text662 663#-------------------------------------------------------------------------664# Rule overrides665#-------------------------------------------------------------------------666 667def rule_intent_override(user_text: str, predicted_label: str) -> str:668    n = normalize_for_route(user_text)669    if re.fullmatch(r"(motivame|motiva|animame|animo|ayudame|que tranza|qué tranza|que tranza)", n):670        return "social"671    return predicted_label672 673#-------------------------------------------------------------------------674# Router675#-------------------------------------------------------------------------676 677def contextual_asnwer(question, label_classes, context_model, cont_tok,678                      tec_model, tec_tok, soc_model, soc_tok, device, gen_params=None, block_web=True):679    context = classify_context(question, label_classes, context_model, cont_tok, device)680    context = rule_intent_override(question, context)681 682    context_icons = {683        "social": "💬", "modelos": "🔧", "evaluación": "📏", "optimización": "⚙️",684        "visualización": "📈", "aprendizaje": "🧠", "vida digital": "🧑‍💻",685        "estadística": "📊", "infraestructura": "🖥", "datos": "📂", "transformación digital": "🌀"}686    icon = context_icons.get(context, "🧠")687 688    if gen_params and "seed" in gen_params:689        set_seeds(gen_params["seed"])690 691    if context == "social":692        return social_asnwer(question, soc_model, soc_tok, device, gen_params=gen_params, block_web=block_web), context693    else:694        return technical_asnwer(question, context, tec_model, tec_tok, device, gen_params=gen_params), context695 696#***************************************************************************697# MAIN698#***************************************************************************699 700if __name__ == '__main__':701 702    # --- Estado que debe persistir en todos los reruns ---703    ss = st.session_state704    ss.setdefault("historial", [])705    ss.setdefault("last_prompt", "")706    ss.setdefault("last_response", "")707    ss.setdefault("just_generated", False)708    709    # Sidebar (control total)710    GEN_PARAMS = sidebar_params()711    GEN_PARAMS["persona"] = st.session_state.persona  # por si acaso712 713    # Setting historial for the current user714    #if "historial" not in st.session_state:715    #    st.session_state.historial = []716 717    # Assigning a new ID to the current user718    if "user_id" not in st.session_state:719        st.session_state["user_id"] = str(uuid.uuid4())[:8]720 721    # Loading classifier encoder classes:722    labels_path = hf_hub_download(repo_id="tecuhtli/assistant-classifier-bert", filename="context_labels.pkl", use_auth_token=HF_TOKEN)723    label_classes = joblib.load(labels_path)724 725    # Loading Saved Models  726    # Modelo Contexto727    context_model = AutoModelForSequenceClassification.from_pretrained("tecuhtli/assistant-classifier-bert", use_auth_token=HF_TOKEN)    728    cont_tok = AutoTokenizer.from_pretrained("tecuhtli/assistant-classifier-bert", use_auth_token=HF_TOKEN)729    730    # Modelo Técnico731    tec_tok = AutoTokenizer.from_pretrained("tecuhtli/assistant-technical-t5", use_auth_token=HF_TOKEN) 732    tec_model = AutoModelForSeq2SeqLM.from_pretrained("tecuhtli/assistant-technical-t5", use_auth_token=HF_TOKEN) 733 734    # Modelo Social735    soc_tok = AutoTokenizer.from_pretrained("tecuhtli/assistant-social-t5", use_auth_token=HF_TOKEN) 736    soc_model = AutoModelForSeq2SeqLM.from_pretrained("tecuhtli/assistant-social-t5", use_auth_token=HF_TOKEN) 737 738    # Available Device739    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")740 741    # Defining Assistant Presentation742    st.title("🤖 Your Personal Assistant 🎓")743 744    st.caption("🙋🏽‍ You can ask me about technical concepts such as visualization, data cleaning, BI, and more.")745    st.caption("🙇🏽 I can *only* understand and answer in Spanish (🦅🇲🇽🌵).")746    st.caption("➡️ At this stage, I can respond to simple questions such as:")747    st.caption("   • ¿Cómo estás?   • ¿Qué es...?   • Explícame algo   • Define algo   • ¿Para qué sirve...?")748 749    st.caption("😊 If you want to know me better, visit: [hazutecuhtli.github.io](https://github.com/hazutecuhtli/LLMs_FineTuned_Chatbot)")750 751    st.markdown("<br>", unsafe_allow_html=True)752 753    st.caption("✏️ Type **'salir'** to exit.")754 755    # 🔁 Limpieza segura antes del formulario756    if st.session_state.pop("_clear_entrada", False):757        if "entrada" in st.session_state:758            del st.session_state["entrada"]759 760    # 🧠 Flash de respuesta (la guardamos, pero la mostraremos después del form)761    _flash = st.session_state.pop("_flash_response", None)762 763 764    with st.form("formulario_assistant"):765        user_question = st.text_area("📝 Escribe tu pregunta aquí", key="entrada", height=100)766        submitted = st.form_submit_button("Responder")767 768    if submitted:769        if not user_question:770            st.info("Chatbot: ¿Podrías repetir eso? No entendí bien 😅")771        else:772            response, context = contextual_asnwer(773                user_question, label_classes, context_model, cont_tok,774                tec_model, tec_tok, soc_model, soc_tok, device,775                gen_params=GEN_PARAMS, block_web=True,776            )777 778            # 🧠 Guarda historial779            hora_actual = dt.datetime.now().strftime("%Y-%m-%d %H:%M:%S")780            st.session_state.historial.append(("Tú", user_question, hora_actual))781 782            hora_actual = dt.datetime.now().strftime("%Y-%m-%d %H:%M:%S")783            st.session_state.historial.append(("Assistant", response, hora_actual))784 785            # 💾 Guarda conversación786            saving_interaction(user_question, response, context, st.session_state["user_id"])787 788            # 🟩 Guarda respuesta para mostrar después del rerun789            st.session_state["_flash_response"] = response790 791            # 🧼 Limpieza del textarea en el próximo ciclo792            st.session_state["_clear_entrada"] = True793 794            # ♻️ Forzar refresh (sidebar verá el nuevo prompt)795            st.rerun()796 797    # -----------------------------------------------------------798    # 💬 Mostrar la respuesta actual (flash) justo aquí ↓↓↓799    # -----------------------------------------------------------800    if _flash:801        st.success(_flash)802 803    # Mostrar último mensaje (opcional, arriba de todo)804    #if st.session_state.get("just_generated"):805    #    if st.session_state["last_response"]:806    #        st.success(st.session_state["last_response"])807    #    st.session_state["just_generated"] = False808 809    # ... formulario y lógica de respuesta ...810 811    # 🔁 Historial con estilo chat y contenedor con scroll812    if st.session_state.historial:813        st.markdown("---")814 815        # 💾 Botón de descarga arriba del historial816        lineas = []817        for msg in reversed(st.session_state.historial):818            if len(msg) == 3:819                autor, texto, hora = msg820                lineas.append(f"[{hora}] {autor}: {texto}")821            else:822                autor, texto = msg823                lineas.append(f"{autor}: {texto}")824        texto_chat = "\n\n".join(lineas)825 826        st.download_button(827            label="💾 Descargar conversación como .txt",828            data=texto_chat,829            file_name="conversacion_assistant.txt",830            mime="text/plain",831            use_container_width=True832        )833 834        # 🪟 Contenedor con scroll y burbujas835        st.markdown(836            """837            <div id="chat-container" style="838                max-height: 400px;839                overflow-y: auto;840                padding: 10px;841                border: 1px solid #333;842                border-radius: 10px;843                background: linear-gradient(180deg, #0e0e0e 0%, #1b1b1b 100%);844                margin-top: 10px;845            ">846            """,847            unsafe_allow_html=True848        )849 850        for msg in reversed(st.session_state.historial):851            if len(msg) == 3:852                autor, texto, _ = msg853            else:854                autor, texto = msg855 856            if autor == "Tú":857                st.markdown(858                    f"""859                    <div style="860                        text-align: right;861                        background-color: #2d2d2d;862                        color: #e6e6e6;863                        padding: 10px 14px;864                        border-radius: 12px;865                        margin: 6px 0;866                        border: 1px solid #3a3a3a;867                        display: inline-block;868                        max-width: 80%;869                        float: right;870                        clear: both;871                    ">872                        🧍‍♂️ <b>{autor}:</b> {texto}873                    </div>874                    """,875                    unsafe_allow_html=True876                )877            else:878                st.markdown(879                    f"""880                    <div style="881                        text-align: left;882                        background-color: #162b1f;883                        color: #d9ead3;884                        padding: 10px 14px;885                        border-radius: 12px;886                        margin: 6px 0;887                        border: 1px solid #264d36;888                        display: inline-block;889                        max-width: 80%;890                        float: left;891                        clear: both;892                    ">893                        🤖 <b>{autor}:</b> {texto}894                    </div>895                    """,896                    unsafe_allow_html=True897                )898 899        st.markdown("</div>", unsafe_allow_html=True)900                901#***************************************************************************902# FIN903#***************************************************************************904