o
    @Tj                     @   s<   d Z ddlZddlZddlmZ eeZG dd dZ	dS )zJInsightFace buffalo_l wrapper for face detection and embedding extraction.    N)FaceAnalysisc                   @   s\   e Zd ZdZddededefdd	Zd
ej	de
ej	 fddZd
ej	dej	dB fddZdS )FaceAnalyzerzKWraps insightface FaceAnalysis (buffalo_l) for detection + 512-d embedding.	buffalo_l333333?
model_namectx_id
det_threshc                 C   sV   t |dgd| _| jj|dd t| jdr"t| jjdr"|| jj_td|| d S )NCPUExecutionProvider)name	providers)  r   )r   det_size	det_modelr	   z2FaceAnalyzer initialized (model=%s, det_thresh=%s))r   _apppreparehasattrr   r	   loggerinfo)selfr   r   r	    r   7/home/ubuntu/htdocs/ootd-api/app/utils/face_analyzer.py__init__   s
   
zFaceAnalyzer.__init__imagereturnc                 C   s>   | j |}|s
g S g }|D ]}|j}||tj q|S )zDetect faces and return list of 512-d embedding vectors.

        Args:
            image: BGR numpy array (OpenCV format).

        Returns:
            List of normalized embedding vectors (one per detected face).
        )r   getnormed_embeddingappendastypenpfloat32)r   r   faces
embeddingsfaceembr   r   r   get_embeddings   s   	zFaceAnalyzer.get_embeddingsNc                 C   s2   | j |}|s
dS t|dd d}|jtjS )zReturn embedding of the largest detected face (by bbox area).

        Args:
            image: BGR numpy array (OpenCV format).

        Returns:
            512-d embedding vector, or None if no face detected.
        Nc                 S   s(   | j d | j d  | j d | j d   S )N   r         )bbox)fr   r   r   <lambda>4   s   ( z9FaceAnalyzer.get_largest_face_embedding.<locals>.<lambda>)key)r   r   maxr   r   r   r    )r   r   r!   largestr   r   r   get_largest_face_embedding(   s
   	z'FaceAnalyzer.get_largest_face_embedding)r   r   r   )__name__
__module____qualname____doc__strintfloatr   r   ndarraylistr%   r/   r   r   r   r   r      s
    r   )
r3   loggingnumpyr   insightface.appr   	getLoggerr0   r   r   r   r   r   r   <module>   s    
