codekingpro/portable-devtools
114k
1import os2from typing import Any, Dict, List, Optional3 4from langchain_core.embeddings import Embeddings5from pydantic import BaseModel, ConfigDict, model_validator6 7 8class AscendEmbeddings(Embeddings, BaseModel):9 """10 Ascend NPU accelerate Embedding model11 12 Please ensure that you have installed CANN and torch_npu.13 14 Example:15 16 from langchain_community.embeddings import AscendEmbeddings17 model = AscendEmbeddings(model_path=<path_to_model>,18 device_id=0,19 query_instruction="Represent this sentence for searching relevant passages: "20 )21 """22 23 """model path"""24 model_path: str25 """Ascend NPU device id."""26 device_id: int = 027 """Unstruntion to used for embedding query."""28 query_instruction: str = ""29 """Unstruntion to used for embedding document."""30 document_instruction: str = ""31 use_fp16: bool = True32 pooling_method: Optional[str] = "cls"33 batch_size: int = 3234 model: Any35 tokenizer: Any36 37 model_config = ConfigDict(protected_namespaces=())38 39 def __init__(self, *args: Any, **kwargs: Any) -> None:40 super().__init__(*args, **kwargs)41 try:42 from transformers import AutoModel, AutoTokenizer43 except ImportError as e:44 raise ImportError(45 "Unable to import transformers, please install with "46 "`pip install -U transformers`."47 ) from e48 try:49 self.model = AutoModel.from_pretrained(self.model_path).npu().eval()50 self.tokenizer = AutoTokenizer.from_pretrained(self.model_path)51 except Exception as e:52 raise Exception(53 f"Failed to load model [self.model_path], due to following error:{e}"54 )55 56 if self.use_fp16:57 self.model.half()58 self.encode([f"warmup {i} times" for i in range(10)])59 60 @model_validator(mode="before")61 @classmethod62 def validate_environment(cls, values: Dict) -> Any:63 if "model_path" not in values:64 raise ValueError("model_path is required")65 if not os.access(values["model_path"], os.F_OK):66 raise FileNotFoundError(67 f"Unable to find valid model path in [{values['model_path']}]"68 )69 try:70 import torch_npu71 except ImportError:72 raise ModuleNotFoundError("torch_npu not found, please install torch_npu")73 except Exception as e:74 raise e75 try:76 torch_npu.npu.set_device(values["device_id"])77 except Exception as e:78 raise Exception(f"set device failed due to {e}")79 return values80 81 def encode(self, sentences: Any) -> Any:82 inputs = self.tokenizer(83 sentences,84 padding=True,85 truncation=True,86 return_tensors="pt",87 max_length=512,88 )89 try:90 import torch91 except ImportError as e:92 raise ImportError(93 "Unable to import torch, please install with `pip install -U torch`."94 ) from e95 last_hidden_state = self.model(96 inputs.input_ids.npu(), inputs.attention_mask.npu(), return_dict=True97 ).last_hidden_state98 tmp = self.pooling(last_hidden_state, inputs["attention_mask"].npu())99 embeddings = torch.nn.functional.normalize(tmp, dim=-1)100 return embeddings.cpu().detach().numpy()101 102 def pooling(self, last_hidden_state: Any, attention_mask: Any = None) -> Any:103 try:104 import torch105 except ImportError as e:106 raise ImportError(107 "Unable to import torch, please install with `pip install -U torch`."108 ) from e109 if self.pooling_method == "cls":110 return last_hidden_state[:, 0]111 elif self.pooling_method == "mean":112 s = torch.sum(113 last_hidden_state * attention_mask.unsqueeze(-1).float(), dim=-1114 )115 d = attention_mask.sum(dim=1, keepdim=True).float()116 return s / d117 else:118 raise NotImplementedError(119 f"Pooling method [{self.pooling_method}] not implemented"120 )121 122 def embed_documents(self, texts: List[str]) -> List[List[float]]:123 try:124 import numpy as np125 except ImportError as e:126 raise ImportError(127 "Unable to import numpy, please install with `pip install -U numpy`."128 ) from e129 embedding_list = []130 for i in range(0, len(texts), self.batch_size):131 texts_ = texts[i : i + self.batch_size]132 emb = self.encode([self.document_instruction + text for text in texts_])133 embedding_list.append(emb)134 return np.concatenate(embedding_list)135 136 def embed_query(self, text: str) -> List[float]:137 return self.encode([self.query_instruction + text])[0]138 