Spaces:
Running on Zero
Running on Zero
File size: 10,695 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 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 | """Causal-LM loading adapted for bidirectional denoising."""
from __future__ import annotations
from collections import Counter, OrderedDict
from typing import Any
import torch
import os
from pathlib import Path
from peft import LoraConfig, PeftModel, TaskType, get_peft_model, prepare_model_for_kbit_training
from transformers import AutoModelForCausalLM
_ATTENTION_MASK_CACHE: OrderedDict[tuple[Any, ...], torch.Tensor] = OrderedDict()
_ATTENTION_MASK_CACHE_SIZE = 4
def bidirectional_attention_mask(padding_mask: torch.Tensor, dtype: torch.dtype) -> torch.Tensor:
"""A 4-D additive full-attention mask accepted unchanged by Transformers >=5.
`padding_mask` is retained for loss/padding bookkeeping, but padding is
intentionally visible to attention. Repeated EOS padding is part of the
configured context-width signal: real and padded queries can attend to all
positions, allowing the model to learn concise answers under wide contexts.
"""
# Every position is deliberately visible, so the additive mask is exactly
# zero. Reuse the most recent shapes instead of allocating and clearing the
# same dense tensor on every forward pass.
shape = (padding_mask.shape[0], 1, padding_mask.shape[1], padding_mask.shape[1])
# Inference tensors cannot later be saved by autograd, so training and
# inference-mode allocations must occupy separate cache entries.
key = (
padding_mask.device.type,
padding_mask.device.index,
dtype,
torch.is_inference_mode_enabled(),
*shape,
)
mask = _ATTENTION_MASK_CACHE.get(key)
if mask is None:
mask = torch.zeros(shape, device=padding_mask.device, dtype=dtype)
_ATTENTION_MASK_CACHE[key] = mask
if len(_ATTENTION_MASK_CACHE) > _ATTENTION_MASK_CACHE_SIZE:
_ATTENTION_MASK_CACHE.popitem(last=False)
else:
_ATTENTION_MASK_CACHE.move_to_end(key)
return mask
def _target_modules(model: torch.nn.Module, requested: list[str]) -> list[str]:
"""Verify every requested LoRA projection exists in the loaded architecture."""
available = {name.rsplit(".", 1)[-1] for name, _ in model.named_modules()}
missing = [name for name in requested if name not in available]
if missing:
raise ValueError(f"LoRA target modules missing from {model.config.model_type}: {missing}; available suffixes include {sorted(available)[:40]}")
return requested
def parameter_audit(model: torch.nn.Module) -> dict[str, Any]:
"""Assert the intended trainable set and return parameter-count diagnostics."""
named = list(model.named_parameters())
trainable = [(name, p) for name, p in named if p.requires_grad]
total = sum(p.numel() for _, p in named)
categories = Counter()
unexpected = []
for name, p in trainable:
if "lora_" in name:
categories["lora"] += p.numel()
elif "norm" in name.lower():
categories["norm"] += p.numel()
else:
categories["other"] += p.numel()
unexpected.append(name)
frozen_embedding = all(not p.requires_grad for n, p in named if any(x in n.lower() for x in ("embed_tokens", "embed_tokens", "wte")))
frozen_lm_head = all(not p.requires_grad for n, p in named if "lm_head" in n)
if not frozen_embedding or not frozen_lm_head or unexpected:
raise AssertionError({"embeddings_frozen": frozen_embedding, "lm_head_frozen": frozen_lm_head, "unexpected_trainable": unexpected})
trainable_count = sum(p.numel() for _, p in trainable)
return {"lora_parameters": categories["lora"], "normalization_parameters": categories["norm"], "other_trainable_parameters": categories["other"], "total_trainable_parameters": trainable_count, "total_model_parameters": total, "trainable_percentage": 100 * trainable_count / total, "trainable_names": [n for n, _ in trainable]}
def load_denoising_model(config: dict[str, Any]) -> tuple[torch.nn.Module, dict[str, Any]]:
"""Load a base CausalLM, attach LoRA, unfreeze norms, and audit it."""
checkpoint = config["model_name_or_path"]
precision = config.get("precision", "bf16")
dtype = {"fp16": torch.float16, "bf16": torch.bfloat16, "fp32": torch.float32}.get(precision)
if dtype is None:
raise ValueError(f"Unknown precision {precision}")
# No device_map: accelerate owns device placement in distributed runs.
model_cache = Path(config.get("base_model_cache_dir", "base_models")); model_cache.mkdir(parents=True, exist_ok=True)
quantization = str(config.get("quantization", "none")).lower()
load_kwargs = dict(dtype=dtype, trust_remote_code=False, token=os.getenv("HF_TOKEN"), cache_dir=str(model_cache))
if quantization in {"4bit", "4-bit", "qlora"}:
try:
from transformers import BitsAndBytesConfig
import bitsandbytes # noqa: F401
except ImportError as exc:
raise ImportError("quantization=4bit requires CUDA bitsandbytes; install with `pip install -e '.[cuda]'`") from exc
if not torch.cuda.is_available():
raise RuntimeError("4-bit bitsandbytes quantization requires an NVIDIA CUDA device")
compute_dtype = {"fp16": torch.float16, "bf16": torch.bfloat16, "fp32": torch.float32}.get(str(config.get("compute_dtype", precision)), dtype)
load_kwargs["quantization_config"] = BitsAndBytesConfig(load_in_4bit=True, bnb_4bit_quant_type=str(config.get("quantization_type", "nf4")), bnb_4bit_compute_dtype=compute_dtype, bnb_4bit_use_double_quant=bool(config.get("double_quant", True)))
elif quantization not in {"none", "off", "false"}:
raise ValueError("quantization must be 'none' or '4bit'")
model = AutoModelForCausalLM.from_pretrained(checkpoint, **load_kwargs)
model.config.use_cache = False
# PEFT's CAUSAL_LM task type only describes adapter integration; it does
# not control attention direction. Make the base-model intent explicit as
# well as supplying the prepared 4-D mask in forward_bidirectional().
model.config.is_causal = False
if hasattr(model.config, "use_bidirectional_attention"):
model.config.use_bidirectional_attention = True
if quantization in {"4bit", "4-bit", "qlora"}:
model = prepare_model_for_kbit_training(model, use_gradient_checkpointing=bool(config.get("gradient_checkpointing", False)))
for parameter in model.parameters():
parameter.requires_grad = False
targets = _target_modules(model, list(config.get("lora_targets", ["q_proj", "v_proj", "o_proj"])))
lora_config = LoraConfig(r=int(config.get("lora_r", 16)), lora_alpha=int(config.get("lora_alpha", 32)), lora_dropout=float(config.get("lora_dropout", 0.05)), target_modules=targets, bias="none", task_type=TaskType.CAUSAL_LM)
resume_adapter = config.get("resume_from_adapter")
if resume_adapter:
adapter_path = Path(resume_adapter)
if not (adapter_path / "adapter_config.json").is_file():
raise ValueError(f"resume_from_adapter is not a saved adapter directory: {adapter_path}")
model = PeftModel.from_pretrained(model, adapter_path, is_trainable=True)
norm_state_path = adapter_path / "normalization_state.pt"
if norm_state_path.is_file():
norm_state = torch.load(norm_state_path, map_location="cpu", weights_only=True)
named = dict(model.named_parameters())
for name, value in norm_state.items():
if name in named:
named[name].data.copy_(value.to(named[name].device, dtype=named[name].dtype))
else:
model = get_peft_model(model, lora_config)
if config.get("train_normalization_layers", True):
for name, parameter in model.named_parameters():
if "norm" in name.lower():
parameter.requires_grad = True
# Accelerate's FP16 GradScaler cannot unscale FP16 gradients.
# Keep trainable normalization parameters in FP32 so their
# gradients are scaler-compatible; frozen base weights remain
# in the configured compute dtype.
if precision == "fp16" and parameter.dtype == torch.float16:
parameter.data = parameter.data.float()
if config.get("gradient_checkpointing", False):
model.gradient_checkpointing_enable()
model.enable_input_require_grads()
audit = parameter_audit(model)
audit.update({"model_name": checkpoint, "resolved_lora_targets": targets})
return model, audit
def forward_bidirectional(model: torch.nn.Module, input_ids: torch.Tensor, padding_mask: torch.Tensor):
"""Run a CausalLM with the project’s explicit bidirectional padding mask."""
# 4-bit bitsandbytes weights are stored as uint8, which cannot represent
# the floating additive attention mask. Use the first floating parameter
# (normally a LoRA or normalization parameter) as the compute dtype.
dtype = getattr(model, "_lad_attention_mask_dtype", None)
if dtype is None:
dtype = next((parameter.dtype for parameter in model.parameters() if parameter.is_floating_point()), torch.float32)
model._lad_attention_mask_dtype = dtype
return model(input_ids=input_ids, attention_mask=bidirectional_attention_mask(padding_mask, dtype), use_cache=False).logits
def forward_bidirectional_selected(
model: torch.nn.Module,
input_ids: torch.Tensor,
padding_mask: torch.Tensor,
selection_mask: torch.Tensor,
):
"""Run the frozen LM head only at positions participating in the objective."""
dtype = getattr(model, "_lad_attention_mask_dtype", None)
if dtype is None:
dtype = next((parameter.dtype for parameter in model.parameters() if parameter.is_floating_point()), torch.float32)
model._lad_attention_mask_dtype = dtype
unwrapped = model.module if hasattr(model, "module") else model
causal_lm = unwrapped.get_base_model() if hasattr(unwrapped, "get_base_model") else unwrapped
backbone = getattr(causal_lm, "model", None)
output_head = causal_lm.get_output_embeddings() if hasattr(causal_lm, "get_output_embeddings") else None
if backbone is None or output_head is None:
raise TypeError(f"Selected-logit optimization is unsupported for {type(causal_lm).__name__}")
outputs = backbone(
input_ids=input_ids,
attention_mask=bidirectional_attention_mask(padding_mask, dtype),
use_cache=False,
)
example_ids, token_ids = selection_mask.nonzero(as_tuple=True)
selected_logits = output_head(outputs.last_hidden_state[example_ids, token_ids])
return selected_logits, example_ids, token_ids
|