MCPcopy Create free account
hub / github.com/AMAP-ML/EMF / create_reference_model

Function create_reference_model

trl/trl/models/modeling_base.py:592–664  ·  view source on GitHub ↗

Creates a static reference copy of a model. Note that model will be in `.eval()` mode. Args: model (`PreTrainedModelWrapper`): The model to be copied. num_shared_layers (`int`, *optional*): The number of initial layers that are shared between both models and kep

(
    model: PreTrainedModelWrapper, num_shared_layers: Optional[int] = None, pattern: Optional[str] = None
)

Source from the content-addressed store, hash-verified

590
591
592def create_reference_model(
593 model: PreTrainedModelWrapper, num_shared_layers: Optional[int] = None, pattern: Optional[str] = None
594) -> PreTrainedModelWrapper:
595 """
596 Creates a static reference copy of a model. Note that model will be in `.eval()` mode.
597
598 Args:
599 model (`PreTrainedModelWrapper`): The model to be copied.
600 num_shared_layers (`int`, *optional*):
601 The number of initial layers that are shared between both models and kept frozen.
602 pattern (`str`, *optional*): The shared layers are selected with a string pattern
603 (e.g. "transformer.h.{layer}" for GPT2) and if a custom pattern is necessary it can be passed here.
604
605 Returns:
606 `PreTrainedModelWrapper`
607 """
608 if is_deepspeed_zero3_enabled():
609 raise ValueError(
610 "DeepSpeed ZeRO-3 is enabled and is not compatible with `create_reference_model()`. Please instantiate your reference model directly with `AutoModelForCausalLM.from_pretrained()`."
611 )
612
613 parameter_names = [n for n, _ in model.named_parameters()]
614 ref_model = deepcopy(model)
615
616 # if no layers are shared, return copy of model
617 if num_shared_layers is None:
618 for param_name in parameter_names:
619 param = ref_model.get_parameter(param_name)
620 param.requires_grad = False
621 return ref_model.eval()
622
623 # identify layer name pattern
624 if pattern is not None:
625 pattern = pattern.format(layer=num_shared_layers)
626 else:
627 for pattern_candidate in LAYER_PATTERNS:
628 pattern_candidate = pattern_candidate.format(layer=num_shared_layers)
629 if any(pattern_candidate in name for name in parameter_names):
630 pattern = pattern_candidate
631 break
632
633 if pattern is None:
634 raise ValueError("Layer pattern could not be matched.")
635
636 # divide parameters in shared and unshared parameter lists
637 shared_param_list = []
638 unshared_param_list = []
639
640 shared_parameter = True
641 for name, _param in model.named_parameters():
642 if pattern in name:
643 shared_parameter = False
644 if shared_parameter:
645 shared_param_list.append(name)
646 else:
647 unshared_param_list.append(name)
648
649 # create reference of the original parameter if they are shared

Callers 6

__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected