hugging-apps/echo-memory
0
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 