import torch from transformers import pipeline, AutoTokenizer, AutoModelForSeq2SeqLM import logging import warnings logging.basicConfig(level=logging.DEBUG, format='%(levelname)s:%(name)s:%(message)s') logger = logging.getLogger(__name__) # Suppress tokenizer warnings warnings.filterwarnings("ignore", category=UserWarning, module="transformers") warnings.filterwarnings("ignore", category=FutureWarning, module="transformers") MODEL_CACHE = {} def get_device(): """Get available device (GPU or CPU)""" if torch.cuda.is_available(): device = 0 logger.info(f"Using GPU: {torch.cuda.get_device_name(0)}") else: device = -1 logger.info("Using CPU") return device class SimpleSeq2SeqWrapper: def __init__(self, model, tokenizer, device): self.model = model self.tokenizer = tokenizer self.device = device def __call__(self, texts, max_length=128, num_beams=2): if isinstance(texts, str): texts = [texts] is_single = True else: is_single = False results = [] for text in texts: try: logger.debug(f"Processing text: {text[:50]}...") inputs = self.tokenizer(text, return_tensors="pt", max_length=512, truncation=True) logger.debug(f"Tokenized inputs shape: {inputs['input_ids'].shape}") if self.device != -1: try: inputs = {k: v.to(f"cuda:{self.device}") for k, v in inputs.items()} except Exception as device_err: logger.warning(f"Failed to move to GPU, using CPU: {device_err}") # Fall back to CPU if hasattr(self.model, 'cpu'): self.model = self.model.cpu() logger.debug(f"Generating with max_length={max_length}, num_beams={num_beams}") with torch.no_grad(): try: outputs = self.model.generate( **inputs, max_length=max_length, num_beams=num_beams, early_stopping=True ) except Exception as gen_param_err: logger.warning(f"Generation with num_beams={num_beams} failed, trying with num_beams=1: {gen_param_err}") outputs = self.model.generate( **inputs, max_length=max_length, num_beams=1 ) logger.debug(f"Generated outputs shape: {outputs.shape}") if outputs is None or len(outputs) == 0: logger.error(f"Model returned empty outputs") results.append({"generated_text": text}) continue decoded = self.tokenizer.decode(outputs[0], skip_special_tokens=True) if not decoded or decoded.strip() == "": logger.warning(f"Model returned empty string, using input") results.append({"generated_text": text}) else: logger.debug(f"Decoded result: {decoded[:50]}...") results.append({"generated_text": decoded}) except Exception as gen_err: logger.error(f"Generation error for text '{text[:50]}...': {type(gen_err).__name__}: {gen_err}", exc_info=True) results.append({"generated_text": text}) return results def load_model_with_fallback(repo_id): """ Load model with fallback strategies for problematic tokenizers. Tries multiple approaches in order of preference. """ device = get_device() # Strategy 1: Try standard pipeline (works for most models) try: logger.info(f"Strategy 1: Trying standard pipeline for {repo_id}") with warnings.catch_warnings(): warnings.simplefilter("ignore") nlp = pipeline( "text2text-generation", model=repo_id, device=device, trust_remote_code=True, model_kwargs={"torch_dtype": torch.float32} ) logger.info("Strategy 1 succeeded") return nlp except Exception as e: logger.warning(f"Strategy 1 failed: {type(e).__name__}: {e}") # Strategy 2: Try with slow tokenizer (use_fast=False) try: logger.info(f"Strategy 2: Trying with slow tokenizer for {repo_id}") with warnings.catch_warnings(): warnings.simplefilter("ignore") nlp = pipeline( "text2text-generation", model=repo_id, device=device, trust_remote_code=True, model_kwargs={"torch_dtype": torch.float32}, tokenizer_kwargs={"use_fast": False} ) logger.info("Strategy 2 succeeded") return nlp except Exception as e: logger.warning(f"Strategy 2 failed: {type(e).__name__}: {e}") # Strategy 3: Load model and tokenizer separately with error handling try: logger.info(f"Strategy 3: Loading model and tokenizer separately for {repo_id}") with warnings.catch_warnings(): warnings.simplefilter("ignore") try: tokenizer = AutoTokenizer.from_pretrained(repo_id, trust_remote_code=True, use_fast=False) except Exception as tok_err: logger.warning(f"Tokenizer loading failed, using base tokenizer: {type(tok_err).__name__}") tokenizer = AutoTokenizer.from_pretrained("facebook/mbart-large-50", use_fast=False) model = AutoModelForSeq2SeqLM.from_pretrained( repo_id, torch_dtype=torch.float32, trust_remote_code=True, device_map="auto" if device != -1 else None ) logger.info("Strategy 3 succeeded") wrapper = SimpleSeq2SeqWrapper(model, tokenizer, device) return wrapper except Exception as e: logger.warning(f"Strategy 3 failed: {type(e).__name__}: {e}") # Strategy 4: Try with mt5 base tokenizer fallback for mt5 models try: logger.info(f"Strategy 4: Trying with mt5 fallback for {repo_id}") with warnings.catch_warnings(): warnings.simplefilter("ignore") try: tokenizer = AutoTokenizer.from_pretrained(repo_id, trust_remote_code=True, use_fast=False) except Exception: logger.warning(f"Tokenizer loading failed, using mt5 base tokenizer") tokenizer = AutoTokenizer.from_pretrained("google/mt5-small", use_fast=False) model = AutoModelForSeq2SeqLM.from_pretrained( repo_id, torch_dtype=torch.float32, trust_remote_code=True, device_map="auto" if device != -1 else None ) logger.info("Strategy 4 succeeded") wrapper = SimpleSeq2SeqWrapper(model, tokenizer, device) return wrapper except Exception as e: logger.warning(f"Strategy 4 failed: {type(e).__name__}: {e}") # Strategy 5: Minimal approach - model only without pipeline try: logger.info(f"Strategy 5: Loading model with default tokenizer for {repo_id}") with warnings.catch_warnings(): warnings.simplefilter("ignore") model = AutoModelForSeq2SeqLM.from_pretrained( repo_id, torch_dtype=torch.float32, trust_remote_code=True, device_map="auto" if device != -1 else None ) tokenizer = AutoTokenizer.from_pretrained(repo_id, trust_remote_code=True) logger.info("Strategy 5 succeeded") wrapper = SimpleSeq2SeqWrapper(model, tokenizer, device) return wrapper except Exception as e: logger.error(f"Strategy 5 failed: {type(e).__name__}: {e}") raise Exception(f"All loading strategies failed for {repo_id}: {e}") def load_model(model_id): """Load a model from HuggingFace with caching and fallback strategies""" # Extract repo_id from URL if needed if model_id.startswith("http"): repo_id = "/".join(model_id.split("/")[-2:]) else: repo_id = model_id if repo_id in MODEL_CACHE: logger.info(f"Using cached model: {repo_id}") return MODEL_CACHE[repo_id] try: logger.info(f"Loading model: {repo_id}") nlp = load_model_with_fallback(repo_id) MODEL_CACHE[repo_id] = nlp logger.info(f"Successfully loaded: {repo_id}") return nlp except Exception as e: logger.error(f"Error loading model {repo_id}: {e}") raise def correct_text(text, model_id, max_length=128): """ Correct Sinhala text using specified model Args: text: Input text (Sinhala sentence) model_id: HuggingFace model URL or ID max_length: Maximum length of generated text Returns: Corrected text """ try: nlp = load_model(model_id) logger.debug(f"Model loaded successfully, type: {type(nlp)}") # Try with text generation try: with warnings.catch_warnings(): warnings.simplefilter("ignore") logger.debug(f"Calling model with text: {text[:50]}...") result = nlp(text, max_length=max_length, num_beams=2) logger.debug(f"Model returned: {result}") if isinstance(result, list) and len(result) > 0: logger.debug(f"Result is list with {len(result)} items") if isinstance(result[0], dict) and 'generated_text' in result[0]: output = result[0]['generated_text'] logger.debug(f"Extracted generated_text: {output[:50]}...") return output else: logger.warning(f"Result[0] doesn't have expected format: {result[0]}") else: logger.warning(f"Result is not a non-empty list: {result}") return text # Fallback to original text except Exception as gen_error: logger.warning(f"Generation failed, returning input: {type(gen_error).__name__}: {gen_error}", exc_info=True) return text # Fallback to input text except Exception as e: logger.error(f"Error correcting text: {type(e).__name__}: {e}", exc_info=True) return f"Error: {str(e)}" def batch_correct_texts(texts, model_id, batch_size=8): """ Correct multiple texts using a model (for evaluation) Args: texts: List of input texts model_id: HuggingFace model URL or ID batch_size: Batch size for processing (smaller for stability) Returns: List of corrected texts """ try: nlp = load_model(model_id) results = [] for i in range(0, len(texts), batch_size): batch = texts[i:i + batch_size] try: with warnings.catch_warnings(): warnings.simplefilter("ignore") batch_results = nlp(batch, max_length=128, num_beams=2) # Extract generated texts for r in batch_results: if isinstance(r, dict) and 'generated_text' in r: results.append(r['generated_text']) else: # Fallback: return original text idx = len(results) if idx < len(batch): results.append(batch[idx]) except Exception as batch_error: logger.warning(f"Batch processing error, using input texts: {batch_error}") # Fallback: return original texts for this batch results.extend(batch) return results except Exception as e: logger.error(f"Error in batch correction: {e}") # Fallback: return input texts return texts def get_model_name(model_url): """Extract model name from URL""" return model_url.split('/')[-1]