import os
import torch
from llmlingua import PromptCompressor

# Définition du répertoire de cache et du fichier de cache pour le modèle
cache_dir = "./cache"
model_cache_file = os.path.join(cache_dir, "prompt_compressor_model.pt")

def compress_prompt(prompt, model, rate=0.33, force_tokens=['\n', '?']):
    """
    Compresse un prompt en utilisant LLMLingua.
    
    Args:
        prompt (str): Le prompt à compresser.
        model: Le modèle de compression à utiliser.
        rate (float): Taux de compression.
        force_tokens (list): Tokens à forcer dans la compression.
        
    Returns:
        str: Le prompt compressé.
    """
    # Compression du prompt avec les paramètres spécifiés
    compressed_prompt = model.compress_prompt(prompt, rate=rate, force_tokens=force_tokens)
    
    # Retour du prompt compressé
    return compressed_prompt['compressed_prompt']

def compress_and_merge_segments(prompt, model, tokenizer, segment_length=512, rate=0.5, force_tokens=['\n', '?']):
    """
    Divise le prompt en segments, les compresse individuellement, puis fusionne les résultats.
    
    Args:
        prompt (str): Le prompt à compresser.
        model: Le modèle de compression à utiliser.
        tokenizer: Le tokenizer utilisé par le modèle.
        segment_length (int): Longueur maximale de chaque segment après tokenisation.
        rate (float): Taux de compression.
        force_tokens (list): Tokens à forcer dans la compression.
        
    Returns:
        str: Le prompt compressé après fusion des segments.
    """
    # Tokenisation du prompt pour déterminer les points de découpe appropriés
    tokens = tokenizer.tokenize(prompt)
    segments = [tokens[i:i+segment_length] for i in range(0, len(tokens), segment_length)]
    compressed_segments = []
    
    # Compression de chaque segment
    for segment_tokens in segments:
        segment = tokenizer.convert_tokens_to_string(segment_tokens)
        compressed_segment = model.compress_prompt(segment, rate=rate, force_tokens=force_tokens)
        compressed_segments.append(compressed_segment['compressed_prompt'])
    
    # Fusion des segments compressés
    compressed_prompt = " ".join(compressed_segments)
    
    return compressed_prompt


def load_model(use_small_model=False):
    """
    Charge le modèle LLMLingua-2 depuis le cache si disponible, sinon l'initialise et le sauvegarde.
    
    Args:
        use_small_model (bool): Si True, utilise le modèle LLMLingua-2-small.
    
    Returns:
        Le modèle chargé ou initialisé.
    """
    model_name = "microsoft/llmlingua-2-xlm-roberta-large-meetingbank"
    if use_small_model:
        model_name = "microsoft/llmlingua-2-bert-base-multilingual-cased-meetingbank"
    
    # Mise à jour du chemin du fichier de cache en fonction du modèle
    global model_cache_file
    model_cache_file = os.path.join(cache_dir, f"{model_name.replace('/', '_')}.pt")
    
    if os.path.exists(model_cache_file):
        # Chargement du modèle depuis le cache si disponible
        print("Chargement du modèle depuis le cache.")
        model = torch.load(model_cache_file)
    else:
        # Création du répertoire de cache s'il n'existe pas
        os.makedirs(cache_dir, exist_ok=True)
        
        # Initialisation du modèle LLMLingua-2 et sauvegarde dans le cache pour une utilisation future
        print("Initialisation et sauvegarde du modèle LLMLingua-2.")
        model = PromptCompressor(model_name=model_name, use_llmlingua2=True)
        torch.save(model, model_cache_file)
    
    return model


model = PromptCompressor(
    model_name="microsoft/llmlingua-2-xlm-roberta-large-meetingbank",
    use_llmlingua2=True, device_map="mps",
)
tokenizer = model.tokenizer  

text = """Vous devez générer 5 mots clés basés sur le texte. Les mots clés doivent être séparés par des virgules, et ne doivent pas inclure le symbole '#'. Il convient de ne fournir que les mots clés, sans aucune autre information ou contenu supplémentaire. TEXTE : [CONTENU]"""
# Exemple d'utilisation du modèle LLMLingua-2
# Pour utiliser le modèle small, passez `use_small_model=True` à `load_model`
compressed_text = compress_and_merge_segments(text, model, tokenizer, segment_length=400)
longueur_originale = len(model.tokenizer.tokenize(text))
longueur_compressée = len(model.tokenizer.tokenize(compressed_text))
pourcentage_gagné = 100 - (longueur_compressée / longueur_originale * 100)

print(f"Prompt original:\n {text}\nLongueur en token: {longueur_originale}")
print("\n")
print(f"Prompt compressé: \n {compressed_text}\nLongueur en token: {longueur_compressée}")
print("\n")
print(f"Réduction: {pourcentage_gagné:.2f}%")
