import torch from model import SupernovaNepaliEncoder, SupernovaEncoderConfig def load_model(model_path='.'): config = SupernovaEncoderConfig.from_pretrained(model_path) model = SupernovaNepaliEncoder(config) # Load weights if necessary, or use from_pretrained return model def get_sana_embeddings(model, input_ids, attention_mask=None): model.eval() with torch.no_grad(): # Returns [batch, seq_len, 2304] embeddings = model(input_ids, attention_mask=attention_mask) return embeddings