(self, q, k, v, mask, use_qknorm=False)
| 130 | return paddle.to_tensor(batch_id_per_token, dtype="int32"), cu_seqlens_q, cu_seqlens_k |
| 131 | |
| 132 | def ref_attention(self, q, k, v, mask, use_qknorm=False): |
| 133 | if use_qknorm: |
| 134 | q = q.reshape([-1, self.head_dim]) |
| 135 | q = fused_rms_norm(q.astype("float32"), self.q_norm_weight_tensor, None, 1e-6)[0].astype(self.dtype) |
| 136 | q = q.reshape([self.bsz, -1, self.num_q_head, self.head_dim]) |
| 137 | q = q.transpose([0, 2, 1, 3]) |
| 138 | if len(k) > 1: |
| 139 | k = paddle.concat(k, axis=1) |
| 140 | else: |
| 141 | k = k[0] |
| 142 | if use_qknorm: |
| 143 | k = k.reshape([-1, self.head_dim]) |
| 144 | k = fused_rms_norm(k.astype("float32"), self.k_norm_weight_tensor, None, 1e-6)[0].astype(self.dtype) |
| 145 | k = k.reshape([self.bsz, -1, self.num_kv_head, self.head_dim]) |
| 146 | k = k.transpose([0, 2, 1, 3]) |
| 147 | if len(v) > 1: |
| 148 | v = paddle.concat(v, axis=1) |
| 149 | else: |
| 150 | v = v[0] |
| 151 | v = v.transpose([0, 2, 1, 3]) |
| 152 | total_len = k.shape[2] |
| 153 | |
| 154 | scores = ( |
| 155 | q.reshape([self.bsz, self.num_kv_head, -1, self.head_dim]) |
| 156 | @ k.transpose([0, 1, 3, 2]) |
| 157 | * (1.0 / math.sqrt(self.head_dim)) |
| 158 | ) |
| 159 | scores = scores.reshape([self.bsz, self.num_q_head, -1, total_len]) |
| 160 | |
| 161 | if mask is not None: |
| 162 | if mask.ndim == 2: |
| 163 | mask = mask.unsqueeze(0).unsqueeze(0) |
| 164 | elif mask.ndim == 3: |
| 165 | mask = mask.unsqueeze(1) |
| 166 | scores = paddle.add(scores, mask) |
| 167 | weights = F.softmax(scores, axis=-1) |
| 168 | |
| 169 | o = weights.reshape([self.bsz, self.num_kv_head, -1, total_len]) @ v |
| 170 | return ( |
| 171 | o.reshape([self.bsz, self.num_q_head, -1, self.head_dim]) |
| 172 | .transpose([0, 2, 1, 3]) |
| 173 | .reshape([-1, self.num_q_head, self.head_dim]) |
| 174 | ) |
| 175 | |
| 176 | def run_append_c16_attention( |
| 177 | self, q_len, kv_len, prefill=False, attn_mask=None, use_qknorm=False, mask_offset=None, qkv=None |
no test coverage detected