Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
stepvideo_text_encoder.py555 linesDownload Raw Back to root
1# Copyright 2025 StepFun Inc. All Rights Reserved.2# 3# Permission is hereby granted, free of charge, to any person obtaining a copy4# of this software and associated documentation files (the "Software"), to deal5# in the Software without restriction, including without limitation the rights6# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell7# copies of the Software, and to permit persons to whom the Software is8# furnished to do so, subject to the following conditions:9#10# The above copyright notice and this permission notice shall be included in all11# copies or substantial portions of the Software.12# ==============================================================================13import os14from typing import Optional15 16import torch17import torch.nn as nn18import torch.nn.functional as F19from .stepvideo_dit import RMSNorm20from safetensors.torch import load_file21from transformers.configuration_utils import PretrainedConfig22from transformers.modeling_utils import PreTrainedModel23from einops import rearrange24import json25from typing import List26from functools import wraps27import warnings28 29 30 31class EmptyInitOnDevice(torch.overrides.TorchFunctionMode):32    def __init__(self, device=None):33        self.device = device34 35    def __torch_function__(self, func, types, args=(), kwargs=None):36        kwargs = kwargs or {}37        if getattr(func, '__module__', None) == 'torch.nn.init':38            if 'tensor' in kwargs:39                return kwargs['tensor']40            else:41                return args[0]42        if self.device is not None and func in torch.utils._device._device_constructors() and kwargs.get('device') is None:43            kwargs['device'] = self.device44        return func(*args, **kwargs)45    46 47def with_empty_init(func):48    @wraps(func)49    def wrapper(*args, **kwargs):50        with EmptyInitOnDevice('cpu'):51            return func(*args, **kwargs)52    return wrapper53 54 55 56class LLaMaEmbedding(nn.Module):57    """Language model embeddings.58 59    Arguments:60        hidden_size: hidden size61        vocab_size: vocabulary size62        max_sequence_length: maximum size of sequence. This63                             is used for positional embedding64        embedding_dropout_prob: dropout probability for embeddings65        init_method: weight initialization method66        num_tokentypes: size of the token-type embeddings. 0 value67                        will ignore this embedding68    """69 70    def __init__(self,71                 cfg,72                 ):73        super().__init__()74        self.hidden_size = cfg.hidden_size75        self.params_dtype = cfg.params_dtype76        self.fp32_residual_connection = cfg.fp32_residual_connection 77        self.embedding_weights_in_fp32 = cfg.embedding_weights_in_fp3278        self.word_embeddings = torch.nn.Embedding(79            cfg.padded_vocab_size, self.hidden_size,80        )81        self.embedding_dropout = torch.nn.Dropout(cfg.hidden_dropout)82 83    def forward(self, input_ids):84        # Embeddings.85        if self.embedding_weights_in_fp32:86            self.word_embeddings = self.word_embeddings.to(torch.float32)87        embeddings = self.word_embeddings(input_ids)88        if self.embedding_weights_in_fp32:89            embeddings = embeddings.to(self.params_dtype)90            self.word_embeddings = self.word_embeddings.to(self.params_dtype)91 92        # Data format change to avoid explicit transposes : [b s h] --> [s b h].93        embeddings = embeddings.transpose(0, 1).contiguous()94 95        # If the input flag for fp32 residual connection is set, convert for float.96        if self.fp32_residual_connection:97            embeddings = embeddings.float()98 99        # Dropout.100        embeddings = self.embedding_dropout(embeddings)101 102        return embeddings103 104 105 106class StepChatTokenizer:107    """Step Chat Tokenizer"""108 109    def __init__(110        self, model_file, name="StepChatTokenizer",111        bot_token="<|BOT|>",  # Begin of Turn112        eot_token="<|EOT|>",  # End of Turn113        call_start_token="<|CALL_START|>",      # Call Start114        call_end_token="<|CALL_END|>",          # Call End115        think_start_token="<|THINK_START|>",    # Think Start116        think_end_token="<|THINK_END|>",        # Think End117        mask_start_token="<|MASK_1e69f|>",      # Mask start118        mask_end_token="<|UNMASK_1e69f|>",      # Mask end119    ):120        import sentencepiece121 122        self._tokenizer = sentencepiece.SentencePieceProcessor(model_file=model_file)123 124        self._vocab = {}125        self._inv_vocab = {}126 127        self._special_tokens = {}128        self._inv_special_tokens = {}129 130        self._t5_tokens = []131 132        for idx in range(self._tokenizer.get_piece_size()):133            text = self._tokenizer.id_to_piece(idx)134            self._inv_vocab[idx] = text135            self._vocab[text] = idx136 137            if self._tokenizer.is_control(idx) or self._tokenizer.is_unknown(idx):138                self._special_tokens[text] = idx139                self._inv_special_tokens[idx] = text140 141        self._unk_id = self._tokenizer.unk_id()142        self._bos_id = self._tokenizer.bos_id()143        self._eos_id = self._tokenizer.eos_id()144 145        for token in [146            bot_token, eot_token, call_start_token, call_end_token,147            think_start_token, think_end_token148        ]:149            assert token in self._vocab, f"Token '{token}' not found in tokenizer"150            assert token in self._special_tokens, f"Token '{token}' is not a special token"151 152        for token in [mask_start_token, mask_end_token]:153            assert token in self._vocab, f"Token '{token}' not found in tokenizer"154 155        self._bot_id = self._tokenizer.piece_to_id(bot_token)156        self._eot_id = self._tokenizer.piece_to_id(eot_token)157        self._call_start_id = self._tokenizer.piece_to_id(call_start_token)158        self._call_end_id = self._tokenizer.piece_to_id(call_end_token)159        self._think_start_id = self._tokenizer.piece_to_id(think_start_token)160        self._think_end_id = self._tokenizer.piece_to_id(think_end_token)161        self._mask_start_id = self._tokenizer.piece_to_id(mask_start_token)162        self._mask_end_id = self._tokenizer.piece_to_id(mask_end_token)163 164        self._underline_id = self._tokenizer.piece_to_id("\u2581")165        166    @property167    def vocab(self):168        return self._vocab169 170    @property171    def inv_vocab(self):172        return self._inv_vocab173 174    @property175    def vocab_size(self):176        return self._tokenizer.vocab_size()177 178    def tokenize(self, text: str) -> List[int]:179        return self._tokenizer.encode_as_ids(text)180 181    def detokenize(self, token_ids: List[int]) -> str:182        return self._tokenizer.decode_ids(token_ids)183 184    185class Tokens:186    def __init__(self, input_ids, cu_input_ids, attention_mask, cu_seqlens, max_seq_len) -> None:187        self.input_ids = input_ids188        self.attention_mask = attention_mask189        self.cu_input_ids = cu_input_ids190        self.cu_seqlens = cu_seqlens191        self.max_seq_len = max_seq_len192    def to(self, device):193        self.input_ids = self.input_ids.to(device)194        self.attention_mask = self.attention_mask.to(device)195        self.cu_input_ids = self.cu_input_ids.to(device)196        self.cu_seqlens = self.cu_seqlens.to(device)197        return self198    199class Wrapped_StepChatTokenizer(StepChatTokenizer):200    def __call__(self, text, max_length=320, padding="max_length", truncation=True, return_tensors="pt"):201        # [bos, ..., eos, pad, pad, ..., pad]202        self.BOS = 1203        self.EOS = 2204        self.PAD = 2205        out_tokens = []206        attn_mask = []207        if len(text) == 0:208            part_tokens = [self.BOS] + [self.EOS]209            valid_size = len(part_tokens)210            if len(part_tokens) < max_length:211                part_tokens += [self.PAD] * (max_length - valid_size)212            out_tokens.append(part_tokens)213            attn_mask.append([1]*valid_size+[0]*(max_length-valid_size))214        else:215            for part in text:216                part_tokens = self.tokenize(part)217                part_tokens = part_tokens[:(max_length - 2)] # leave 2 space for bos and eos218                part_tokens = [self.BOS] + part_tokens + [self.EOS]219                valid_size = len(part_tokens)220                if len(part_tokens) < max_length:221                    part_tokens += [self.PAD] * (max_length - valid_size)222                out_tokens.append(part_tokens)223                attn_mask.append([1]*valid_size+[0]*(max_length-valid_size))224 225        out_tokens = torch.tensor(out_tokens, dtype=torch.long)226        attn_mask = torch.tensor(attn_mask, dtype=torch.long)227 228        # padding y based on tp size229        padded_len = 0230        padded_flag = True if padded_len > 0 else False231        if padded_flag:232            pad_tokens = torch.tensor([[self.PAD] * max_length], device=out_tokens.device)233            pad_attn_mask = torch.tensor([[1]*padded_len+[0]*(max_length-padded_len)], device=attn_mask.device)234            out_tokens = torch.cat([out_tokens, pad_tokens], dim=0)235            attn_mask = torch.cat([attn_mask, pad_attn_mask], dim=0)236        237        # cu_seqlens238        cu_out_tokens = out_tokens.masked_select(attn_mask != 0).unsqueeze(0)239        seqlen = attn_mask.sum(dim=1).tolist()240        cu_seqlens = torch.cumsum(torch.tensor([0]+seqlen), 0).to(device=out_tokens.device,dtype=torch.int32)241        max_seq_len = max(seqlen)242        return Tokens(out_tokens, cu_out_tokens, attn_mask, cu_seqlens, max_seq_len)243 244 245 246def flash_attn_func(q, k, v, dropout_p=0.0, softmax_scale=None, causal=True,247                    return_attn_probs=False, tp_group_rank=0, tp_group_size=1):248    softmax_scale = q.size(-1) ** (-0.5) if softmax_scale is None else softmax_scale249    if hasattr(torch.ops.Optimus, "fwd"):250        results = torch.ops.Optimus.fwd(q, k, v, None, dropout_p, softmax_scale, causal, return_attn_probs, None, tp_group_rank, tp_group_size)[0]251    else:252        warnings.warn("Cannot load `torch.ops.Optimus.fwd`. Using `torch.nn.functional.scaled_dot_product_attention` instead.")253        results = torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), is_causal=True, scale=softmax_scale).transpose(1, 2)254    return results255 256 257class FlashSelfAttention(torch.nn.Module):258    def __init__(259        self,260        attention_dropout=0.0,261    ):262        super().__init__()263        self.dropout_p = attention_dropout264 265 266    def forward(self, q, k, v, cu_seqlens=None, max_seq_len=None):267        if cu_seqlens is None:268            output = flash_attn_func(q, k, v, dropout_p=self.dropout_p)269        else:270            raise ValueError('cu_seqlens is not supported!')271 272        return output273 274 275    276def safediv(n, d):277    q, r = divmod(n, d)278    assert r == 0279    return q280 281 282class MultiQueryAttention(nn.Module):283    def __init__(self, cfg, layer_id=None):284        super().__init__()285 286        self.head_dim = cfg.hidden_size // cfg.num_attention_heads287        self.max_seq_len = cfg.seq_length288        self.use_flash_attention = cfg.use_flash_attn289        assert self.use_flash_attention, 'FlashAttention is required!'290 291        self.n_groups = cfg.num_attention_groups292        self.tp_size = 1293        self.n_local_heads = cfg.num_attention_heads294        self.n_local_groups = self.n_groups295 296        self.wqkv = nn.Linear(297            cfg.hidden_size,298            cfg.hidden_size + self.head_dim * 2 * self.n_groups,299            bias=False,300        )301        self.wo = nn.Linear(302            cfg.hidden_size,303            cfg.hidden_size,304            bias=False,305        )306 307        assert self.use_flash_attention, 'non-Flash attention not supported yet.'308        self.core_attention = FlashSelfAttention(attention_dropout=cfg.attention_dropout)309        310        self.layer_id = layer_id311 312    def forward(313        self,314        x: torch.Tensor,315        mask: Optional[torch.Tensor],316        cu_seqlens: Optional[torch.Tensor],317        max_seq_len: Optional[torch.Tensor],318    ):319        seqlen, bsz, dim = x.shape320        xqkv = self.wqkv(x)321 322        xq, xkv = torch.split(323            xqkv,324            (dim // self.tp_size,325             self.head_dim*2*self.n_groups // self.tp_size326            ),327            dim=-1,328        )329 330        # gather on 1st dimension331        xq = xq.view(seqlen, bsz, self.n_local_heads, self.head_dim)332        xkv = xkv.view(seqlen, bsz, self.n_local_groups, 2 * self.head_dim)333        xk, xv = xkv.chunk(2, -1)334 335        # rotary embedding + flash attn336        xq = rearrange(xq, "s b h d -> b s h d")337        xk = rearrange(xk, "s b h d -> b s h d")338        xv = rearrange(xv, "s b h d -> b s h d")339 340        q_per_kv = self.n_local_heads // self.n_local_groups341        if q_per_kv > 1:342            b, s, h, d = xk.size()343            if h == 1:344                xk = xk.expand(b, s, q_per_kv, d)345                xv = xv.expand(b, s, q_per_kv, d)346            else:347                ''' To cover the cases where h > 1, we have348                    the following implementation, which is equivalent to:349                        xk = xk.repeat_interleave(q_per_kv, dim=-2)350                        xv = xv.repeat_interleave(q_per_kv, dim=-2)351                    but can avoid calling aten::item() that involves cpu.352                '''353                idx = torch.arange(q_per_kv * h, device=xk.device).reshape(q_per_kv, -1).permute(1, 0).flatten()354                xk = torch.index_select(xk.repeat(1, 1, q_per_kv, 1), 2, idx).contiguous()355                xv = torch.index_select(xv.repeat(1, 1, q_per_kv, 1), 2, idx).contiguous()356 357        if self.use_flash_attention:358            output = self.core_attention(xq, xk, xv,359                                      cu_seqlens=cu_seqlens,360                                      max_seq_len=max_seq_len)361            # reduce-scatter only support first dimension now362            output = rearrange(output, "b s h d -> s b (h d)").contiguous()363        else:364            xq, xk, xv = [365                rearrange(x, "b s ... -> s b ...").contiguous()366                for x in (xq, xk, xv)367            ]368            output = self.core_attention(xq, xk, xv, mask)369        output = self.wo(output)370        return output371 372 373 374class FeedForward(nn.Module):375    def __init__(376        self,377        cfg,378        dim: int,379        hidden_dim: int,380        layer_id: int,381        multiple_of: int=256,382    ):383        super().__init__()384 385        hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)386        def swiglu(x):387            x = torch.chunk(x, 2, dim=-1)388            return F.silu(x[0]) * x[1]389        self.swiglu = swiglu390            391        self.w1 = nn.Linear(392            dim,393            2 * hidden_dim,394            bias=False,395        )396        self.w2 = nn.Linear(397            hidden_dim,398            dim,399            bias=False,400        )401 402    def forward(self, x):403        x = self.swiglu(self.w1(x))404        output = self.w2(x)405        return output406 407 408 409class TransformerBlock(nn.Module):410    def __init__(411        self, cfg, layer_id: int412    ):413        super().__init__()414 415        self.n_heads = cfg.num_attention_heads416        self.dim = cfg.hidden_size417        self.head_dim = cfg.hidden_size // cfg.num_attention_heads418        self.attention = MultiQueryAttention(419            cfg,420            layer_id=layer_id,421        )422 423        self.feed_forward = FeedForward(424            cfg,425            dim=cfg.hidden_size,426            hidden_dim=cfg.ffn_hidden_size,427            layer_id=layer_id,428        )429        self.layer_id = layer_id430        self.attention_norm = RMSNorm(431            cfg.hidden_size,432            eps=cfg.layernorm_epsilon,433        )434        self.ffn_norm = RMSNorm(435            cfg.hidden_size,436            eps=cfg.layernorm_epsilon,437        )438 439    def forward(440        self,441        x: torch.Tensor,442        mask: Optional[torch.Tensor],443        cu_seqlens: Optional[torch.Tensor],444        max_seq_len: Optional[torch.Tensor],445    ):446        residual = self.attention.forward(447            self.attention_norm(x), mask,448            cu_seqlens, max_seq_len449        )450        h = x + residual451        ffn_res = self.feed_forward.forward(self.ffn_norm(h))452        out = h + ffn_res453        return out454 455 456class Transformer(nn.Module):457    def __init__(458        self,459        config,460        max_seq_size=8192,461    ):462        super().__init__()463        self.num_layers = config.num_layers464        self.layers = self._build_layers(config)465 466    def _build_layers(self, config):467        layers = torch.nn.ModuleList()468        for layer_id in range(self.num_layers):469            layers.append(470                TransformerBlock(471                    config,472                    layer_id=layer_id + 1 ,473                )474            )475        return layers476 477    def forward(478        self,479        hidden_states,480        attention_mask,481        cu_seqlens=None,482        max_seq_len=None,483    ):484 485        if max_seq_len is not None and not isinstance(max_seq_len, torch.Tensor):486            max_seq_len = torch.tensor(max_seq_len, dtype=torch.int32, device="cpu")487 488        for lid, layer in enumerate(self.layers):489            hidden_states = layer(490                                    hidden_states,491                                    attention_mask,492                                    cu_seqlens,493                                    max_seq_len,494                                )495        return hidden_states496 497 498class Step1Model(PreTrainedModel):499    config_class=PretrainedConfig500    @with_empty_init501    def __init__(502        self,503        config,504    ):505        super().__init__(config)506        self.tok_embeddings = LLaMaEmbedding(config)507        self.transformer = Transformer(config)508 509    def forward(510        self,511        input_ids=None,512        attention_mask=None,513    ):514 515        hidden_states = self.tok_embeddings(input_ids)516 517        hidden_states = self.transformer(518            hidden_states,519            attention_mask,520        )521        return hidden_states522    523    524 525class STEP1TextEncoder(torch.nn.Module):526    def __init__(self, model_dir, max_length=320):527        super(STEP1TextEncoder, self).__init__()528        self.max_length = max_length529        self.text_tokenizer = Wrapped_StepChatTokenizer(os.path.join(model_dir, 'step1_chat_tokenizer.model'))530        text_encoder = Step1Model.from_pretrained(model_dir)531        self.text_encoder = text_encoder.eval().to(torch.bfloat16)532 533    @staticmethod534    def from_pretrained(path, torch_dtype=torch.bfloat16):535        model = STEP1TextEncoder(path).to(torch_dtype)536        return model537        538    @torch.no_grad539    def forward(self, prompts, with_mask=True, max_length=None, device="cuda"):540        self.device = device541        with torch.no_grad(), torch.amp.autocast(dtype=torch.bfloat16, device_type=device):542            if type(prompts) is str:543                prompts = [prompts]544            545            txt_tokens = self.text_tokenizer(546                prompts, max_length=max_length or self.max_length, padding="max_length", truncation=True, return_tensors="pt"547            )548            y = self.text_encoder(549                txt_tokens.input_ids.to(self.device), 550                attention_mask=txt_tokens.attention_mask.to(self.device) if with_mask else None551            )552            y_mask = txt_tokens.attention_mask553        return y.transpose(0,1), y_mask554 555