File size: 5,708 Bytes
a0e2620
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
"""Compatibility classes for trusted full-model checkpoints from the legacy app."""
from __future__ import annotations

import sys
import types

import torch
import torch.nn as nn
from transformers import PreTrainedModel, PretrainedConfig


class LegacyCustomTransformerConfig(PretrainedConfig):
    """Pickle-compatible replacement for ``model_config.CustomTransformerConfig``."""

    def __init__(self, vocab_size=128256, hidden_size=4096, num_layers=32, num_heads=32,
                 prediction_chunk=256, dropout=0, max_position_embeddings=4096,
                 masking_type="bidirectional", **kwargs):
        super().__init__(**kwargs)
        self.vocab_size = vocab_size
        self.hidden_size = hidden_size
        self.num_layers = num_layers
        self.num_heads = num_heads
        self.dropout = dropout
        self.prediction_chunk = prediction_chunk
        self.max_position_embeddings = max_position_embeddings
        self.input_size = prediction_chunk
        self.masking_type = masking_type


class LegacyCustomTransformerModel(PreTrainedModel):
    """Pickle-compatible legacy wrapper that supplies full bidirectional attention."""

    config_class = LegacyCustomTransformerConfig

    def forward(self, input_ids, labels=None, **kwargs):
        batch_size, seq_len = input_ids.shape
        masking_type = getattr(self.config, "masking_type", "bidirectional")
        if masking_type == "bidirectional":
            base_mask = torch.ones(seq_len, seq_len, dtype=torch.bool, device=input_ids.device)
        elif masking_type == "bidirectional_masked":
            base_mask = torch.ones(seq_len, seq_len, dtype=torch.bool, device=input_ids.device)
            base_mask.fill_diagonal_(False)
        elif masking_type == "unidirectional":
            base_mask = torch.tril(torch.ones(seq_len, seq_len, dtype=torch.bool, device=input_ids.device))
        else:
            raise ValueError(f"Unknown masking type: {masking_type}")
        llama = getattr(self.llama, "base_model", self.llama)
        compute_dtype = next(
            (parameter.dtype for parameter in llama.parameters() if parameter.is_floating_point()),
            torch.float32,
        )
        # The legacy checkpoint is commonly loaded in FP16 on Colab. SDPA
        # requires an additive attention bias to have the same dtype as the
        # query, so avoid the old unconditional float32 mask here.
        attention_mask = base_mask.unsqueeze(0).unsqueeze(1).expand(batch_size, 1, seq_len, seq_len).to(dtype=compute_dtype)
        # The hosted full checkpoint was serialized with peft==0.15.1.  Newer
        # PEFT's outer PeftModel.forward expects attributes absent from that
        # old pickled object.  Its base_model is the already-injected LoraModel
        # (and therefore retains the trained LoRA layers), so call it directly
        # when present rather than relying on version-sensitive PEFT hooks.
        outputs = llama(input_ids, attention_mask=attention_mask, output_hidden_states=True, use_cache=False, **kwargs)
        logits = outputs.logits[:, :, :self.config.vocab_size].view(batch_size, seq_len, self.config.vocab_size)
        if labels is None:
            return {"logits": logits}
        loss = nn.CrossEntropyLoss()(logits.view(-1, self.config.vocab_size), labels.view(-1))
        return {"loss": loss, "logits": logits}


_MISSING = object()


def install_legacy_pickle_modules() -> dict[str, object]:
    """Temporarily register the historical class locations expected by torch.load."""
    previous: dict[str, object] = {name: sys.modules.get(name) for name in ("model_config", "models")}
    config_module = types.ModuleType("model_config")
    config_module.CustomTransformerConfig = LegacyCustomTransformerConfig
    model_module = types.ModuleType("models")
    model_module.CustomTransformerModel = LegacyCustomTransformerModel
    sys.modules["model_config"] = config_module
    sys.modules["models"] = model_module
    # Some notebook-created full checkpoints pickle these classes under
    # ``__main__`` rather than their original source modules.
    main_module = sys.modules["__main__"]
    for name, value in {
        "CustomTransformerConfig": LegacyCustomTransformerConfig,
        "CustomTransformerModel": LegacyCustomTransformerModel,
    }.items():
        previous[f"__main__.{name}"] = getattr(main_module, name, _MISSING)
        setattr(main_module, name, value)
    return previous


def restore_legacy_pickle_modules(previous: dict[str, object]) -> None:
    """Restore module registrations changed for one trusted checkpoint load."""
    for name in ("model_config", "models"):
        module = previous[name]
        if module is None:
            sys.modules.pop(name, None)
        else:
            sys.modules[name] = module  # type: ignore[assignment]
    main_module = sys.modules["__main__"]
    for name in ("CustomTransformerConfig", "CustomTransformerModel"):
        previous_value = previous[f"__main__.{name}"]
        if previous_value is _MISSING:
            delattr(main_module, name)
        else:
            setattr(main_module, name, previous_value)


def patch_legacy_lora_modules(model: nn.Module) -> int:
    """Add fields expected by newer PEFT LoRA forwards to an old pickle.

    The hosted checkpoint predates PEFT's adapter-variant mechanism.  Its
    injected LoRA linears remain ordinary LoRA modules; an empty mapping makes
    current PEFT take that unchanged vanilla-LoRA branch.
    """
    patched = 0
    for module in model.modules():
        if hasattr(module, "lora_A") and not hasattr(module, "lora_variant"):
            module.lora_variant = {}
            patched += 1
    return patched