
    Nj              
           d dl Z d dlZd dlmZ 	 	 d
dej                  dej                  deej                  z  deez  fdZ G d d	ej                        Zy)    N)nnqkoffset
max_periodc                    | j                   \  }}}}|j                   \  }}	}
}|||f||	|fk(  sJ |dkD  sJ |dz  dk(  sJ |dkD  sJ t        j                  |dz  | j                  t        j                        }t        j
                  |t        j                  |       dz  |z  z        }t        j                  || j                  t        j                        }||z  }|j                  ddd      }| j                  ||||dz  d      } |j                  |||
|dz  d      }| d   j                         }| d   j                         }|d   j                         }|d   j                         }t        j                  ||z        }t        j                  ||z        }||z  ||z  z
  }||z  ||z  z   }||z  ||z  z
  }||z  ||z  z   }| j                  }t        j                  |j                  |      |j                  |      gd      }t        j                  |j                  |      |j                  |      gd      }|j                  ||||      |j                  |||
|      fS )	z
    Args:
        q (torch.Tensor): Queries, shape `[B, T, H, D]`.
        k (torch.Tensor): Keys, shape `[B, T, H, D]`.
        offset (int): Current offset, e.g. when streaming.
        max_period (float): Maximum period for the cos and sin.
    r      )devicedtype   ).r   ).r   )dim)shapetorcharanger
   float32expmathlogviewfloatcossinr   stackto)r   r   r   r   BTHDBkTkHkDkdsfreqstsqrqikrkirotrrotiqorqoikorkoir   qokos                               h/Users/ahmed/devFolder/Ultron/claude-voice/.venv/lib/python3.12/site-packages/pocket_tts/modules/rope.py
apply_roper4      sC    JAq!QWWNBBq!9R$$$q5L5q5A::>>	a1fQXXU]]	CBIIbTXXj11A59:;E 
a	>B&LB	Q	B	q!QQ"A	q!Ra#A 
6	B	
6	B	
6	B	
6	B99URZ D99URZ D
t)b4i
C
t)b4i
C
t)b4i
C
t)b4i
CGGE	cffUmSVVE]3	<B	cffUmSVVE]3	<B771aA1b! 444    c                        e Zd ZdZddeez  f fdZdej                  dej                  dej                  ez  fdZ	 xZ
S )	RotaryEmbeddingzRotary positional embedding (RoPE) from [Su et al 2022](https://arxiv.org/abs/2104.09864).

    Args:
        max_period (float): Maximum period of the rotation frequencies.
    r   c                 0    t         |           || _        y )N)super__init__r   )selfr   	__class__s     r3   r:   zRotaryEmbedding.__init__D   s    $r5   r   r   r   c                 2    t        |||| j                        S )z+Apply rope rotation to query or key tensor.)r4   r   )r;   r   r   r   s       r3   forwardzRotaryEmbedding.forwardH   s    !Q88r5   )g     @)__name__
__module____qualname____doc__r   intr:   r   Tensorr>   __classcell__)r<   s   @r3   r7   r7   =   sC    %53; %9 9%,, 9s@R 9r5   r7   )r   i'  )	r   r   r   rD   rC   r   r4   Moduler7    r5   r3   <module>rH      se       "#$	35||35||35 %,,35 e	35l9bii 9r5   