Spaces:
Running on Zero
Running on Zero
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
|