tecuhtli/assistant-t5-qa-data-processing
0
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 