Team Ai
Modelpublic

staticlabs/dlm-code0.6b-exp

sourceHugging Faceupdated 1mo agoView on Hugging Face
0likes126downloads
generation_utils.py249 linesDownload Raw Back to root
1# veomni/models/transformers/qwen2/generation_utils.py2 3import warnings4import copy5from dataclasses import dataclass6from typing import Any, Dict, Optional, Tuple, Union7 8import torch9import torch.distributions as dists10from torch.nn import functional as F11from transformers import __version__12from transformers.generation.configuration_utils import GenerationConfig13from transformers.utils import ModelOutput, is_torchdynamo_compiling, logging14 15logger = logging.get_logger(__name__)16 17 18def top_p_logits(logits, top_p=None):19    sorted_logits, sorted_indices = torch.sort(logits, descending=True)20    cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)21    sorted_indices_to_remove = cumulative_probs > top_p22    sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()23    sorted_indices_to_remove[..., 0] = 024    mask = torch.zeros_like(logits, dtype=torch.bool, device=logits.device)25    mask = mask.scatter_(-1, sorted_indices, sorted_indices_to_remove)26    logits = logits.masked_fill(mask, torch.finfo(logits.dtype).min)27    return logits28 29def top_k_logits(logits, top_k=None):30    if top_k is None or top_k == 0:31        return logits32    top_k = min(top_k, logits.size(-1))33    indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None]34    logits = logits.masked_fill(indices_to_remove, torch.finfo(logits.dtype).min)35    return logits36 37def sample_tokens(logits, temperature=0.0, top_p=None, top_k=None, margin_confidence=False, neg_entropy=False):38    if temperature > 0:39        logits = logits / temperature40    if top_p is not None and top_p < 1:41        logits = top_p_logits(logits, top_p)42    if top_k is not None:43        logits = top_k_logits(logits, top_k)44    probs = torch.softmax(logits.float(), dim=-1)45    if temperature > 0:46        x0 = dists.Categorical(probs=probs).sample()47    else:48        _, x0 = probs.max(dim=-1)49    50    confidence = torch.gather(probs, -1, x0.unsqueeze(-1)).squeeze(-1)51 52    if margin_confidence:53        sorted_probs, _ = torch.sort(probs, dim=-1, descending=True)54        top1_probs = sorted_probs[..., 0]55        top2_probs = sorted_probs[..., 1]56        confidence = top1_probs - top2_probs57    elif neg_entropy:58        log_probs = torch.log(probs.clamp(min=1e-10))59        confidence = (probs * log_probs).sum(dim=-1)60    61    return confidence, x062 63 64@dataclass65class MDMModelOutput(ModelOutput):66    sequences: torch.LongTensor = None67    history: Optional[Tuple[torch.FloatTensor]] = None68 69class MDMGenerationConfig(GenerationConfig):70    def __init__(self, **kwargs):71        super().__init__(**kwargs)72        self.temperature: float = kwargs.pop("temperature", 0.0)73        self.top_p: Optional[float] = kwargs.pop("top_p", None)74        self.top_k: Optional[int] = kwargs.pop("top_k", None)75        self.eps: float = kwargs.pop("eps", 1e-3)76        self.steps: int = kwargs.pop("steps", 512)77        self.alg: str = kwargs.pop("alg", 'entropy')78        self.alg_temp: Optional[float] = kwargs.pop("alg_temp", 0.0)79        self.output_history: bool = kwargs.pop("output_history", False)80        self.mask_token_id = kwargs.pop("mask_token_id", None)81 82 83class MDMGenerationMixin:84    """85    Mixin class for Masked Diffusion Model generation, adapted from the Dream model's generation utils.86    """87    @staticmethod88    def _expand_inputs_for_generation(89        expand_size: int = 1,90        input_ids: Optional[torch.LongTensor] = None,91        attention_mask: Optional[torch.LongTensor] = None92    ) -> Tuple[torch.LongTensor, Dict[str, Any]]:93        if expand_size == 1:94            return input_ids, attention_mask95        96        if input_ids is not None:97            input_ids = input_ids.repeat_interleave(expand_size, dim=0)98        if attention_mask is not None:99            attention_mask = attention_mask.repeat_interleave(expand_size, dim=0)100        return input_ids, attention_mask101 102    def _prepare_generation_config(103        self, generation_config: Optional[GenerationConfig], **kwargs104    ) -> MDMGenerationConfig:105        if generation_config is None:106            generation_config = self.generation_config107        108        # Use MDMGenerationConfig as the target class109        if not isinstance(generation_config, MDMGenerationConfig):110            generation_config = MDMGenerationConfig.from_dict(generation_config.to_dict())111 112        # Update with kwargs113        generation_config.update(**kwargs)114        return generation_config115 116    @torch.no_grad()117    def diffusion_generate(118        self,119        inputs: Optional[torch.Tensor] = None,120        generation_config: Optional[MDMGenerationConfig] = None,121        **kwargs,122    ) -> Union[MDMModelOutput, torch.LongTensor]:123        124        # 1. Prepare generation config125        generation_config = self._prepare_generation_config(generation_config, **kwargs)126 127        # 2. Prepare inputs128        input_ids = inputs129        attention_mask = kwargs.get("attention_mask", None)130 131        if input_ids is None:132            raise ValueError("`inputs` must be provided for diffusion generation.")133 134        if generation_config.max_new_tokens is not None:135            generation_config.max_length = input_ids.shape[-1] + generation_config.max_new_tokens136        137        # 3. Expand inputs for multi-sequence generation138        input_ids, attention_mask = self._expand_inputs_for_generation(139            expand_size=generation_config.num_return_sequences,140            input_ids=input_ids,141            attention_mask=attention_mask142        )143        # 4. Run the sampling loop144        return self._sample(145            input_ids,146            attention_mask=attention_mask,147            generation_config=generation_config148        )149 150    def _sample(151        self,152        input_ids: torch.LongTensor,153        attention_mask: Optional[torch.LongTensor],154        generation_config: MDMGenerationConfig155    ) -> Union[MDMModelOutput, torch.LongTensor]:156        157        # Extract params from config158        max_length = generation_config.max_length159        mask_token_id = generation_config.mask_token_id160        if mask_token_id is None:161            raise ValueError("`mask_token_id` must be set in the generation config.")162 163        steps = generation_config.steps164        eps = generation_config.eps165        alg = generation_config.alg166        alg_temp = generation_config.alg_temp167        temperature = generation_config.temperature168        top_p = generation_config.top_p169        top_k = generation_config.top_k170 171        histories = [] if generation_config.output_history else None172 173        # Pad input_ids to max_length with mask tokens174        x = F.pad(input_ids, (0, max_length - input_ids.shape[1]), value=mask_token_id)175 176        # The model expects a bidirectional mask, so we just use the presence of pad_token_id177        # for the attention mask during generation.178        gen_attention_mask = (x != self.config.pad_token_id).long() if self.config.pad_token_id is not None else None179 180        timesteps = torch.linspace(1, eps, steps + 1, device=x.device)181 182        for i in range(steps):183            mask_index = (x == mask_token_id)184            if not mask_index.any(): # Stop if no tokens are masked185                break186            # is_causal=False is crucial for bidirectional attention187            outputs = self(input_ids=x, attention_mask=gen_attention_mask, is_causal=False)188            logits = outputs.logits189            190            # CRITICAL: Shift logits to predict the next token, aligning with training191            logits = torch.cat([logits[:, :1], logits[:, :-1]], dim=1)192 193            mask_logits = logits[mask_index]194            t = timesteps[i]195            s = timesteps[i + 1]196 197            if alg == 'origin':198                p_transfer = 1 - s / t if i < steps - 1 else 1199                x0 = torch.full_like(x[mask_index], fill_value=mask_token_id, device=self.device, dtype=torch.long)200                transfer_index_t_s = torch.rand(*x0.shape, device=self.device) < p_transfer201                _, sampled_tokens = sample_tokens(mask_logits[transfer_index_t_s], temperature=temperature, top_p=top_p, top_k=top_k)202                x0[transfer_index_t_s] = sampled_tokens203                x[mask_index] = x0204            else:205                # Confidence-based sampling (maskgit, entropy, etc.)206                confidence_alg_map = {'maskgit_plus': False, 'topk_margin': True, 'entropy': True}207                is_margin_conf = confidence_alg_map.get(alg, False)208                is_neg_entropy = alg == 'entropy'209                210                confidence, x0 = sample_tokens(mask_logits, temperature, top_p, top_k, margin_confidence=is_margin_conf, neg_entropy=is_neg_entropy)211 212                num_masked = mask_index.sum(dim=-1, keepdim=True)213                gamma = 1 - s / t214                num_to_unmask = (num_masked * gamma).long()215 216                # Place confidence scores back into a full tensor to find top-k across the sequence217                full_confidence = torch.full_like(x, -torch.inf, device=self.device, dtype=confidence.dtype)218                full_confidence[mask_index] = confidence219 220                if (alg_temp is not None and alg_temp > 0):221                    # Temperature-based sampling of which tokens to unmask222                    unmask_probs = F.softmax(full_confidence / alg_temp, dim=-1)223                    unmask_indices = torch.multinomial(unmask_probs, num_samples=num_to_unmask.max(), replacement=False)224                else:225                    # Top-k confidence sampling226                    _, unmask_indices = torch.topk(full_confidence, k=num_to_unmask.max(), dim=-1)227 228                # Create a mask for the tokens we are going to unmask229                rows = torch.arange(x.size(0), device=x.device).unsqueeze(1)230                unmask_selection_mask = torch.zeros_like(x, dtype=torch.bool)231                unmask_selection_mask[rows, unmask_indices] = True232                233                # Filter indices based on per-row `num_to_unmask`234                unmask_selection_mask = unmask_selection_mask & (torch.cumsum(unmask_selection_mask.long(), dim=-1) <= num_to_unmask)235 236                # Place the newly generated tokens (x0) into a full tensor237                x_unmasked_proposals = torch.full_like(x, fill_value=mask_token_id)238                x_unmasked_proposals[mask_index] = x0239 240                # Update the main tensor `x` with the unmasked tokens241                x[unmask_selection_mask] = x_unmasked_proposals[unmask_selection_mask]242 243            if histories is not None:244                histories.append(x.clone())245 246        if generation_config.return_dict_in_generate:247            return MDMModelOutput(sequences=x, history=histories)248        else:249            return x