MCPcopy Create free account
hub / github.com/THUDM/LongWriter / forward_impl

Method forward_impl

train/patch/modeling_chatglm.py:103–128  ·  view source on GitHub ↗

Enhanced Transformer with Rotary Position Embedding. Derived from: https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/master/labml_nn/ transformers/rope/__init__.py. MIT License: https://github.com/labmlai/annotated_deep_learning_paper_implementati

(
            self, seq_len: int, n_elem: int, dtype: torch.dtype, device: torch.device, base: int = 10000
    )

Source from the content-addressed store, hash-verified

101 self.rope_ratio = rope_ratio
102
103 def forward_impl(
104 self, seq_len: int, n_elem: int, dtype: torch.dtype, device: torch.device, base: int = 10000
105 ):
106 """Enhanced Transformer with Rotary Position Embedding.
107
108 Derived from: https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/master/labml_nn/
109 transformers/rope/__init__.py. MIT License:
110 https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/master/license.
111 """
112 # $\Theta = {\theta_i = 10000^{\frac{2(i-1)}{d}}, i \in [1, 2, ..., \frac{d}{2}]}$
113
114 base = base * self.rope_ratio
115 theta = 1.0 / (base ** (torch.arange(0, n_elem, 2, dtype=torch.float, device=device) / n_elem))
116
117 # Create position indexes `[0, 1, ..., seq_len - 1]`
118 seq_idx = torch.arange(seq_len, dtype=torch.float, device=device)
119
120 # Calculate the product of position index and $\theta_i$
121 idx_theta = torch.outer(seq_idx, theta).float()
122
123 cache = torch.stack([torch.cos(idx_theta), torch.sin(idx_theta)], dim=-1)
124
125 # this is to mimic the behaviour of complex32, else we will get different results
126 if dtype in (torch.float16, torch.bfloat16, torch.int8):
127 cache = cache.bfloat16() if dtype == torch.bfloat16 else cache.half()
128 return cache
129
130 def forward(self, max_seq_len, offset=0):
131 return self.forward_impl(

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected