# coding=utf-8 # Copyright 2024 The Dream team, HKUNLP Group and the HuggingFace Inc. team. All rights reserved. # # This code is based on EleutherAI's GPT-NeoX library and the GPT-NeoX # and OPT and Qwen implementations in this library. It has been modified from its # original forms to accommodate minor architectural differences compared # to GPT-NeoX and OPT and Qwen used by the Meta AI and Qwen team that trained the model. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. """PyTorch Dream model.""" from .modeling_sensevoice import AudioEncoder from .resampler_projector import ResamplerProjector import random from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel # import torch._dynamo # torch._dynamo.config.suppress_errors = True import sys import pdb class ForkedPdb(pdb.Pdb): """ PDB Subclass for debugging multi-processed code Suggested in: https://stackoverflow.com/questions/4716533/how-to-attach-debugger-to-a-python-subproccess """ def interaction(self, *args, **kwargs): _stdin = sys.stdin try: sys.stdin = open('/dev/stdin') pdb.Pdb.interaction(self, *args, **kwargs) finally: sys.stdin = _stdin import math from typing import List, Optional, Tuple, Union import os import torch import torch.utils.checkpoint from torch import nn from transformers.activations import ACT2FN from transformers.cache_utils import Cache, DynamicCache from transformers.modeling_outputs import ( # BaseModelOutput, MaskedLMOutput, ModelOutput ) from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS from transformers.modeling_utils import PreTrainedModel from transformers.utils import ( add_start_docstrings, add_start_docstrings_to_model_forward, is_flash_attn_2_available, is_flash_attn_greater_or_equal_2_10, logging, ) from transformers import PretrainedConfig from .configuration_dream import DreamConfig from .generation_utils import DreamGenerationMixin, DreamGenerationConfig from dataclasses import dataclass from typing import Any if is_flash_attn_2_available(): from transformers.modeling_flash_attention_utils import _flash_attention_forward from transformers.modeling_outputs import CausalLMOutputWithPast from .modeling_sensevoice import AudioEncoder from .resampler_projector import ResamplerProjector logger = logging.get_logger(__name__) # def forward_process(bsz, seq_len, device, first_non_neg_idx_list, last_non_neg_idx_list, eps=1e-3): # b, l = bsz, seq_len # b → batch_size,l → 序列长度 # # 初始化掩码输出 # masked_indices = torch.zeros((b, l), device=device, dtype=torch.bool) # p_mask = torch.rand(b, device=device) # shape: [b],生成一个随机数作为掩码比例 # # 映射到 (eps, 1) 区间,保证最小值不低于 eps # p_mask = (1 - eps) * p_mask + eps # shape: [b] # p_mask = p_mask[:, None] # shape: [b, 1],方便广播 # # 针对每个样本的有效部分生成掩码 # for i in range(b): # first_non_neg_idx = first_non_neg_idx_list[i] # last_non_neg_idx = last_non_neg_idx_list[i] # # 如果无有效区间,跳过 # if first_non_neg_idx is None or last_non_neg_idx is None: # continue # valid_length = last_non_neg_idx - first_non_neg_idx + 1 # if valid_length <= 0: # continue # # 生成当前样本的掩码概率阈值 # t = torch.rand(valid_length, device=device) # shape: [valid_length] # mask_threshold = (1 - eps) * t + eps # 计算该样本的掩码概率 # # 计算该样本掩码的上限 # mask_cutoff = torch.max(mask_threshold, torch.min(torch.rand(valid_length, device=device))) # shape: [valid_length] # # 在有效部分生成掩码 # masked_indices[i, first_non_neg_idx:last_non_neg_idx+1] = torch.rand(valid_length, device=device) <= mask_cutoff # return masked_indices, p_mask def forward_process( bsz: int, seq_len: int, device: torch.device, labels: torch.Tensor, # [b, l] 的标签 tensor eps: float = 1e-3, special_token_id: int = 151643, # 要“优待”的 token id special_mask_ratio: float = 0.1 # special token 只按原阈值的 10% 掩 ) -> Tuple[torch.Tensor, torch.Tensor]: """ 生成掩码,并打印统计: - 总共被掩的 token 数 - special_token_id 被掩的数量 参数: - bsz: batch size - seq_len: 序列长度 - device: torch 设备 - labels: [b, l] 的标签 tensor(-100 表示无效位置) - eps: 阈值下限 - special_token_id: 要降低掩码率的特殊 token id - special_mask_ratio: special token 的掩码率缩放因子 返回: - masked_indices: [b, l] 的 bool 掩码矩阵 - p_mask: [b, 1] 每个样本的整体掩码比例 """ b, l = bsz, seq_len # 初始化掩码矩阵 & 每个样本的整体掩码比例 masked_indices = torch.zeros((b, l), device=device, dtype=torch.bool) p_mask = torch.rand(b, device=device) p_mask = (1 - eps) * p_mask + eps p_mask = p_mask.unsqueeze(1) # [b,1] # 先为每条序列计算第一个和最后一个非 -100 的位置 first_idxs = [] last_idxs = [] for i in range(b): nonneg = (labels[i] != -100).nonzero(as_tuple=True)[0] if nonneg.numel() == 0: first_idxs.append(None) last_idxs.append(None) else: first_idxs.append(int(nonneg[0])) last_idxs.append(int(nonneg[-1])) # 针对每条序列的有效区间生成掩码 for i in range(b): start = first_idxs[i] end = last_idxs[i] if start is None or end is None or end < start: continue valid_len = end - start + 1 # 为每个位置生成基础阈值 t = torch.rand(valid_len, device=device) mask_threshold = (1 - eps) * t + eps # [valid_len] # 生成随机判定值 rand_vals = torch.rand(valid_len, device=device) # 普通 token 的掩码决定 normal_mask = rand_vals <= mask_threshold # special token 的掩码阈值更低 special_thresh = mask_threshold * special_mask_ratio special_mask = rand_vals <= special_thresh labels_slice = labels[i, start : end + 1] # 最终掩码:特殊 token 用 special_mask,其他用 normal_mask final_mask = torch.where( labels_slice == special_token_id, special_mask, normal_mask ) masked_indices[i, start : end + 1] = final_mask # 打印统计信息 total_masked = int(masked_indices.sum().item()) special_masked = int((masked_indices & (labels == special_token_id)).sum().item()) # print(f"Total masked tokens: {total_masked}") # print(f"Special token_id={special_token_id} masked count: {special_masked}") return masked_indices, p_mask # def forward_process(bsz, seq_len, device, eps=1e-3): # b, l = bsz, seq_len # b → batch_size,l → 序列长度 # # 1) 为 batch 中的每个样本生成一个 0~1 的随机数 t # t = torch.rand(b, device=device) # # 2) 把 t 映射到 (eps, 1) 区间,保证最小值不低于 eps # # p_mask 相当于给每个样本定一个「掩码概率阈值」 # p_mask = (1 - eps) * t + eps # shape: [b] # # 3) 扩展出维度 [b, 1],方便后续广播 # p_mask = p_mask[:, None] # shape: [b, 1] # # 4) 针对 batch 中的每个 token 再生成一次随机数 # masked_indices = torch.rand((b, l), device=device) # shape: [b, l] # # 5) 计算当前样本要用的“掩码上限”: # # - masked_indices.min(-1).values → 每个样本里随机矩阵的最小值(保证至少有一个 token 会被掩掉) # # - torch.max(p_mask, 该最小值) → 二者取大,得到最终 cutoff # mask_cutoff = torch.max(p_mask, # masked_indices.min(-1, keepdim=True).values) # shape: [b, 1] # # 6) 生成最终布尔掩码:随机值 ≤ cutoff 的 token 被置 True # masked_indices = masked_indices <= mask_cutoff # shape: [b, l],dtype=bool # # 7) (可选)把 True 位置替换成 [MASK] token(示例注释里用 126336 表示) # # noisy_batch = torch.where(masked_indices, 126336, input_ids) # # 返回: # # masked_indices → [b, l] 的布尔矩阵,告诉你哪些 token 需要被掩码 # # p_mask → [b, 1] 的阈值,记录每条样本的“目标掩码比例” # return masked_indices, p_mask def generate_attention_mask(labels): batch_size, seq_len = labels.shape attention_mask = torch.zeros(batch_size, seq_len, seq_len, device=labels.device) # 用于存储每个 batch 的 first_non_neg_idx 和 last_non_neg_idx first_non_neg_idx_list = [] last_non_neg_idx_list = [] for i in range(batch_size): label = labels[i] # assert label.dtype in [torch.int64, torch.int32], f"label dtype is {label.dtype}" # assert not torch.isnan(label.float()).any(), "label has NaN" # assert not torch.isinf(label.float()).any(), "label has inf" try: non_neg_idx = (label != -100).nonzero(as_tuple=True)[0] except Exception as e: label_cpu = label.detach().cpu() # 先搬到 CPU print("label (unique) =", label_cpu.unique(), "shape =", label_cpu.shape) print('label.device:', label.device) print('label.shape:', label.shape) # 先拷到 CPU 再打印 try: print('label (cpu):', label.cpu()) except Exception as e2: print('label.cpu() 也出错:', e2) print('Exception:', e) # continue # continue # 跳过这个样本 # assert label.dtype in [torch.int64, torch.int32], f"label dtype is {label.dtype}" # assert not torch.isnan(label).any(), "label has NaN" # assert not torch.isinf(label).any(), "label has inf" # try: # non_neg_idx = (label != -100).nonzero(as_tuple=True)[0] # except: # print('label is :',label) if non_neg_idx.numel() == 0: # 全是-100,无法分区,给默认值或raise first_non_neg_idx = None last_non_neg_idx = None # 你可以选择跳过或全0/全1 # attention_mask[i] = 0 # 或者1 else: first_non_neg_idx = non_neg_idx[0].item() last_non_neg_idx = non_neg_idx[-1].item() # 第一部分只能看到自己 attention_mask[i, :first_non_neg_idx, :first_non_neg_idx] = 1 # 第二部分能看到第一部分和自己 attention_mask[i, first_non_neg_idx:last_non_neg_idx + 1, :first_non_neg_idx] = 1 attention_mask[i, first_non_neg_idx:last_non_neg_idx + 1, first_non_neg_idx:last_non_neg_idx + 1] = 1 # 第三部分能看到所有部分 attention_mask[i, last_non_neg_idx + 1:, :] = 1 first_non_neg_idx_list.append(first_non_neg_idx) last_non_neg_idx_list.append(last_non_neg_idx) return attention_mask, first_non_neg_idx_list, last_non_neg_idx_list def update_labels(input_ids, labels, eos_id, max_n=20): batch_size, seq_len = input_ids.shape first_occurrence_indices = [] # 记录每个 batch 中 eos_id 首次出现的位置 for idx in range(batch_size): eos_positions = (input_ids[idx] == eos_id).nonzero(as_tuple=True)[0] if len(eos_positions) > 0: first_occurrence_indices.append(eos_positions[0].item()) else: first_occurrence_indices.append(-1) # 如果没有 eos_id,则记录为 -1 # 从 first_idx 开始,按顺序选择 n 个位置来更新 for i in range(batch_size): first_idx = first_occurrence_indices[i] if first_idx == -1: continue # 跳过没有 eos 的样本 # 确保不会超过序列长度 max_possible = seq_len - first_idx # 如果 max_possible==0,说明 eos 刚好在最后一个位置,也跳过 if max_possible <= 0: continue num_to_select = random.randint(1, min(max_n, max_possible)) selected_indices = torch.arange(first_idx, first_idx + num_to_select) # 将这些位置的 labels 更新为 eos_id labels[i, selected_indices] = eos_id return labels import torch import random def update_labels_and_inputs(input_ids, labels, eos_id, max_n=20, pad_token_id=0, pad_label_id=-100): batch_size, seq_len = input_ids.shape input_ids = input_ids.clone() labels = labels.clone() new_input_ids = [] new_labels = [] for idx in range(batch_size): eos_positions = (input_ids[idx] == eos_id).nonzero(as_tuple=True)[0] if len(eos_positions) > 0: first_idx = eos_positions[0].item() cur_input_ids = input_ids[idx] cur_labels = labels[idx] else: # 扩展 max_n 个 eos_id random_max_n = random.randint(1, max_n) eos_ids = torch.full((random_max_n,), eos_id, device=input_ids.device, dtype=input_ids.dtype) cur_input_ids = torch.cat([input_ids[idx], eos_ids]) # labels 扩展 max_n 个 pad_label_id pad_labels = torch.full((random_max_n,), eos_id, device=labels.device, dtype=labels.dtype) cur_labels = torch.cat([labels[idx], pad_labels]) # first_idx = len(cur_input_ids) - random_max_n new_input_ids.append(cur_input_ids) new_labels.append(cur_labels) # pad到同一长度 max_len = max(len(x) for x in new_input_ids) padded_input_ids = torch.stack([ torch.cat([x, torch.full((max_len - len(x),), pad_token_id, device=x.device, dtype=x.dtype)]) for x in new_input_ids ]) padded_labels = torch.stack([ torch.cat([x, torch.full((max_len - len(x),), pad_label_id, device=x.device, dtype=x.dtype)]) for x in new_labels ]) return padded_input_ids, padded_labels @dataclass class MaskedLMOutput(ModelOutput): """ Base class for masked language models outputs. Args: loss (`torch.FloatTensor` of shape `(1,)`, *optional*, returned when `labels` is provided): Masked language modeling (MLM) loss. logits (`torch.FloatTensor` of shape `(batch_size, sequence_length, config.vocab_size)`): Prediction scores of the language modeling head (scores for each vocabulary token before SoftMax). hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`): Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, + one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`. Hidden-states of the model at the output of each layer plus the optional initial embedding outputs. attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`): Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length, sequence_length)`. Attentions weights after the attention softmax, used to compute the weighted average in the self-attention heads. """ loss: Optional[torch.FloatTensor] = None # loss_tqa: Optional[torch.FloatTensor] = None # loss_sqa: Optional[torch.FloatTensor] = None # loss_asr: Optional[torch.FloatTensor] = None # loss_tts: Optional[torch.FloatTensor] = None # loss_vqa: Optional[torch.FloatTensor] = None # loss_svqa: Optional[torch.FloatTensor] = None # loss_t2i: Optional[torch.FloatTensor] = None # loss_s2i: Optional[torch.FloatTensor] = None logits: torch.FloatTensor = None hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None attentions: Optional[Tuple[torch.FloatTensor, ...]] = None past_key_values: Optional[Tuple[Tuple[torch.FloatTensor]]] = None _CHECKPOINT_FOR_DOC = "Dream-7B" _CONFIG_FOR_DOC = "DreamConfig" import os ENFORCE_NUM_ITEMIN_BATCH = os.environ.get("ENFORCE_NUM_ITEMIN_BATCH", False) @dataclass class BaseModelOutput(ModelOutput): """ Base class for model's outputs, with potential hidden states and attentions. Args: last_hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`): Sequence of hidden-states at the output of the last layer of the model. hidden_states (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed or when `config.output_hidden_states=True`): Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, + one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`. Hidden-states of the model at the output of each layer plus the optional initial embedding outputs. attentions (`tuple(torch.FloatTensor)`, *optional*, returned when `output_attentions=True` is passed or when `config.output_attentions=True`): Tuple of `torch.FloatTensor` (one for each layer) of shape `(batch_size, num_heads, sequence_length, sequence_length)`. Attentions weights after the attention softmax, used to compute the weighted average in the self-attention heads. """ last_hidden_state: torch.FloatTensor = None hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None attentions: Optional[Tuple[torch.FloatTensor, ...]] = None past_key_values: Optional[Cache] = None # Copied from transformers.models.llama.modeling_llama.LlamaRMSNorm with Llama->Dream class DreamRMSNorm(nn.Module): def __init__(self, hidden_size, eps=1e-6): """ DreamRMSNorm is equivalent to T5LayerNorm """ super().__init__() self.weight = nn.Parameter(torch.ones(hidden_size)) self.variance_epsilon = eps def forward(self, hidden_states): input_dtype = hidden_states.dtype hidden_states = hidden_states.to(torch.float32) variance = hidden_states.pow(2).mean(-1, keepdim=True) hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) return self.weight * hidden_states.to(input_dtype) def extra_repr(self): return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}" # Copied from transformers.models.llama.modeling_llama.LlamaRotaryEmbedding with Llama->Dream class DreamRotaryEmbedding(nn.Module): def __init__( self, dim=None, max_position_embeddings=2048, base=10000, device=None, scaling_factor=1.0, rope_type="default", config: Optional[DreamConfig] = None, ): super().__init__() # TODO (joao): remove the `if` below, only used for BC self.rope_kwargs = {} if config is None: logger.warning_once( "`DreamRotaryEmbedding` can now be fully parameterized by passing the model config through the " "`config` argument. All other arguments will be removed in v4.46" ) self.rope_kwargs = { "rope_type": rope_type, "factor": scaling_factor, "dim": dim, "base": base, "max_position_embeddings": max_position_embeddings, } self.rope_type = rope_type self.max_seq_len_cached = max_position_embeddings self.original_max_seq_len = max_position_embeddings else: # BC: "rope_type" was originally "type" if config.rope_scaling is not None: self.rope_type = config.rope_scaling.get("rope_type", config.rope_scaling.get("type")) else: self.rope_type = "default" self.max_seq_len_cached = config.max_position_embeddings self.original_max_seq_len = config.max_position_embeddings self.config = config self.rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type] inv_freq, self.attention_scaling = self.rope_init_fn(self.config, device, **self.rope_kwargs) self.register_buffer("inv_freq", inv_freq, persistent=False) self.original_inv_freq = self.inv_freq def reset_parameters(self): inv_freq, self.attention_scaling = self.rope_init_fn(self.config, self.inv_freq.device, **self.rope_kwargs) self.register_buffer("inv_freq", inv_freq, persistent=False) self.original_inv_freq = self.inv_freq def _dynamic_frequency_update(self, position_ids, device): """ dynamic RoPE layers should recompute `inv_freq` in the following situations: 1 - growing beyond the cached sequence length (allow scaling) 2 - the current sequence length is in the original scale (avoid losing precision with small sequences) """ seq_len = torch.max(position_ids) + 1 if seq_len > self.max_seq_len_cached: # growth inv_freq, self.attention_scaling = self.rope_init_fn( self.config, device, seq_len=seq_len, **self.rope_kwargs ) self.register_buffer("inv_freq", inv_freq, persistent=False) # TODO joao: may break with compilation self.max_seq_len_cached = seq_len if seq_len < self.original_max_seq_len and self.max_seq_len_cached > self.original_max_seq_len: # reset self.register_buffer("inv_freq", self.original_inv_freq, persistent=False) self.max_seq_len_cached = self.original_max_seq_len @torch.no_grad() def forward(self, x, position_ids): if "dynamic" in self.rope_type: self._dynamic_frequency_update(position_ids, device=x.device) # Core RoPE block inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1) position_ids_expanded = position_ids[:, None, :].float() # Force float32 (see https://github.com/huggingface/transformers/pull/29285) device_type = x.device.type device_type = device_type if isinstance(device_type, str) and device_type != "mps" else "cpu" with torch.autocast(device_type=device_type, enabled=False): freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2) emb = torch.cat((freqs, freqs), dim=-1) cos = emb.cos() sin = emb.sin() # Advanced RoPE types (e.g. yarn) apply a post-processing scaling factor, equivalent to scaling attention cos = cos * self.attention_scaling sin = sin * self.attention_scaling return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype) # Copied from transformers.models.llama.modeling_llama.rotate_half def rotate_half(x): """Rotates half the hidden dims of the input.""" x1 = x[..., : x.shape[-1] // 2] x2 = x[..., x.shape[-1] // 2 :] return torch.cat((-x2, x1), dim=-1) # Copied from transformers.models.llama.modeling_llama.apply_rotary_pos_emb def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1): """Applies Rotary Position Embedding to the query and key tensors. Args: q (`torch.Tensor`): The query tensor. k (`torch.Tensor`): The key tensor. cos (`torch.Tensor`): The cosine part of the rotary embedding. sin (`torch.Tensor`): The sine part of the rotary embedding. position_ids (`torch.Tensor`, *optional*): Deprecated and unused. unsqueeze_dim (`int`, *optional*, defaults to 1): The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2. Returns: `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding. """ cos = cos.unsqueeze(unsqueeze_dim) sin = sin.unsqueeze(unsqueeze_dim) q_embed = (q * cos) + (rotate_half(q) * sin) k_embed = (k * cos) + (rotate_half(k) * sin) return q_embed, k_embed # Copied from transformers.models.mistral.modeling_mistral.MistralMLP with Mistral->Dream class DreamMLP(nn.Module): def __init__(self, config): super().__init__() self.hidden_size = config.hidden_size self.intermediate_size = config.intermediate_size self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) self.act_fn = ACT2FN[config.hidden_act] def forward(self, hidden_state): return self.down_proj(self.act_fn(self.gate_proj(hidden_state)) * self.up_proj(hidden_state)) # Copied from transformers.models.llama.modeling_llama.repeat_kv def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor: """ This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch, num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim) """ batch, num_key_value_heads, slen, head_dim = hidden_states.shape if n_rep == 1: return hidden_states hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim) return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim) class DreamAttention(nn.Module): """ Multi-headed attention from 'Attention Is All You Need' paper. Modified to use sliding window attention: Longformer and "Generating Long Sequences with Sparse Transformers". """ def __init__(self, config: DreamConfig, layer_idx: Optional[int] = None): super().__init__() self.config = config self.layer_idx = layer_idx if layer_idx is None: logger.warning_once( f"Instantiating {self.__class__.__name__} without passing `layer_idx` is not recommended and will " "to errors during the forward call, if caching is used. Please make sure to provide a `layer_idx` " "when creating this class." ) self.hidden_size = config.hidden_size self.num_heads = config.num_attention_heads self.head_dim = self.hidden_size // self.num_heads self.num_key_value_heads = config.num_key_value_heads self.num_key_value_groups = self.num_heads // self.num_key_value_heads self.max_position_embeddings = config.max_position_embeddings self.rope_theta = config.rope_theta self.is_causal = False self.attention_dropout = config.attention_dropout if (self.head_dim * self.num_heads) != self.hidden_size: raise ValueError( f"hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size}" f" and `num_heads`: {self.num_heads})." ) self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.head_dim, bias=True) self.k_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=True) self.v_proj = nn.Linear(self.hidden_size, self.num_key_value_heads * self.head_dim, bias=True) self.o_proj = nn.Linear(self.num_heads * self.head_dim, self.hidden_size, bias=False) self.rotary_emb = DreamRotaryEmbedding(config=self.config) def forward( self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_value: Optional[Cache] = None, output_attentions: bool = False, use_cache: bool = False, cache_position: Optional[torch.LongTensor] = None, position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # will become mandatory in v4.46 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]: bsz, q_len, _ = hidden_states.size() query_states = self.q_proj(hidden_states) key_states = self.k_proj(hidden_states) value_states = self.v_proj(hidden_states) query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2) key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2) value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2) if position_embeddings is None: logger.warning_once( "The attention layers in this model are transitioning from computing the RoPE embeddings internally " "through `position_ids` (2D tensor with the indexes of the tokens), to using externally computed " "`position_embeddings` (Tuple of tensors, containing cos and sin). In v4.46 `position_ids` will be " "removed and `position_embeddings` will be mandatory." ) cos, sin = self.rotary_emb(value_states, position_ids) else: cos, sin = position_embeddings query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin) if past_key_value is not None: cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} # Specific to RoPE models key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs) # repeat k/v heads if n_kv_heads < n_heads key_states = repeat_kv(key_states, self.num_key_value_groups) value_states = repeat_kv(value_states, self.num_key_value_groups) attn_weights = torch.matmul(query_states, key_states.transpose(2, 3)) / math.sqrt(self.head_dim) if attention_mask is not None: # no matter the length, we just slice it causal_mask = attention_mask[:, :, :, : key_states.shape[-2]] attn_weights = attn_weights + causal_mask # upcast attention to fp32 attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query_states.dtype) attn_weights = nn.functional.dropout(attn_weights, p=self.attention_dropout, training=self.training) attn_output = torch.matmul(attn_weights, value_states) if attn_output.size() != (bsz, self.num_heads, q_len, self.head_dim): raise ValueError( f"`attn_output` should be of size {(bsz, self.num_heads, q_len, self.head_dim)}, but is" f" {attn_output.size()}" ) attn_output = attn_output.transpose(1, 2).contiguous() attn_output = attn_output.reshape(bsz, q_len, self.hidden_size) attn_output = self.o_proj(attn_output) if not output_attentions: attn_weights = None return attn_output, attn_weights, past_key_value class DreamSdpaAttention(DreamAttention): """ Dream attention module using torch.nn.functional.scaled_dot_product_attention. This module inherits from `DreamAttention` as the weights of the module stays untouched. The only changes are on the forward pass to adapt to SDPA API. """ # Adapted from DreamAttention.forward def forward( self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_value: Optional[Cache] = None, output_attentions: bool = False, use_cache: bool = False, cache_position: Optional[torch.LongTensor] = None, position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # will become mandatory in v4.46 ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]: if output_attentions: # TODO: Improve this warning with e.g. `model.config.attn_implementation = "manual"` once this is implemented. logger.warning_once( "DreamModel is using DreamSdpaAttention, but `torch.nn.functional.scaled_dot_product_attention` does not support `output_attentions=True`. Falling back to the manual attention implementation, " 'but specifying the manual implementation will be required from Transformers version v5.0.0 onwards. This warning can be removed using the argument `attn_implementation="eager"` when loading the model.' ) return super().forward( hidden_states=hidden_states, attention_mask=attention_mask, position_ids=position_ids, past_key_value=past_key_value, output_attentions=output_attentions, use_cache=use_cache, ) # breakpoint() # ForkedPdb().set_trace() bsz, q_len, _ = hidden_states.size() query_states = self.q_proj(hidden_states) key_states = self.k_proj(hidden_states) value_states = self.v_proj(hidden_states) query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2) key_states = key_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2) value_states = value_states.view(bsz, q_len, self.num_key_value_heads, self.head_dim).transpose(1, 2) if position_embeddings is None: logger.warning_once( "The attention layers in this model are transitioning from computing the RoPE embeddings internally " "through `position_ids` (2D tensor with the indexes of the tokens), to using externally computed " "`position_embeddings` (Tuple of tensors, containing cos and sin). In v4.46 `position_ids` will be " "removed and `position_embeddings` will be mandatory." ) cos, sin = self.rotary_emb(value_states, position_ids) else: cos, sin = position_embeddings query_states, key_states = apply_rotary_pos_emb(query_states, key_states, cos, sin) if past_key_value is not None: cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} # Specific to RoPE models key_states, value_states = past_key_value.update(key_states, value_states, self.layer_idx, cache_kwargs) key_states = repeat_kv(key_states, self.num_key_value_groups) value_states = repeat_kv(value_states, self.num_key_value_groups) # causal_mask = attention_mask # if attention_mask is not None: # no matter the length, we just slice it # causal_mask = attention_mask[:, :, :, : key_states.shape[-2]] # SDPA with memory-efficient backend is currently (torch==2.1.2) bugged with non-contiguous inputs with custom attn_mask, # Reference: https://github.com/pytorch/pytorch/issues/112577. if query_states.device.type == "cuda" and attention_mask is not None: query_states = query_states.contiguous() key_states = key_states.contiguous() value_states = value_states.contiguous() # We dispatch to SDPA's Flash Attention or Efficient kernels via this `is_causal` if statement instead of an inline conditional assignment # in SDPA to support both torch.compile's dynamic shapes and full graph options. An inline conditional prevents dynamic shapes from compiling. # The q_len > 1 is necessary to match with AttentionMaskConverter.to_causal_4d that does not create a causal mask in case q_len == 1. # is_causal = True if causal_mask is None and q_len > 1 else False bool_mask = attention_mask.to(torch.bool) #原始的 # ForkedPdb().set_trace() attn_output = torch.nn.functional.scaled_dot_product_attention( query_states, key_states, value_states, attn_mask=bool_mask , dropout_p=self.attention_dropout if self.training else 0.0, is_causal=False, # hard coded ) attn_output = attn_output.transpose(1, 2).contiguous() attn_output = attn_output.view(bsz, q_len, self.hidden_size) attn_output = self.o_proj(attn_output) return attn_output, None, past_key_value #换成flash attn # attention_interface = ALL_ATTENTION_FUNCTIONS["flash_attention_2"] # # ForkedPdb().set_trace() # attn_output, attn_weights = attention_interface( # self, # query_states, # key_states, # value_states, # attention_mask, # dropout=0.0 if not self.training else self.attention_dropout, # scaling=self.head_dim**-0.5, # sliding_window=None, # position_ids=position_ids, # output_attentions= output_attentions, # use_cache = use_cache # # 其他参数 # ) # # attn_output = attn_output.transpose(1, 2).contiguous() # attn_output = attn_output.view(bsz, q_len, self.hidden_size) # attn_output = self.o_proj(attn_output) # return attn_output, attn_weights, past_key_value class DreamDecoderLayer(nn.Module): def __init__(self, config: DreamConfig, layer_idx: int): super().__init__() self.hidden_size = config.hidden_size if config.sliding_window and config._attn_implementation != "flash_attention_2": logger.warning_once( f"Sliding Window Attention is enabled but not implemented for `{config._attn_implementation}`; " "unexpected results may be encountered." ) # self.self_attn = Dream_ATTENTION_CLASSES[config._attn_implementation](config, layer_idx) self.self_attn = DreamSdpaAttention(config, layer_idx) self.mlp = DreamMLP(config) self.input_layernorm = DreamRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.post_attention_layernorm = DreamRMSNorm(config.hidden_size, eps=config.rms_norm_eps) # @torch.compile def forward( self, hidden_states: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_value: Optional[Tuple[torch.Tensor]] = None, output_attentions: Optional[bool] = False, use_cache: Optional[bool] = False, cache_position: Optional[torch.LongTensor] = None, position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, # will become mandatory in v4.46 **kwargs, ) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]: """ Args: hidden_states (`torch.FloatTensor`): input to the layer of shape `(batch, seq_len, embed_dim)` attention_mask (`torch.FloatTensor`, *optional*): attention mask of size `(batch, sequence_length)` where padding elements are indicated by 0. output_attentions (`bool`, *optional*): Whether or not to return the attentions tensors of all attention layers. See `attentions` under returned tensors for more detail. use_cache (`bool`, *optional*): If set to `True`, `past_key_values` key value states are returned and can be used to speed up decoding (see `past_key_values`). past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states cache_position (`torch.LongTensor` of shape `(sequence_length)`, *optional*): Indices depicting the position of the input sequence tokens in the sequence. position_embeddings (`Tuple[torch.FloatTensor, torch.FloatTensor]`, *optional*): Tuple containing the cosine and sine positional embeddings of shape `(batch_size, seq_len, head_dim)`, with `head_dim` being the embedding dimension of each attention head. kwargs (`dict`, *optional*): Arbitrary kwargs to be ignored, used for FSDP and other methods that injects code into the model """ residual = hidden_states hidden_states = self.input_layernorm(hidden_states) # Self Attention # ForkedPdb().set_trace() hidden_states, self_attn_weights, present_key_value = self.self_attn( hidden_states=hidden_states, attention_mask=attention_mask, position_ids=position_ids, past_key_value=past_key_value, output_attentions=output_attentions, use_cache=use_cache, cache_position=cache_position, position_embeddings=position_embeddings, ) hidden_states = residual + hidden_states # Fully Connected residual = hidden_states hidden_states = self.post_attention_layernorm(hidden_states) hidden_states = self.mlp(hidden_states) hidden_states = residual + hidden_states outputs = (hidden_states,) if output_attentions: outputs += (self_attn_weights,) if use_cache: outputs += (present_key_value,) return outputs class DreamPreTrainedModel(PreTrainedModel): config_class = DreamConfig base_model_prefix = "model" supports_gradient_checkpointing = True _no_split_modules = ["DreamDecoderLayer"] _skip_keys_device_placement = "past_key_values" _supports_flash_attn_2 = True _supports_sdpa = True _supports_cache_class = True _supports_quantized_cache = True _supports_static_cache = True def _init_weights(self, module): std = self.config.initializer_range if isinstance(module, nn.Linear): module.weight.data.normal_(mean=0.0, std=std) if module.bias is not None: module.bias.data.zero_() elif isinstance(module, nn.Embedding): module.weight.data.normal_(mean=0.0, std=std) if module.padding_idx is not None: module.weight.data[module.padding_idx].zero_() @classmethod def from_pretrained( cls, pretrained_model_name_or_path: Optional[Union[str, os.PathLike]], *model_args, config: Optional[Union[PretrainedConfig, str, os.PathLike]] = None, cache_dir: Optional[Union[str, os.PathLike]] = None, ignore_mismatched_sizes: bool = False, force_download: bool = False, local_files_only: bool = False, token: Optional[Union[str, bool]] = None, revision: str = "main", use_safetensors: Optional[bool] = None, weights_only: bool = True, **kwargs, ): _model,_ = super().from_pretrained( pretrained_model_name_or_path, *model_args, config=config, cache_dir=cache_dir, ignore_mismatched_sizes=ignore_mismatched_sizes, force_download=force_download, local_files_only=local_files_only, token=token, revision=revision, use_safetensors=use_safetensors, weights_only=weights_only, **kwargs, ) # _model[0].generation_config # ForkedPdb().set_trace() # NOTE(Lin): we need to override the generation config # because the generation config loaded in `from_pretrained` # does not include all the attributes of DreamGenerationConfig resume_download = kwargs.get("resume_download", None) proxies = kwargs.get("proxies", None) subfolder = kwargs.get("subfolder", "") from_auto_class = kwargs.get("_from_auto", False) from_pipeline = kwargs.get("_from_pipeline", None) _model.generation_config= DreamGenerationConfig.from_pretrained( pretrained_model_name_or_path, cache_dir=cache_dir, force_download=force_download, resume_download=resume_download, proxies=proxies, local_files_only=local_files_only, token=token, revision=revision, subfolder=subfolder, _from_auto=from_auto_class, _from_pipeline=from_pipeline, ) return _model,_ class DreamPrefixLMCache(Cache): def __init__(self): super().__init__() self.past_key_values = {} # this will not be updated beyond the prefilling phase def update( self, key_states: torch.Tensor, value_states: torch.Tensor, layer_idx: int, cache_kwargs = None, ) -> Tuple[torch.Tensor, torch.Tensor]: if layer_idx in self.past_key_values: past_key, past_value = self.past_key_values[layer_idx] key_states = torch.cat((past_key, key_states), dim=-2) value_states = torch.cat((past_value, value_states), dim=-2) return key_states,value_states else: self.past_key_values[layer_idx] = (key_states, value_states) return key_states, value_states def get_seq_length(self, layer_idx: Optional[int] = 0) -> int: """Returns the sequence length of the cached states. A layer index can be optionally passed.""" # TODO: deprecate this function in favor of `cache_position` if len(self.past_key_values) == 0: return 0 else: return self.past_key_values[0][0].shape[-2] def get_max_cache_shape(self) -> Optional[int]: return None import deepspeed class DreamBaseModel(DreamPreTrainedModel):# """ Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`DreamDecoderLayer`] Args: config: DreamConfig """ def __init__(self, config: DreamConfig): super().__init__(config) self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx) self.layers = nn.ModuleList( [DreamDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)] ) self._attn_implementation = config._attn_implementation self.norm = DreamRMSNorm(config.hidden_size, eps=config.rms_norm_eps) self.rotary_emb = DreamRotaryEmbedding(config=config) self.gradient_checkpointing = False # Initialize weights and apply final processing self.audio_model = AudioEncoder() self.audio_projection = ResamplerProjector(512, config.hidden_size) self.post_init() def get_input_embeddings(self): return self.embed_tokens def set_input_embeddings(self, value): self.embed_tokens = value def forward( self, input_ids: torch.LongTensor = None, attention_mask: Optional[torch.Tensor] = None, audios: Optional[torch.FloatTensor] = None, audio_indices: Optional[torch.LongTensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[List[torch.FloatTensor]] = None, inputs_embeds: Optional[torch.FloatTensor] = None, use_cache: Optional[bool] = None, output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, cache_position: Optional[torch.LongTensor] = None, ) -> Union[Tuple, BaseModelOutput]: # ForkedPdb().set_trace() if (past_key_values is None or len(past_key_values) == 0) and audios is not None: audio_embeds, audio_lengths = self.audio_model(audios) # if torch.distributed.get_rank() == 0: # print(f"audio_embeds {audio_embeds.size()}") assert audio_embeds.shape[0] == len(audios) fake_audios = None audio_embeds = self.audio_projection(audio_embeds) # torch.set_printoptions(threshold=100_000) # if torch.distributed.get_rank() == 0: # print(f"audio_embeds {audio_embeds.size()}") # print(f"audio_embeds {audio_embeds.sum()}") # print(f"audios {[x.size() for x in audios]}") # print(f"audios {[x.sum() for x in audios]}") # print(f"input_ids {input_ids.size()}") # print(f"input_ids {input_ids.sum()}") # # print(f"input_ids {input_ids}") # print(f"audio_indices {[x.size() for x in audio_indices]}") # print(f"audio_indices {[x.sum() for x in audio_indices]}") # # print(f"audio_indices {audio_indices}") elif self.training: device = self.get_input_embeddings().weight.data.device dtype = self.get_input_embeddings().weight.data.dtype fake_audios = torch.ones((1, 1, 560), dtype=dtype, device=device) audio_embeds, audio_lengths = self.audio_model(fake_audios) audio_embeds = self.audio_projection(audio_embeds) else: fake_audios = None audio_embeds = None output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions output_hidden_states = ( output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states ) use_cache = use_cache if use_cache is not None else self.config.use_cache return_dict = return_dict if return_dict is not None else self.config.use_return_dict if (input_ids is None) ^ (inputs_embeds is not None): raise ValueError("You must specify exactly one of input_ids or inputs_embeds") if self.gradient_checkpointing and self.training: if use_cache: logger.warning_once( "`use_cache=True` is incompatible with gradient checkpointing. Setting `use_cache=False`..." ) use_cache = False if inputs_embeds is None: inputs_embeds = self.embed_tokens(input_ids) if fake_audios is not None: inputs_embeds = inputs_embeds + audio_embeds.mean() * 0.0 elif audio_embeds is not None: inputs_embeds = inputs_embeds.clone() for audio_embeds_, audio_lengths_, audio_indices_ in zip(audio_embeds, audio_lengths, audio_indices,): # print(f"{audio_embeds_.size()=} {audio_lengths_=} {audio_indices_.size()=}") audio_embeds_ = audio_embeds_[:audio_lengths_, ...] audio_embeds_ = audio_embeds_.to(inputs_embeds.device) indices_b, indices_s = audio_indices_.to(inputs_embeds.device).unbind(dim=0) inputs_embeds[indices_b.view(-1), indices_s.view(-1)] = audio_embeds_.view(-1, audio_embeds_.shape[-1]) # inputs_embeds = inputs_embeds + audio_embeds.mean() * 0.0 if use_cache and past_key_values is None: past_key_values = DreamPrefixLMCache() if cache_position is None: past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0 cache_position = torch.arange( past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device ) if position_ids is None: position_ids = cache_position.unsqueeze(0) hidden_states = inputs_embeds # create position embeddings to be shared across the decoder layers position_embeddings = self.rotary_emb(hidden_states, position_ids) # decoder layers all_hidden_states = () if output_hidden_states else None all_self_attns = () if output_attentions else None for decoder_layer in self.layers: if output_hidden_states: all_hidden_states += (hidden_states,) if self.gradient_checkpointing and self.training: layer_outputs = deepspeed.checkpointing.checkpoint( decoder_layer, hidden_states, attention_mask, position_ids, past_key_values, output_attentions, use_cache, cache_position, position_embeddings, ) else: layer_outputs = decoder_layer( hidden_states, attention_mask=attention_mask, position_ids=position_ids, past_key_value=past_key_values, output_attentions=output_attentions, use_cache=use_cache, cache_position=cache_position, position_embeddings=position_embeddings, ) # breakpoint() if isinstance(layer_outputs,torch.Tensor): layer_outputs = (layer_outputs,None) hidden_states = layer_outputs[0] if output_attentions: all_self_attns += (layer_outputs[1],) hidden_states = self.norm(hidden_states) # add hidden states from the last decoder layer if output_hidden_states: all_hidden_states += (hidden_states,) if not return_dict: return tuple(v for v in [hidden_states, all_hidden_states, all_self_attns] if v is not None) return BaseModelOutput( last_hidden_state=hidden_states, hidden_states=all_hidden_states, attentions=all_self_attns, past_key_values=past_key_values, ) class DreamModel(DreamGenerationMixin, DreamPreTrainedModel): _tied_weights_keys = ["lm_head.weight"] def __init__(self, config): super().__init__(config) self.model = DreamBaseModel(config) self.vocab_size = config.vocab_size self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) # Initialize weights and apply final processing self.tokenizer = None self.post_init() def reset_rope_parameters(self): self.model.rotary_emb.reset_parameters() for layer in self.model.layers: layer.self_attn.rotary_emb.reset_parameters() def get_input_embeddings(self): return self.model.embed_tokens def set_input_embeddings(self, value): self.model.embed_tokens = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embeddings): self.lm_head = new_embeddings def set_decoder(self, decoder): self.model = decoder def get_decoder(self): return self.model def forward( self, input_ids: torch.LongTensor = None, attention_mask: Optional[torch.Tensor] = None, audios: Optional[torch.FloatTensor] = None, audio_indices: Optional[torch.LongTensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[List[torch.FloatTensor]] = None, inputs_embeds: Optional[torch.FloatTensor] = None, labels: Optional[torch.LongTensor] = None, use_cache: Optional[bool] = None, output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, cache_position: Optional[torch.LongTensor] = None, num_logits_to_keep: int = 0, num_items_in_batch: int = None, **loss_kwargs, ) -> Union[Tuple, MaskedLMOutput]: # eos_id = 151643 # 自定义 # mask_id = 151666 # 自定义 # import pdb; pdb.set_trace() # ⚠ 调试断点,如无需要可删 # raw_inputs_ids = input_ids # 保留原始 ID,后续需要对齐 labels # --------------------------------------------------------- # 1. 将 位置从注意力 & labels 中临时移除(参见 Sec B.1) # --------------------------------------------------------- #最终的输出也需要 eos,或者说输入的最后就应该全是eos # non_padding = ~(raw_inputs_ids == eos_id) # 强制让 位在 attention_mask 中视为 *可被注意*(True) # attention_mask[raw_inputs_ids == eos_id] = True # labels 位置恢复成 eos_id,避免被 -100 忽略 # 更新labels,让模型能学到eos,但是又不想弄太多eos来影响训练 # input_ids, labels = update_labels_and_inputs(input_ids,labels,eos_id,300) # new_attention_mask, first_non_neg_idx_list, last_non_neg_idx_list = generate_attention_mask(new_lables) # first_non_neg_idx_list里面可能有None # --------------------------------------------------------- # 3. 若存在 labels(训练模式),进行 Forward‑Process: # • 采样需要 Mask 的 token 下标 (masked_indices) # • 为每个样本构造互补分支 (masked / inverse masked) # • 拼接两条分支,得到 2×batch 的输入 / labels # --------------------------------------------------------- # if labels is not None: # # audio 这一块应该没有这个必要? 不过训练也行吧 # labels_mask = ~(labels == -100) # label != -100,assitant 部分 # # noise_embeddings = self.get_input_embeddings()(torch.tensor([mask_id]).to(raw_inputs_ids)) # (1, D) # bsz, seq_len = labels_mask.shape # # noise_embeddings = noise_embeddings.view(1, 1, -1) # mask token 的embedding # # 生成masked_indices, p_mask # masked_indices, p_mask = forward_process( # bsz, seq_len, raw_inputs_ids.device, labels # ) # # ForkedPdb().set_trace() # # 只mask有效token # final_masked_indices = masked_indices & labels_mask # final_masked_indices_inv = (~masked_indices) & labels_mask # # mask_id要和input_ids类型、设备一致 # mask_id_tensor = torch.full_like(input_ids, mask_id) # input_ids = torch.where(final_masked_indices, mask_id_tensor, input_ids) # # new_labels是labels的clone # new_labels = labels.clone() # new_labels[final_masked_indices_inv] = -100 # final_masked_indices_inv = (~masked_indices) & labels_mask #assistant并且没有被mask部分 # 使用 torch.where 将目标 token 替换为噪声向量 # #这里改成把输入换成mask tokne的id就行 # inputs_embeds_inv = torch.where(final_masked_indices_inv.view(bsz, seq_len, 1), # noise_embeddings, inputs_embeds) #没被mask部分 # inputs_embeds = torch.where(final_masked_indices.view(bsz, seq_len, 1), # noise_embeddings, inputs_embeds) #mask部分 # ForkedPdb().set_trace() # 构造两份 labels:各自只在对应分支需要预测的位置保留真值,其余填 -100 # labels_inv = labels.clone() # labels_inv[~final_masked_indices_inv] = -100 # labels[~final_masked_indices] = -100 # 将两条分支沿 batch 维度拼接: # 文章里面说的是,视觉元素可能出现在没被mask的地方,导致训了没啥用 # inputs_embeds = torch.cat([inputs_embeds, inputs_embeds_inv]) # labels = torch.cat([labels, labels_inv]) # final_masked_indices = torch.cat([final_masked_indices, final_masked_indices_inv]) # Debug: 打印序列长度 # seq_len = labels.shape[-1] # print(f"[forward] seq_len={seq_len}") # --------------------------------------------------------- # 4. (可选) DPO ‑style 正/反样本前向;此处暂未实现 # --------------------------------------------------------- # if dpo_forward: # raise NotImplementedError("DPO forward 尚未实现,请按需补充") # ForkedPdb().set_trace() # --------------------------------------------------------- # 5. 常规前向 — 调用基类实现 # --------------------------------------------------------- # attention_mask = None # ⚠ 此处把 mask 置空,让基类自己处理(或依赖 ALiBi) #import pdb; pdb.set_trace() #import time #print(f"begin forward - {time.time()} - {input_ids.device}") num_items_in_batch = None if ENFORCE_NUM_ITEMIN_BATCH: num_items_in_batch = labels.ne(-100).sum() num_items_in_batch = torch.distributed.reduce(num_items_in_batch) # ForkedPdb().set_trace() # lables = new_labels # new_attention_mask = None # attention_mask = new_attention_mask is_new = position_ids == 0 # is_new[0] = True segment_id = torch.cumsum(is_new.long(), dim=1) - 1 new_attention_mask = (segment_id.unsqueeze(1) == segment_id.unsqueeze(2)).long() # ForkedPdb().set_trace() mask = attention_mask.unsqueeze(-1) # [bs, len, 1] new_attention_mask = new_attention_mask * mask # [bs, len, len] * [bs, len, 1],自动broadcast if self.config.chunk_size > 0: item_start_id = torch.where(position_ids[0] == 0)[0] im_start_id = torch.where(input_ids[0] == self.tokenizer.encode("<|im_start|>")[0])[0].tolist() chunk_mask = torch.zeros_like(new_attention_mask) for item_i in range(len(item_start_id)): im_start = item_start_id[item_i] im_end = item_start_id[item_i + 1] if item_i != len(item_start_id) - 1 else input_ids.shape[-1] im_index = im_start_id.index(im_start) chunk_begin = im_start while im_index < len(im_start_id) and im_start_id[im_index] < im_end: if self.tokenizer.decode(input_ids[0, im_start_id[im_index]+1]) == "assistant": chunk_id = 1 ans_begin = im_start_id[im_index] ans_end = im_start_id[im_index + 1] if im_index != len(im_start_id) - 1 else input_ids.shape[-1] while 1: chunk_end = min(ans_begin + chunk_id * self.config.chunk_size, ans_end) chunk_mask[:, chunk_begin: chunk_end, im_start: chunk_end] = 1 if chunk_end == ans_end: break chunk_id += 1 chunk_begin = chunk_end chunk_begin = chunk_end im_index += 1 else: im_index += 1; continue new_attention_mask = new_attention_mask * chunk_mask # visualization # import matplotlib.pyplot as plt # mask_np = new_attention_mask[0].detach().cpu().numpy() # plt.figure(figsize=(10, 8)); plt.imshow(mask_np); plt.savefig("tmp.png"); plt.close() # if chunk_num != (position_ids == 0).sum(): import pdb; pdb.set_trace() output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions output_hidden_states = ( output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states ) return_dict = return_dict if return_dict is not None else self.config.use_return_dict # import pdb;pdb.set_trace() # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn) # position_ids = torch.arange(input_ids.size(1), dtype=torch.long).unsqueeze(0) # position_ids = torch.arange( # input_ids.size(1), # dtype=torch.long, # device=input_ids.device # ).unsqueeze(0).expand(input_ids.size(0), -1) # print(input_ids.shape,labels.shape) #import time #print(f"self.model forward - {time.time()} - {input_ids.device}") outputs = self.model( input_ids=input_ids, attention_mask=new_attention_mask, audios=audios, audio_indices=audio_indices, position_ids=position_ids, past_key_values=past_key_values, inputs_embeds=inputs_embeds, use_cache=use_cache, output_attentions=output_attentions, output_hidden_states=output_hidden_states, return_dict=return_dict, cache_position=cache_position, ) hidden_states = outputs[0] # Only compute necessary logits, and do not upcast them to float if we are not computing the loss logits = self.lm_head(hidden_states[:, -num_logits_to_keep:, :]) # import pdb;pdb.set_trace() loss = None if labels is not None: if ENFORCE_NUM_ITEMIN_BATCH: assert num_items_in_batch is not None, "num_items_in_batch must be provided if ENFORCE_NUM_ITEMIN_BATCH is True" # ForkedPdb().set_trace() loss = self.loss_function(logits, labels, self.vocab_size,num_items_in_batch=num_items_in_batch, **loss_kwargs) if not return_dict: output = (logits,) + outputs[1:] return (loss,) + output if loss is not None else output # ForkedPdb().set_trace() #import time #print(f"forward finish - {time.time()} - {input_ids.device}") # loss_t2i = None # loss_s2i = None # loss_vqa = None # loss_svqa = None # loss_asr = None # loss_tts = None # loss_tqa = None # loss_sqa = None # input_text = self.tokenizer.decode(input_ids[0]) # t2i_prompt = get_t2i_prompt() # for p in t2i_prompt: # if p in input_text: loss_t2i = loss.detach().copy(); break # if "Convert the speech to text." in input_text: loss_asr = loss.detach().copy() # elif "Convert the text to speech." in input_text: loss_tts = loss.detach().copy() # elif "<|image" not in input_text and "<|audio" not in input_text and loss_t2i is None: loss_tqa = loss.detach().copy() # elif "<|image|>" in input_text and loss_t2i is None: in input_text: loss_vqa = loss.detach().copy() # elif "Please response the input audio." in input_text: loss_sqa = loss.detach().copy() # elif "Please generate an image based on the input audio." in input_text: loss_s2i = loss.detach().copy() # elif "Please response the input audio based on the given image." in input_text: loss_svqa = loss.detach().copy() return MaskedLMOutput( loss=loss, # loss_asr=loss_asr, # loss_tts=loss_tts, # loss_tqa=loss_tqa, # loss_sqa=loss_sqa, # loss_t2i=loss_t2i, # loss_s2i=loss_s2i, # loss_vqa=loss_vqa, # loss_svqa=loss_svqa, logits=logits, hidden_states=outputs.hidden_states, attentions=outputs.attentions, past_key_values=outputs.past_key_values ) def forward_dream( self, input_ids: torch.LongTensor = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[List[torch.FloatTensor]] = None, inputs_embeds: Optional[torch.FloatTensor] = None, labels: Optional[torch.LongTensor] = None, use_cache: Optional[bool] = None, output_attentions: Optional[bool] = None, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, cache_position: Optional[torch.LongTensor] = None, num_logits_to_keep: int = 0, **loss_kwargs, ) -> Union[Tuple, MaskedLMOutput]: output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions output_hidden_states = ( output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states ) return_dict = return_dict if return_dict is not None else self.config.use_return_dict attention_mask = None # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn) # import pdb;pdb.set_trace() outputs = self.model( input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, past_key_values=past_key_values, inputs_embeds=inputs_embeds, use_cache=use_cache, output_attentions=output_attentions, output_hidden_states=output_hidden_states, return_dict=return_dict, cache_position=cache_position, ) hidden_states = outputs[0] # Only compute necessary logits, and do not upcast them to float if we are not computing the loss logits = self.lm_head(hidden_states[:, -num_logits_to_keep:, :]) loss = None if labels is not None: loss = self.loss_function(logits, labels, self.vocab_size, **loss_kwargs) if not return_dict: output = (logits,) + outputs[1:] return (loss,) + output if loss is not None else output return MaskedLMOutput( loss=loss, logits=logits, hidden_states=outputs.hidden_states, attentions=outputs.attentions, past_key_values=outputs.past_key_values, ) @torch.no_grad() def generate( self, input_ids: Optional[torch.Tensor] = None, audios: Optional[torch.FloatTensor] = None, audio_indices: Optional[torch.LongTensor] = None, max_new_tokens=512, steps=512, temperature=0.2, top_p=0.95, alg_temp=0., alg="entropy", output_history=False, **kwargs, ): # modalities = kwargs.pop("modalities", None) if "modalities" in kwargs and modalities is None else modalities position_ids = kwargs.pop("position_ids", None) attention_mask = kwargs.pop("attention_mask", None) if "inputs_embeds" in kwargs: raise NotImplementedError("`inputs_embeds` is not supported") # import pdb;pdb.set_trace() # if images is not None: # (inputs, position_ids, attention_mask, _, inputs_embeds, _) = self.prepare_inputs_labels_for_multimodal(inputs, position_ids, attention_mask, None, None, images, modalities, image_sizes=image_sizes) # else: # # breakpoint() # inputs_embeds = self.get_model().embed_tokens(inputs) #return super().generate(position_ids=position_ids, attention_mask=attention_mask, inputs_embeds=inputs_embeds, **kwargs) #return llada_generate(self.get_model(),inputs_embeds=inputs_embeds,position_ids=position_ids,attention_mask=attention_mask,**kwargs) # breakpoint() # ForkedPdb().set_trace() if audios is not None: audio_embeds, audio_lengths = self.model.audio_model(audios) # if torch.distributed.get_rank() == 0: # print(f"audio_embeds {audio_embeds.size()}") assert audio_embeds.shape[0] == len(audios) fake_audios = None audio_embeds = self.model.audio_projection(audio_embeds) # torch.set_printoptions(threshold=100_000) # if torch.distributed.get_rank() == 0: # print(f"audio_embeds {audio_embeds.size()}") # print(f"audio_embeds {audio_embeds.sum()}") # print(f"audios {[x.size() for x in audios]}") # print(f"audios {[x.sum() for x in audios]}") # print(f"input_ids {input_ids.size()}") # print(f"input_ids {input_ids.sum()}") # # print(f"input_ids {input_ids}") # print(f"audio_indices {[x.size() for x in audio_indices]}") # print(f"audio_indices {[x.sum() for x in audio_indices]}") # # print(f"audio_indices {audio_indices}") elif self.training: device = self.model.get_input_embeddings().weight.data.device dtype = self.model.get_input_embeddings().weight.data.dtype fake_audios = torch.ones((1, 1, 560), dtype=dtype, device=device) audio_embeds, audio_lengths = self.model.audio_model(fake_audios) audio_embeds = self.model.audio_projection(audio_embeds) else: fake_audios = None audio_embeds = None # if inputs_embeds is None: inputs_embeds = self.model.embed_tokens(input_ids) if fake_audios is not None: inputs_embeds = inputs_embeds + audio_embeds.mean() * 0.0 elif audio_embeds is not None: inputs_embeds = inputs_embeds.clone() for audio_embeds_, audio_lengths_, audio_indices_ in zip(audio_embeds, audio_lengths, audio_indices,): # print(f"{audio_embeds_.size()=} {audio_lengths_=} {audio_indices_.size()=}") audio_embeds_ = audio_embeds_[:audio_lengths_, ...] audio_embeds_ = audio_embeds_.to(inputs_embeds.device) indices_b, indices_s = audio_indices_.to(inputs_embeds.device).unbind(dim=0) inputs_embeds[indices_b.view(-1), indices_s.view(-1)] = audio_embeds_.view(-1, audio_embeds_.shape[-1]) # inputs_embeds = inputs_embeds + audio_embeds.mean() * 0.0 return self.diffusion_generate( None, inputs_embeds=inputs_embeds, max_new_tokens=max_new_tokens, output_history=output_history, return_dict_in_generate=True, steps=steps, temperature=temperature, top_p=top_p, alg=alg, alg_temp=alg_temp, **kwargs ) # class LlavaDreamForMaskedDiffusion(DreamModel,DreamPreTrainedModel): # # config_class = LlavaDreamConfig # supports_gradient_checkpointing = True # def __init__(self, config: DreamConfig, model: Optional[DreamModel] = None, init_params: bool = False,vision_kwargs=None,**kwargs): # DreamModel.__init__(self, config) # # configure default generation settings # config.model_type = "llava_dream" # # config.rope_scaling = None # # if not model: # self.model = DreamModel(config) # # else: # # self.model = model # #self.model.set_activation_checkpointing('whole_layer') # self.post_init() # TODO # def get_model(self): # return self.model