| import time |
| import base64 |
| import io |
| import re |
| from PIL import Image |
| import torch |
| from .model_loader import model_loader |
| from .persona_manager import persona_manager |
| from .config import MAX_NEW_TOKENS, TEMPERATURE, TOP_K, TOP_P |
|
|
| def predict( |
| pictogram_sequence: list[str], |
| persona_id: str, |
| context: str = "daily", |
| audio_transcript: str = None, |
| image_base64: str = None, |
| n_alternatives: int = 3 |
| ) -> dict: |
| start_time = time.time() |
| |
| |
| persona = persona_manager.get_persona(persona_id) |
| if not persona: |
| return {"alternatives": ["Persona not found"], "primary": "Persona not found", "inference_time_ms": 0, "persona_id": persona_id} |
| |
| system_prompt = persona.to_system_prompt() |
| |
| |
| pictogram_str = ", ".join(pictogram_sequence) |
| audio_note = f"\nChild vocalization heard: '{audio_transcript}' — use this to refine the prediction" if audio_transcript else "" |
| image_note = "\nImage provided for visual grounding." if image_base64 else "" |
| |
| prompt = f"""<system> |
| {system_prompt} |
| </system> |
| <model_ack> |
| Understood. Ready. |
| </model_ack> |
| <user> |
| Pictogram sequence: [{pictogram_str}] |
| Context: {context}{audio_note}{image_note} |
| Predict {n_alternatives} most likely sentences. Use child's language mix. Short sentences. Numbered list. |
| </user> |
| <model> |
| """ |
| |
| |
| image = None |
| if image_base64: |
| try: |
| image_data = base64.b64decode(image_base64) |
| image = Image.open(io.BytesIO(image_data)) |
| except Exception as e: |
| print(f"Error decoding image: {e}") |
| |
| |
| if not model_loader.is_ready(): |
| return {"alternatives": ["Model not loaded"], "primary": "Model not loaded", "inference_time_ms": 0, "persona_id": persona_id} |
|
|
| |
| try: |
| model_loader.load_persona_adapter(persona_id) |
| except Exception as e: |
| return {"alternatives": [f"Error loading adapter: {str(e)}"], "primary": "Error loading adapter", "inference_time_ms": 0, "persona_id": persona_id} |
|
|
| try: |
| |
| if image: |
| inputs = model_loader.processor(text=prompt, images=image, return_tensors="pt").to(model_loader.model.device) |
| else: |
| inputs = model_loader.processor(text=prompt, return_tensors="pt").to(model_loader.model.device) |
| |
| |
| with torch.no_grad(): |
| output_ids = model_loader.model.generate( |
| **inputs, |
| max_new_tokens=MAX_NEW_TOKENS, |
| temperature=TEMPERATURE, |
| top_k=TOP_K, |
| top_p=TOP_P, |
| do_sample=True, |
| pad_token_id=model_loader.processor.tokenizer.pad_token_id or model_loader.processor.tokenizer.eos_token_id |
| ) |
| |
| |
| input_len = inputs.input_ids.shape[1] |
| response = model_loader.processor.decode(output_ids[0][input_len:], skip_special_tokens=True) |
| |
| |
| alternatives = [] |
| |
| matches = re.findall(r'\d+\.\s*(.+)', response) |
| if matches: |
| alternatives = [m.strip() for m in matches[:n_alternatives]] |
| else: |
| |
| alternatives = [line.strip() for line in response.split('\n') if line.strip()][:n_alternatives] |
| |
| |
| if not alternatives: |
| alternatives = [response.strip()] |
| |
| inference_time_ms = (time.time() - start_time) * 1000 |
| |
| return { |
| "alternatives": alternatives, |
| "primary": alternatives[0] if alternatives else "", |
| "inference_time_ms": inference_time_ms, |
| "persona_id": persona_id |
| } |
| except Exception as e: |
| print(f"Prediction error: {e}") |
| return { |
| "alternatives": ["Prediction failed — please try again"], |
| "primary": "Prediction failed", |
| "inference_time_ms": (time.time() - start_time) * 1000, |
| "persona_id": persona_id |
| } |
|
|