| 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 | |