| 1085 | |
| 1086 | |
| 1087 | class PlasAttentionConfig: |
| 1088 | def __init__( |
| 1089 | self, |
| 1090 | args, |
| 1091 | ): |
| 1092 | self.plas_encoder_top_k_left: int = None |
| 1093 | self.plas_encoder_top_k_right: int = None |
| 1094 | "The sparse topk of encoder attention is located at [plas_encoder_top_k_left, plas_encoder top_k_right]" |
| 1095 | self.plas_decoder_top_k_left: int = None |
| 1096 | self.plas_decoder_top_k_right: int = None |
| 1097 | "The sparse topk of decoder attention is located at [plas_decoder_top_k_left, plas_decoder top_k_right]" |
| 1098 | self.plas_use_encoder_seq_limit: int = None |
| 1099 | "When the number of encdoer token is less than plas_use_encoder_seq_limit, it is not sparse" |
| 1100 | self.plas_use_decoder_seq_limit: int = None |
| 1101 | "When the number of decdoer token is less than plas_use_decoder_seq_limit, it is not sparse" |
| 1102 | self.plas_block_size: int = 128 |
| 1103 | self.mlp_weight_name: str = "plas_attention_mlp_weight.safetensors" |
| 1104 | self.plas_max_seq_length: int = 128 * 1024 |
| 1105 | if args is not None: |
| 1106 | for key, value in args.items(): |
| 1107 | if hasattr(self, key): |
| 1108 | setattr(self, key, value) |
| 1109 | if self.plas_use_encoder_seq_limit is None and self.plas_encoder_top_k_left is not None: |
| 1110 | self.plas_use_encoder_seq_limit = self.plas_encoder_top_k_left * self.plas_block_size |
| 1111 | if self.plas_use_decoder_seq_limit is None and self.plas_decoder_top_k_left is not None: |
| 1112 | self.plas_use_decoder_seq_limit = self.plas_decoder_top_k_left * self.plas_block_size |
| 1113 | self.check_legality_parameters() |
| 1114 | |
| 1115 | def check_legality_parameters( |
| 1116 | self, |
| 1117 | ) -> None: |
| 1118 | if self.plas_encoder_top_k_left is not None: |
| 1119 | assert self.plas_encoder_top_k_left > 0, "plas_encoder_top_k_left must large than 0" |
| 1120 | |
| 1121 | if self.plas_encoder_top_k_right is not None: |
| 1122 | assert self.plas_encoder_top_k_right > 0, "plas_encoder_top_k_right must large than 0" |
| 1123 | assert ( |
| 1124 | self.plas_encoder_top_k_right >= self.plas_encoder_top_k_left |
| 1125 | ), "plas_encoder_top_k_right must large than plas_encoder_top_k_left" |
| 1126 | |
| 1127 | if self.plas_decoder_top_k_left is not None: |
| 1128 | assert self.plas_decoder_top_k_left > 0, "plas_decoder_top_k_left must large than 0" |
| 1129 | |
| 1130 | if self.plas_decoder_top_k_right is not None: |
| 1131 | assert self.plas_decoder_top_k_right > 0, "plas_decoder_top_k_right must large than 0" |
| 1132 | assert ( |
| 1133 | self.plas_decoder_top_k_right >= self.plas_decoder_top_k_left |
| 1134 | ), "plas_decoder_top_k_right must large than plas_decoder_top_k_left" |
| 1135 | |
| 1136 | if self.plas_use_encoder_seq_limit is not None and self.plas_encoder_top_k_left is not None: |
| 1137 | assert self.plas_use_encoder_seq_limit >= self.plas_encoder_top_k_left * self.plas_block_size |
| 1138 | if self.plas_use_decoder_seq_limit is not None and self.plas_decoder_top_k_left is not None: |
| 1139 | assert self.plas_use_decoder_seq_limit >= self.plas_decoder_top_k_left * self.plas_block_size |
| 1140 | |
| 1141 | def to_json_string(self): |
| 1142 | """ |
| 1143 | Convert plas_attention_config to json string. |
| 1144 | """ |
no outgoing calls
no test coverage detected