(module, prefix='', level=0)
| 220 | all_ds_ids = {} |
| 221 | |
| 222 | def load_module_recursive(module, prefix='', level=0): |
| 223 | for name, child in module.named_children(): |
| 224 | if child.__class__ in layer_policies: |
| 225 | checking_key = prefix + name + '.' |
| 226 | if not any(checking_key in item for item in sd[0].keys()): |
| 227 | if hasattr(child, 'weight') and \ |
| 228 | (hasattr(child.weight, 'ds_id') and \ |
| 229 | child.weight.ds_id in all_ds_ids): |
| 230 | prefix1 = all_ds_ids[child.weight.ds_id] |
| 231 | if child.__class__ is nn.Linear: |
| 232 | child = LinearLayer.from_weights(weight=all_ds_ids[child.weight.ds_id]) |
| 233 | setattr(module, name, child) |
| 234 | continue |
| 235 | child_params = list(child.parameters()) |
| 236 | if len(child_params) > 0 and (child_params[0].numel() == 0 or child_params[0].is_meta): |
| 237 | if child.weight.is_meta: |
| 238 | ds_shape = child.weight.shape |
| 239 | else: |
| 240 | ds_shape = child.weight.ds_shape |
| 241 | if child.__class__ is nn.LayerNorm: |
| 242 | child = Normalize(dim=ds_shape[-1], dtype=child.weight.dtype, eps=child.eps) |
| 243 | setattr(module, name, child) |
| 244 | elif child.__class__ in [nn.Linear, ColumnParallelLinear, RowParallelLinear]: |
| 245 | child = LinearLayer.from_weights(weight_shape=child.weight.shape, |
| 246 | dtype=child.weight.dtype, |
| 247 | bias=child.bias) |
| 248 | setattr(module, name, child) |
| 249 | elif child.__class__ is OPTLearnedPositionalEmbedding: |
| 250 | child = OPTEmbedding(weight_shape=ds_shape) |
| 251 | setattr(module, name, child) |
| 252 | elif child.__class__ in [LlamaRMSNorm, RMSNorm]: |
| 253 | child = RMSNormalize(dim=ds_shape[-1], |
| 254 | dtype=child.weight.dtype, |
| 255 | eps=child.eps if hasattr(child, 'eps') else child.variance_epsilon) |
| 256 | setattr(module, name, child) |
| 257 | else: |
| 258 | ds_id = None |
| 259 | if hasattr(child.weight, 'ds_id'): |
| 260 | ds_id = child.weight.ds_id |
| 261 | child = EmbeddingLayer(weight_shape=ds_shape, dtype=child.weight.dtype) |
| 262 | if ds_id is not None: |
| 263 | all_ds_ids[ds_id] = child.weight |
| 264 | setattr(module, name, child) |
| 265 | layer_policies[child.__class__](child, prefix + name + '.') |
| 266 | else: |
| 267 | load_module_recursive( |
| 268 | child, |
| 269 | prefix if (level == 0 and ckpt_type == 'pp') and skip_level_0_prefix else \ |
| 270 | prefix + name + '.', |
| 271 | level + 1) |
| 272 | |
| 273 | load_module_recursive(r_module) |
| 274 |
no test coverage detected