
    g
                         d Z ddlZddlZddlZddlmZ  G d d          Z G d de          Z G d d	e          ZdS )
z=
This contrib module contains Pytorch code for quantization.
    N)
clusteringc                   &    e Zd Zd Zd Zd Zd ZdS )	Quantizerc                 "    || _         || _        dS )za
        d: dimension of vectors
        code_size: nb of bytes of the code (per vector)
        N)d	code_size)selfr   r   s      ^/var/www/html/mpstechhub/venv/lib/python3.11/site-packages/faiss/contrib/torch/quantization.py__init__zQuantizer.__init__   s    
 "    c                     dS )z;
        takes a n-by-d array and peforms training
        N r	   xs     r
   trainzQuantizer.train   	     	r   c                     dS )zV
        takes a n-by-d float array, encodes to an n-by-code_size uint8 array
        Nr   r   s     r
   encodezQuantizer.encode!   r   r   c                     dS )zL
        takes a n-by-code_size uint8 array, returns a n-by-d array
        Nr   )r	   codess     r
   decodezQuantizer.decode'   r   r   N__name__
__module____qualname__r   r   r   r   r   r   r
   r   r      sP        # # #        r   r   c                       e Zd Zd Zd ZdS )VectorQuantizerc                     t          t          j        t          j        |          dz                      }t
                              ||           || _        d S )N   )intmathceiltorchlog2r   r   k)r	   r   r%   r   s       r
   r   zVectorQuantizer.__init__0   sG    	%*Q--!"34455	1i(((r   c                     d S Nr   r   s     r
   r   zVectorQuantizer.train6   s    r   N)r   r   r   r   r   r   r   r
   r   r   .   s2              r   r   c                   &    e Zd Zd Zd Zd Zd ZdS )ProductQuantizerc                     ||z  dk    sJ |dk    sJ t          t          j        ||z  dz                      }t                              | ||           || _        || _        || _        dS )zj M: number of subvectors, d%M == 0
        nbits: number of bits that each vector is encoded into
        r   r   N)r    r!   r"   r   r   Mnbitsr   )r	   r   r+   r,   r   s        r
   r   zProductQuantizer.__init__;   sr     1uzzzzzzzz	!e)a-0011	4I...
"r   c                    d| j         z  }| j        | j        z  }|j        }|j        }t          j        | j        ||f||          | _        t          | j                  D ]~}|d d || j        z  | j        z  |dz   | j        z  | j        z  f         }t          j
        |                                          }t          j        d| j         z  |          | j        |<   d S )N   )devicedtype   )r,   r   r+   r/   r0   r#   zeroscodebookranger   DatasetAssign
contiguouskmeans)	r	   r   ncsddevr0   mxsubdatas	            r
   r   zProductQuantizer.trainG   s    $*_VtvhTVR$4SNNNtv 	H 	HAQQQDF
df,q1u.>$&.HHHID+DOO,=,=>>D)0dj$GGDM!	H 	Hr   c                    t          j        |j        d         | j        ft           j                  }t          | j                  D ]}|d d || j        z  | j        z  |dz   | j        z  | j        z  f         }t          j	        |
                                | j        |         d          \  }}|                                |d d |f<   |S )Nr   )r0   r1   )r#   r2   shaper   uint8r4   r+   r   faissknnr6   r3   ravel)r	   r   r   r;   r<   _Is          r
   r   zProductQuantizer.encodeR   s    QWQZ8LLLtv 	$ 	$AQQQDF
df,a!etv-=-GGGHD9T__..a0@!DDDAq''))E!!!Q$KKr   c                     fdt           j                  D              fdt           j                  D             }t          j        |d          } j        j        d         }|                    d| j        z            }|S )Nc                 L    g | ] }d d |f                                          !S r'   )long).0r;   r   s     r
   
<listcomp>z+ProductQuantizer.decode.<locals>.<listcomp>[   s1    :::qaaad  "":::r   c                 @    g | ]}j         ||         d d f         S r'   )r3   )rI   r;   idxsr	   s     r
   rJ   z+ProductQuantizer.decode.<locals>.<listcomp>\   s.    GGGA4=DGQQQ/GGGr   r1   )dim)r4   r+   r#   stackr3   r?   reshape)r	   r   vectorsstacked_vectorscbdx_recrL   s   ``    @r
   r   zProductQuantizer.decodeZ   s    ::::E$&MM:::GGGGGtvGGG+g1555m!"%''C$&L99r   Nr   r   r   r
   r)   r)   :   sS        
# 
# 
#	H 	H 	H      r   r)   )	__doc__r#   rA   r!   faiss.contrib.torchr   r   r   r)   r   r   r
   <module>rW      s        * * * * * *       :	 	 	 	 	i 	 	 	& & & & &y & & & & &r   