2024-07-21 18:15:28 +08:00
import os
2025-02-23 16:38:56 +08:00
import json
import requests
from FlagEmbedding import FlagModel
2025-03-03 21:57:27 +08:00
from zhipuai import ZhipuAI
2024-07-09 05:04:20 +08:00
2025-03-20 19:51:46 +08:00
from src import config
2025-03-04 13:49:00 +08:00
from src . utils import hashstr , logger , get_docker_safe_url
2024-07-09 05:04:20 +08:00
2025-03-04 13:49:00 +08:00
class BaseEmbeddingModel :
2025-03-03 21:57:27 +08:00
embed_state = { }
2025-03-20 19:51:46 +08:00
def get_dimension ( self ) :
if hasattr ( self , " dimension " ) :
return self . dimension
2025-04-06 20:47:45 +08:00
if hasattr ( self , " embed_model_fullname " ) :
2025-04-06 20:33:15 +08:00
return config . embed_model_names [ self . embed_model_fullname ] . get ( " dimension " , None )
2025-03-20 19:51:46 +08:00
2025-04-06 20:33:15 +08:00
return config . embed_model_names [ self . model ] . get ( " dimension " , None )
2025-03-04 13:49:00 +08:00
def encode ( self , message ) :
return self . predict ( message )
def encode_queries ( self , queries ) :
return self . predict ( queries )
2025-03-03 21:57:27 +08:00
def batch_encode ( self , messages , batch_size = 20 ) :
2025-03-07 01:05:50 +08:00
logger . info ( f " Batch encoding { len ( messages ) } messages " )
2025-03-03 21:57:27 +08:00
data = [ ]
if len ( messages ) > batch_size :
task_id = hashstr ( messages )
self . embed_state [ task_id ] = {
' status ' : ' in-progress ' ,
' total ' : len ( messages ) ,
' progress ' : 0
}
for i in range ( 0 , len ( messages ) , batch_size ) :
group_msg = messages [ i : i + batch_size ]
logger . info ( f " Encoding { i } to { i + batch_size } with { len ( messages ) } messages " )
2025-03-11 16:26:55 +08:00
response = self . encode ( group_msg )
logger . debug ( f " Response: { len ( response ) =} , { len ( group_msg ) =} , { len ( response [ 0 ] ) =} " )
2025-03-03 21:57:27 +08:00
data . extend ( response )
if len ( messages ) > batch_size :
self . embed_state [ task_id ] [ ' progress ' ] = len ( messages )
self . embed_state [ task_id ] [ ' status ' ] = ' completed '
return data
2025-03-04 13:49:00 +08:00
class LocalEmbeddingModel ( FlagModel , BaseEmbeddingModel ) :
2025-02-23 16:38:56 +08:00
def __init__ ( self , config , * * kwargs ) :
2025-04-06 20:33:15 +08:00
info = config . embed_model_names [ config . embed_model ]
2024-07-09 05:04:20 +08:00
2025-03-04 13:49:00 +08:00
self . model = config . model_local_paths . get ( info [ " name " ] , info . get ( " local_path " ) )
self . model = self . model or info [ " name " ]
2025-03-20 19:51:46 +08:00
self . dimension = info [ " dimension " ]
self . embed_model_fullname = config . embed_model
2025-03-04 13:49:00 +08:00
2025-03-11 16:26:55 +08:00
if os . path . exists ( _path := os . path . join ( os . getenv ( " MODEL_DIR " ) , self . model ) ) :
self . model = _path
2025-03-07 01:05:50 +08:00
logger . info ( f " Loading local model ` { info [ ' name ' ] } ` from ` { self . model } ` with device ` { config . device } ` " )
2025-03-04 13:49:00 +08:00
super ( ) . __init__ ( self . model ,
2025-02-23 16:38:56 +08:00
query_instruction_for_retrieval = info . get ( " query_instruction " , None ) ,
2025-03-07 01:05:50 +08:00
use_fp16 = False ,
device = config . device ,
* * kwargs )
2024-07-09 05:04:20 +08:00
2025-02-23 16:38:56 +08:00
logger . info ( f " Embedding model { info [ ' name ' ] } loaded " )
2024-07-17 18:52:20 +08:00
2024-07-21 18:15:28 +08:00
2025-03-04 13:49:00 +08:00
class ZhipuEmbedding ( BaseEmbeddingModel ) :
2024-08-25 12:34:35 +08:00
2025-02-23 16:38:56 +08:00
def __init__ ( self , config ) - > None :
self . config = config
2025-04-06 20:33:15 +08:00
self . model = config . embed_model_names [ config . embed_model ] [ " name " ]
self . dimension = config . embed_model_names [ config . embed_model ] [ " dimension " ]
2025-02-23 16:38:56 +08:00
self . client = ZhipuAI ( api_key = os . getenv ( " ZHIPUAI_API_KEY " ) )
2025-03-20 19:51:46 +08:00
self . embed_model_fullname = config . embed_model
2024-09-09 17:07:03 +08:00
2025-02-23 16:38:56 +08:00
def predict ( self , message ) :
response = self . client . embeddings . create (
model = self . model ,
input = message ,
)
data = [ a . embedding for a in response . data ]
2024-08-25 12:34:35 +08:00
return data
2024-07-22 00:00:54 +08:00
2025-03-04 13:49:00 +08:00
class OllamaEmbedding ( BaseEmbeddingModel ) :
def __init__ ( self , config ) - > None :
2025-04-06 20:33:15 +08:00
self . info = config . embed_model_names [ config . embed_model ]
2025-03-04 13:49:00 +08:00
self . model = self . info [ " name " ]
self . url = self . info . get ( " url " , " http://localhost:11434/api/embed " )
self . url = get_docker_safe_url ( self . url )
2025-03-20 19:51:46 +08:00
self . dimension = self . info . get ( " dimension " , None )
self . embed_model_fullname = config . embed_model
2024-07-22 00:00:54 +08:00
2025-03-04 13:49:00 +08:00
def predict ( self , message : list [ str ] | str ) :
if isinstance ( message , str ) :
message = [ message ]
2024-07-22 00:00:54 +08:00
2025-03-04 13:49:00 +08:00
payload = {
" model " : self . model ,
" input " : message ,
}
response = requests . request ( " POST " , self . url , json = payload )
response = json . loads ( response . text )
assert response . get ( " embeddings " ) , f " Ollama Embedding failed: { response } "
return response [ " embeddings " ]
class OtherEmbedding ( BaseEmbeddingModel ) :
2025-02-23 16:38:56 +08:00
def __init__ ( self , config ) - > None :
2025-04-06 20:33:15 +08:00
self . info = config . embed_model_names [ config . embed_model ]
2025-03-20 19:51:46 +08:00
self . embed_model_fullname = config . embed_model
self . dimension = self . info . get ( " dimension " , None )
2025-03-04 13:49:00 +08:00
self . model = self . info [ " name " ]
self . api_key = os . getenv ( self . info [ " api_key " ] , None )
self . url = get_docker_safe_url ( self . info [ " url " ] )
assert self . url and self . model , f " URL and model are required. Cur embed model: { config . embed_model } "
2025-02-23 16:38:56 +08:00
self . headers = {
2025-03-04 13:49:00 +08:00
" Authorization " : f " Bearer { self . api_key } " ,
2025-02-23 16:38:56 +08:00
" Content-Type " : " application/json "
}
2025-03-04 13:49:00 +08:00
def predict ( self , message ) :
2025-02-23 16:38:56 +08:00
payload = self . build_payload ( message )
response = requests . request ( " POST " , self . url , json = payload , headers = self . headers )
response = json . loads ( response . text )
2025-03-04 13:49:00 +08:00
assert response [ " data " ] , f " Other Embedding failed: { response } "
2025-02-23 16:38:56 +08:00
data = [ a [ " embedding " ] for a in response [ " data " ] ]
return data
def build_payload ( self , message ) :
return {
" model " : self . model ,
" input " : message ,
}
2024-07-22 00:00:54 +08:00
def get_embedding_model ( config ) :
2024-07-31 20:22:05 +08:00
if not config . enable_knowledge_base :
return None
2025-02-23 16:38:56 +08:00
provider , model_name = config . embed_model . split ( ' / ' , 1 )
2025-04-06 20:33:15 +08:00
assert config . embed_model in config . embed_model_names . keys ( ) , f " Unsupported embed model: { config . embed_model } , only support { config . embed_model_names . keys ( ) } "
2025-02-23 16:38:56 +08:00
logger . debug ( f " Loading embedding model { config . embed_model } " )
if provider == " local " :
model = LocalEmbeddingModel ( config )
2024-08-25 20:29:24 +08:00
2025-03-04 13:49:00 +08:00
elif provider == " zhipu " :
2025-02-23 16:38:56 +08:00
model = ZhipuEmbedding ( config )
2024-08-25 20:29:24 +08:00
2025-03-04 13:49:00 +08:00
elif provider == " ollama " :
model = OllamaEmbedding ( config )
else :
model = OtherEmbedding ( config )
2024-08-25 20:29:24 +08:00
2024-09-11 01:07:19 +08:00
return model
def handle_local_model ( paths , model_name , default_path ) :
model_path = paths . get ( model_name , default_path )
return model_path