""" Load BioBBC model from Hugging Face Hub and use for prediction """ import torch import json from huggingface_hub import hf_hub_download from transformers import AutoTokenizer from model_biobbc import BioBBCModel from pos_embeddings import POSTagger import argparse class BioBBCPredictor: """Predictor that loads model from Hugging Face Hub""" def __init__(self, repo_id, device=None): """ Load BioBBC model from Hugging Face Hub Args: repo_id: Hugging Face repository ID (e.g., 'username/biobbc-biomarker-ner') device: Device to run on ('cuda' or 'cpu') """ self.device = device if device else ('cuda' if torch.cuda.is_available() else 'cpu') print(f"Loading BioBBC model from: {repo_id}") print(f"Device: {self.device}") # Download model files print("\nDownloading model files...") model_path = hf_hub_download(repo_id=repo_id, filename="pytorch_model.bin") config_path = hf_hub_download(repo_id=repo_id, filename="config.json") vocab_path = hf_hub_download(repo_id=repo_id, filename="vocabularies.bin") # Load configuration print("Loading configuration...") with open(config_path, 'r') as f: self.config = json.load(f) # Load vocabularies print("Loading vocabularies...") vocabs = torch.load(vocab_path, map_location='cpu') self.char_vocab = vocabs['char_vocab'] self.pos_vocab = vocabs['pos_vocab'] self.domain_vocab = vocabs['domain_vocab'] # Initialize tokenizer print("Loading tokenizer...") self.tokenizer = AutoTokenizer.from_pretrained(self.config['bert_model']) # Initialize POS tagger print("Initializing POS tagger...") self.pos_tagger = POSTagger(use_biomedical=False) # Use general tagger # Create model print("Creating model...") self.model = BioBBCModel( bert_model_name=self.config['bert_model'], hidden_dim=self.config['hidden_dim'], num_layers=self.config['num_layers'], num_labels=self.config['num_labels'], char_vocab_size=len(self.char_vocab), pos_vocab_size=len(self.pos_vocab), domain_vocab=self.domain_vocab, use_char=self.config['use_char_embeddings'], use_pos=self.config['use_pos_embeddings'], use_domain=self.config['use_domain_embeddings'], char_embedding_dim=self.config['char_embedding_dim'], char_hidden_dim=self.config['char_hidden_dim'], pos_embedding_dim=self.config['pos_embedding_dim'], domain_embedding_dim=self.config['domain_embedding_dim'], dropout=self.config['dropout'] ) # Load weights print("Loading model weights...") state_dict = torch.load(model_path, map_location='cpu') self.model.load_state_dict(state_dict) self.model = self.model.to(self.device) self.model.eval() # Label mapping self.id2label = {int(k): v for k, v in self.config['id2label'].items()} print(f"\n✓ Model loaded successfully!") print(f" F1 Score: {self.config.get('best_f1', 'N/A')}") print(f" Trained Epochs: {self.config.get('total_epochs_trained', 'N/A')}") print(f" Labels: {self.config['labels']}") def predict(self, text, return_entities=True): """ Predict biomarkers in text Args: text: Input text string return_entities: If True, return entity strings; if False, return tokens with labels Returns: List of (entity, start_char, end_char) if return_entities=True List of (token, label) if return_entities=False """ # Tokenize tokens = text.split() # Get BERT encodings encoding = self.tokenizer( tokens, is_split_into_words=True, max_length=self.config['max_len'], padding='max_length', truncation=True, return_tensors='pt' ) # Get character IDs char_ids = self._get_char_ids(tokens) # Get POS IDs pos_ids = self._get_pos_ids(tokens) # Get domain IDs domain_ids = self._get_domain_ids(tokens) # Move to device input_ids = encoding['input_ids'].to(self.device) attention_mask = encoding['attention_mask'].to(self.device) char_ids = char_ids.to(self.device) pos_ids = pos_ids.to(self.device) domain_ids = domain_ids.to(self.device) # Predict with torch.no_grad(): predictions = self.model( input_ids=input_ids, attention_mask=attention_mask, char_ids=char_ids, pos_ids=pos_ids, domain_ids=domain_ids ) # Convert predictions to labels pred_labels = [self.id2label[p] for p in predictions[0][:len(tokens)]] if return_entities: # Extract entities entities = [] current_entity = [] start_idx = 0 for i, (token, label) in enumerate(zip(tokens, pred_labels)): if label.startswith('B-'): if current_entity: entity_text = ' '.join(current_entity) entities.append((entity_text, start_idx, start_idx + len(entity_text))) current_entity = [token] start_idx = text.find(token, start_idx if not entities else entities[-1][2]) elif label.startswith('I-') and current_entity: current_entity.append(token) else: if current_entity: entity_text = ' '.join(current_entity) entities.append((entity_text, start_idx, start_idx + len(entity_text))) current_entity = [] if current_entity: entity_text = ' '.join(current_entity) entities.append((entity_text, start_idx, start_idx + len(entity_text))) return entities else: return list(zip(tokens, pred_labels)) def _get_char_ids(self, tokens): """Convert tokens to character IDs""" max_word_len = 20 char_ids = [] for token in tokens: token_chars = [] for char in token[:max_word_len]: char_id = self.char_vocab.char2id.get(char, 0) token_chars.append(char_id) # Pad to max_word_len while len(token_chars) < max_word_len: token_chars.append(0) char_ids.append(token_chars) # Pad to max sequence length while len(char_ids) < self.config['max_len']: char_ids.append([0] * max_word_len) return torch.tensor([char_ids[:self.config['max_len']]]) def _get_pos_ids(self, tokens): """Convert tokens to POS tag IDs""" pos_ids = self.pos_tagger.tag_tokens(tokens) # Pad to max sequence length while len(pos_ids) < self.config['max_len']: pos_ids.append(0) return torch.tensor([pos_ids[:self.config['max_len']]]) def _get_domain_ids(self, tokens): """Convert tokens to domain word IDs""" domain_ids = [self.domain_vocab.word2id.get(token.lower(), 0) for token in tokens] # Pad to max sequence length while len(domain_ids) < self.config['max_len']: domain_ids.append(0) return torch.tensor([domain_ids[:self.config['max_len']]]) def main(): parser = argparse.ArgumentParser(description='Predict biomarkers using BioBBC model from HF Hub') parser.add_argument('--repo_id', type=str, required=True, help='Hugging Face repository ID (e.g., username/biobbc-biomarker-ner)') parser.add_argument('--text', type=str, help='Text to analyze (or use --file)') parser.add_argument('--file', type=str, help='File containing text to analyze') parser.add_argument('--interactive', action='store_true', help='Interactive mode - enter text repeatedly') args = parser.parse_args() # Load model predictor = BioBBCPredictor(args.repo_id) # Interactive mode if args.interactive: print("\n" + "="*70) print("Interactive Biomarker Prediction") print("="*70) print("Enter text to analyze (or 'quit' to exit)\n") while True: text = input("Text: ").strip() if text.lower() in ['quit', 'exit', 'q']: break if text: entities = predictor.predict(text) if entities: print(f"\n✓ Found {len(entities)} biomarker(s):") for entity, start, end in entities: print(f" - {entity} (char {start}-{end})") else: print("\n✗ No biomarkers found") print() # File mode elif args.file: with open(args.file, 'r') as f: text = f.read() print(f"\nAnalyzing file: {args.file}") print(f"Text: {text[:100]}...\n") entities = predictor.predict(text) print(f"\n✓ Found {len(entities)} biomarker(s):") for entity, start, end in entities: print(f" - {entity} (char {start}-{end})") # Single text mode elif args.text: print(f"\nAnalyzing: {args.text}\n") entities = predictor.predict(args.text) if entities: print(f"✓ Found {len(entities)} biomarker(s):") for entity, start, end in entities: print(f" - {entity} (char {start}-{end})") else: print("✗ No biomarkers found") else: # Demo examples print("\n" + "="*70) print("Demo: Biomarker Prediction") print("="*70 + "\n") examples = [ "Elevated levels of IL-6 and TNF-alpha were observed in patients.", "Serum glucose and HbA1c levels were measured.", "Cardiac troponin and BNP levels indicate myocardial injury.", "The study measured CRP, IL-10, and interferon-gamma concentrations." ] for i, text in enumerate(examples, 1): print(f"\nExample {i}: {text}") entities = predictor.predict(text) if entities: print(f" ✓ Biomarkers: {', '.join([e[0] for e in entities])}") else: print(" ✗ No biomarkers found") if __name__ == '__main__': main()