(model)
| 173 | self.norm_type, self.scale_grad_by_freq, self.sparse) |
| 174 | |
| 175 | def replace_layer(model): |
| 176 | for name, module in model.named_children(): |
| 177 | if isinstance(module,quantized_layer.QRMSNorm): |
| 178 | continue |
| 179 | if isinstance(module, torch.nn.Linear): |
| 180 | with init_weights_on_device(): |
| 181 | new_layer = quantized_layer.QLinear(module.in_features,module.out_features) |
| 182 | new_layer.weight = module.weight |
| 183 | if module.bias is not None: |
| 184 | new_layer.bias = module.bias |
| 185 | setattr(model, name, new_layer) |
| 186 | elif isinstance(module, RMSNorm): |
| 187 | if hasattr(module,"quantized"): |
| 188 | continue |
| 189 | module.quantized= True |
| 190 | new_layer = quantized_layer.QRMSNorm(module) |
| 191 | setattr(model, name, new_layer) |
| 192 | elif isinstance(module,torch.nn.Embedding): |
| 193 | rows, cols = module.weight.shape |
| 194 | new_layer = quantized_layer.QEmbedding( |
| 195 | num_embeddings=rows, |
| 196 | embedding_dim=cols, |
| 197 | _weight=module.weight, |
| 198 | # _freeze=module.freeze, |
| 199 | padding_idx=module.padding_idx, |
| 200 | max_norm=module.max_norm, |
| 201 | norm_type=module.norm_type, |
| 202 | scale_grad_by_freq=module.scale_grad_by_freq, |
| 203 | sparse=module.sparse) |
| 204 | setattr(model, name, new_layer) |
| 205 | else: |
| 206 | replace_layer(module) |
| 207 | |
| 208 | replace_layer(self) |
| 209 |
nothing calls this directly
no test coverage detected