(
self,
query_layer,
key_layer,
value_layer,
attention_mask,
attention_dropout=None,
log_attention_weights=None,
scaling_attention_score=True,
**kwargs,
)
| 352 | return None |
| 353 | |
| 354 | def attention_fn( |
| 355 | self, |
| 356 | query_layer, |
| 357 | key_layer, |
| 358 | value_layer, |
| 359 | attention_mask, |
| 360 | attention_dropout=None, |
| 361 | log_attention_weights=None, |
| 362 | scaling_attention_score=True, |
| 363 | **kwargs, |
| 364 | ): |
| 365 | attention_fn_default = HOOKS_DEFAULT["attention_fn"] |
| 366 | |
| 367 | if self.pnp: |
| 368 | query_layer = self.rotary(query_layer, **kwargs) |
| 369 | key_layer = self.rotary(key_layer, **kwargs) |
| 370 | if self.rot_v: |
| 371 | value_layer = self.rotary(value_layer) |
| 372 | else: |
| 373 | query_layer = torch.cat( |
| 374 | ( |
| 375 | query_layer[ |
| 376 | :, |
| 377 | :, |
| 378 | : kwargs["text_length"], |
| 379 | ], |
| 380 | self.rotary( |
| 381 | query_layer[ |
| 382 | :, |
| 383 | :, |
| 384 | kwargs["text_length"] :, |
| 385 | ] |
| 386 | ), |
| 387 | ), |
| 388 | dim=2, |
| 389 | ) |
| 390 | key_layer = torch.cat( |
| 391 | ( |
| 392 | key_layer[ |
| 393 | :, |
| 394 | :, |
| 395 | : kwargs["text_length"], |
| 396 | ], |
| 397 | self.rotary( |
| 398 | key_layer[ |
| 399 | :, |
| 400 | :, |
| 401 | kwargs["text_length"] :, |
| 402 | ] |
| 403 | ), |
| 404 | ), |
| 405 | dim=2, |
| 406 | ) |
| 407 | if self.rot_v: |
| 408 | value_layer = torch.cat( |
| 409 | ( |
| 410 | value_layer[ |
| 411 | :, |
nothing calls this directly
no test coverage detected