[diffusion] ERNIE-Image: fuse rotate-half RoPE + GELU-mul and hoist rope cos/sin (denoise -16.2% H100 / -12.7% H200, bit-exact) (#34306)
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Fable 5
parent
ba3dc16401
commit
071f0f1e9d
@@ -200,6 +200,17 @@ def silu_and_mul_with_activation_rounding_(input: torch.Tensor) -> torch.Tensor:
|
||||
return input[..., :hidden_size]
|
||||
|
||||
|
||||
def gelu_and_mul_with_activation_rounding(
|
||||
input: torch.Tensor,
|
||||
out: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
hidden_size = input.shape[-1] // 2
|
||||
if out is None:
|
||||
out = input.new_empty(*input.shape[:-1], hidden_size)
|
||||
_run_activation_with_rounding_inplace("gelu", input, out)
|
||||
return out
|
||||
|
||||
|
||||
def gelu_and_mul(
|
||||
input: torch.Tensor,
|
||||
out: Optional[torch.Tensor] = None,
|
||||
|
||||
@@ -0,0 +1,154 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Bit-exact fused rotate-half RoPE for bf16 ``(B, S, H, D)`` activations.
|
||||
|
||||
Replaces the eager ERNIE-Image per-projection chain
|
||||
|
||||
``cos/sin -> chunk -> cat(-x2, x1) -> two muls + add -> cat(tail)``
|
||||
|
||||
(~7 kernels per q/k, including two full-width concats) with one Triton
|
||||
kernel, reproducing every aten bf16 rounding boundary bit for bit:
|
||||
|
||||
- ``out[i] = round(round(x1 * cos1) + round(-x2 * sin1))``
|
||||
- ``out[i + R/2] = round(round(x2 * cos2) + round( x1 * sin2))``
|
||||
- columns past the rotary span are copied through unchanged (the eager
|
||||
path concatenates them back untouched).
|
||||
|
||||
``cos``/``sin`` are precomputed once per forward as ``(B * S, rot_dim)``
|
||||
bf16 rows — the same values the eager chain materializes per layer via
|
||||
``torch.cos(freqs).to(dtype)`` — so the per-layer trigonometry disappears
|
||||
as well. Negation, the fp32 products and the single-rounded add match
|
||||
aten elementwise semantics exactly (no reductions are involved), which is
|
||||
what makes a lossless default-on mount possible; callers still verify the
|
||||
first call against the eager chain and fall back on any mismatch (see
|
||||
``ernie_image.py``).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import triton # type: ignore
|
||||
import triton.language as tl # type: ignore
|
||||
|
||||
from sglang.kernels.ops.diffusion.triton.numerics import round_bf16_to_fp32
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _rope_rotate_half_kernel(
|
||||
out_ptr,
|
||||
x_ptr,
|
||||
cos_ptr,
|
||||
sin_ptr,
|
||||
heads,
|
||||
D: tl.constexpr,
|
||||
ROT: tl.constexpr,
|
||||
HALF: tl.constexpr,
|
||||
H_BLOCK: tl.constexpr,
|
||||
HALF_BLOCK: tl.constexpr,
|
||||
TAIL_BLOCK: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0).to(tl.int64) # one program per (batch, seq) row
|
||||
base = row * heads * D
|
||||
hs = tl.arange(0, H_BLOCK)[:, None]
|
||||
hmask = hs < heads
|
||||
cols = tl.arange(0, HALF_BLOCK)[None, :]
|
||||
cmask = cols < HALF
|
||||
m = hmask & cmask
|
||||
|
||||
off1 = base + hs * D + cols
|
||||
off2 = off1 + HALF
|
||||
x1 = tl.load(x_ptr + off1, mask=m, other=0.0).to(tl.float32)
|
||||
x2 = tl.load(x_ptr + off2, mask=m, other=0.0).to(tl.float32)
|
||||
cos1 = tl.load(cos_ptr + row * ROT + cols, mask=cmask, other=0.0).to(tl.float32)
|
||||
cos2 = tl.load(cos_ptr + row * ROT + HALF + cols, mask=cmask, other=0.0).to(
|
||||
tl.float32
|
||||
)
|
||||
sin1 = tl.load(sin_ptr + row * ROT + cols, mask=cmask, other=0.0).to(tl.float32)
|
||||
sin2 = tl.load(sin_ptr + row * ROT + HALF + cols, mask=cmask, other=0.0).to(
|
||||
tl.float32
|
||||
)
|
||||
|
||||
# Each product is rounded to bf16 like the eager mul; the store rounds
|
||||
# the fp32 add exactly once, like the eager add.
|
||||
out1 = round_bf16_to_fp32(x1 * cos1) + round_bf16_to_fp32(-x2 * sin1)
|
||||
out2 = round_bf16_to_fp32(x2 * cos2) + round_bf16_to_fp32(x1 * sin2)
|
||||
tl.store(out_ptr + off1, out1, mask=m)
|
||||
tl.store(out_ptr + off2, out2, mask=m)
|
||||
|
||||
if D > ROT:
|
||||
tcols = ROT + tl.arange(0, TAIL_BLOCK)[None, :]
|
||||
tmask = hmask & (tcols < D)
|
||||
toff = base + hs * D + tcols
|
||||
tail = tl.load(x_ptr + toff, mask=tmask, other=0.0)
|
||||
tl.store(out_ptr + toff, tail, mask=tmask)
|
||||
|
||||
|
||||
def can_use_fused_rope_rotate_half(
|
||||
x: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
) -> bool:
|
||||
if x.dtype is not torch.bfloat16 or not x.is_cuda:
|
||||
return False
|
||||
if x.dim() != 4 or not x.is_contiguous():
|
||||
return False
|
||||
rows = x.shape[0] * x.shape[1]
|
||||
rot = cos.shape[-1]
|
||||
return (
|
||||
cos.dtype is torch.bfloat16
|
||||
and sin.dtype is torch.bfloat16
|
||||
and cos.is_cuda
|
||||
and cos.device == x.device
|
||||
and sin.device == x.device
|
||||
and cos.shape == (rows, rot)
|
||||
and sin.shape == (rows, rot)
|
||||
and cos.is_contiguous()
|
||||
and sin.is_contiguous()
|
||||
and rot % 2 == 0
|
||||
and 0 < rot <= x.shape[-1]
|
||||
)
|
||||
|
||||
|
||||
def _fake_rope_rotate_half(
|
||||
x: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
return torch.empty_like(x)
|
||||
|
||||
|
||||
@register_custom_op(
|
||||
op_name="triton_fused_rope_rotate_half_bitexact",
|
||||
mutates_args=[],
|
||||
fake_impl=_fake_rope_rotate_half,
|
||||
)
|
||||
def fused_rope_rotate_half_bitexact(
|
||||
x: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Rotate-half RoPE over the leading ``cos.shape[-1]`` columns of ``x``.
|
||||
|
||||
``x`` is ``(B, S, H, D)``; ``cos``/``sin`` are ``(B * S, rot_dim)`` rows.
|
||||
Bit-exact vs the eager chunk/neg/cat/mul/add chain.
|
||||
"""
|
||||
batch, seq_len, heads, head_dim = x.shape
|
||||
rot = cos.shape[-1]
|
||||
half = rot // 2
|
||||
out = torch.empty_like(x)
|
||||
tail = head_dim - rot
|
||||
with torch.cuda.device(x.device):
|
||||
_rope_rotate_half_kernel[(batch * seq_len,)](
|
||||
out,
|
||||
x,
|
||||
cos,
|
||||
sin,
|
||||
heads,
|
||||
D=head_dim,
|
||||
ROT=rot,
|
||||
HALF=half,
|
||||
H_BLOCK=triton.next_power_of_2(heads),
|
||||
HALF_BLOCK=triton.next_power_of_2(half),
|
||||
TAIL_BLOCK=triton.next_power_of_2(max(tail, 1)),
|
||||
)
|
||||
return out
|
||||
@@ -19,6 +19,9 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from diffusers.models.embeddings import TimestepEmbedding, Timesteps
|
||||
|
||||
from sglang.kernels.ops.activation.activation import (
|
||||
gelu_and_mul_with_activation_rounding,
|
||||
)
|
||||
from sglang.kernels.ops.diffusion.bitexact_gate import (
|
||||
BitExactFusionGate,
|
||||
tensors_equal,
|
||||
@@ -30,6 +33,10 @@ from sglang.kernels.ops.diffusion.triton.rmsnorm_scale_shift_bitexact import (
|
||||
fused_rmsnorm_scale_shift_bitexact,
|
||||
fused_scale_residual_rmsnorm_scale_shift_bitexact,
|
||||
)
|
||||
from sglang.kernels.ops.diffusion.triton.rope_rotate_half_bitexact import (
|
||||
can_use_fused_rope_rotate_half,
|
||||
fused_rope_rotate_half_bitexact,
|
||||
)
|
||||
from sglang.multimodal_gen.configs.models.dits.ernie_image import (
|
||||
ErnieImageDitConfig,
|
||||
)
|
||||
@@ -58,6 +65,8 @@ logger = init_logger(__name__)
|
||||
|
||||
_ERNIE_NORM = BitExactFusionGate("ERNIE fused-norm")
|
||||
_ERNIE_GATED_NORM = BitExactFusionGate("ERNIE fused gated-norm")
|
||||
_ERNIE_ROPE = BitExactFusionGate("ERNIE fused RoPE")
|
||||
_ERNIE_GEGLU = BitExactFusionGate("ERNIE fused GELU-mul")
|
||||
|
||||
|
||||
def _eager_norm_scale_shift(
|
||||
@@ -269,7 +278,8 @@ class ErnieImageSelfAttention(nn.Module):
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
rotary_pos_emb: torch.Tensor,
|
||||
rope_cos: torch.Tensor,
|
||||
rope_sin: torch.Tensor,
|
||||
attn_mask: torch.Tensor | None = None,
|
||||
attn_mask_meta: dict | None = None,
|
||||
) -> torch.Tensor:
|
||||
@@ -292,8 +302,8 @@ class ErnieImageSelfAttention(nn.Module):
|
||||
self.head_dim,
|
||||
)
|
||||
|
||||
q = _apply_rotary_bshd(q, rotary_pos_emb)
|
||||
k = _apply_rotary_bshd(k, rotary_pos_emb)
|
||||
q = _ernie_rope(q, rope_cos, rope_sin)
|
||||
k = _ernie_rope(k, rope_cos, rope_sin)
|
||||
|
||||
attn_out = self.attn(
|
||||
q, k, v, attn_mask=attn_mask, attn_mask_meta=attn_mask_meta
|
||||
@@ -328,8 +338,7 @@ class ErnieImageMLP(nn.Module):
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
gate_up, _ = self.gate_up_proj(x)
|
||||
gate, up = gate_up.chunk(2, dim=-1)
|
||||
x = up * F.gelu(gate)
|
||||
x = _ernie_geglu(gate_up)
|
||||
x, _ = self.linear_fc2(x)
|
||||
return x
|
||||
|
||||
@@ -363,7 +372,8 @@ class ErnieImageSharedAdaLNBlock(nn.Module):
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
rotary_pos_emb: torch.Tensor,
|
||||
rope_cos: torch.Tensor,
|
||||
rope_sin: torch.Tensor,
|
||||
shift_msa: torch.Tensor,
|
||||
scale_msa: torch.Tensor,
|
||||
gate_msa: torch.Tensor,
|
||||
@@ -376,7 +386,11 @@ class ErnieImageSharedAdaLNBlock(nn.Module):
|
||||
residual = x
|
||||
x = _ernie_norm_scale_shift(self.adaLN_sa_ln, x, scale_msa, shift_msa)
|
||||
attn_out = self.self_attention(
|
||||
x, rotary_pos_emb, attn_mask=attn_mask, attn_mask_meta=attn_mask_meta
|
||||
x,
|
||||
rope_cos,
|
||||
rope_sin,
|
||||
attn_mask=attn_mask,
|
||||
attn_mask_meta=attn_mask_meta,
|
||||
)
|
||||
x, residual = _ernie_gated_norm_scale_shift(
|
||||
self.adaLN_mlp_ln, residual, attn_out, gate_msa, scale_mlp, shift_mlp
|
||||
@@ -386,21 +400,114 @@ class ErnieImageSharedAdaLNBlock(nn.Module):
|
||||
return x
|
||||
|
||||
|
||||
def _apply_rotary_bshd(x: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor:
|
||||
freqs = freqs.permute(1, 0, 2, 3)
|
||||
rot_dim = freqs.shape[-1]
|
||||
x_rot, x_pass = x[..., :rot_dim], x[..., rot_dim:]
|
||||
def _precompute_rope_cos_sin(
|
||||
freqs: torch.Tensor, dtype: torch.dtype
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""cos/sin of the rotary embedding, computed once per forward.
|
||||
|
||||
cos_ = torch.cos(freqs).to(x.dtype)
|
||||
sin_ = torch.sin(freqs).to(x.dtype)
|
||||
``freqs`` is the ``(S, B, 1, rot_dim)`` output of :class:`EmbedND3`; the
|
||||
eager chain recomputed ``torch.cos(freqs).to(dtype)`` per layer per
|
||||
projection. Returns bit-identical ``(B * S, rot_dim)`` rows.
|
||||
"""
|
||||
freqs = freqs.permute(1, 0, 2, 3)
|
||||
cos_ = torch.cos(freqs).to(dtype)
|
||||
sin_ = torch.sin(freqs).to(dtype)
|
||||
rot_dim = freqs.shape[-1]
|
||||
return cos_.reshape(-1, rot_dim), sin_.reshape(-1, rot_dim)
|
||||
|
||||
|
||||
def _apply_rotary_bshd_eager(
|
||||
x: torch.Tensor, cos_: torch.Tensor, sin_: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""Reference rotate-half chain on precomputed cos/sin (bit-exact vs the
|
||||
original per-layer version, which materialized the same cos/sin)."""
|
||||
batch, seq_len = x.shape[0], x.shape[1]
|
||||
rot_dim = cos_.shape[-1]
|
||||
cos_b = cos_.view(batch, seq_len, 1, rot_dim)
|
||||
sin_b = sin_.view(batch, seq_len, 1, rot_dim)
|
||||
x_rot, x_pass = x[..., :rot_dim], x[..., rot_dim:]
|
||||
|
||||
x1, x2 = x_rot.chunk(2, dim=-1)
|
||||
x_rotated = torch.cat((-x2, x1), dim=-1)
|
||||
|
||||
x_rot = x_rot * cos_ + x_rotated * sin_
|
||||
x_rot = x_rot * cos_b + x_rotated * sin_b
|
||||
return torch.cat((x_rot, x_pass), dim=-1)
|
||||
|
||||
|
||||
def _ernie_rope(
|
||||
x: torch.Tensor, cos_: torch.Tensor, sin_: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
"""Single-kernel rotate-half RoPE, bit-exact vs the eager chain.
|
||||
|
||||
Pure elementwise math, so the Triton kernel reproduces every aten bf16
|
||||
rounding boundary exactly; the first call still verifies ``torch.equal``
|
||||
against the eager chain and disables the fast path on any mismatch.
|
||||
"""
|
||||
verified = _ERNIE_ROPE.verified
|
||||
if (
|
||||
not _ERNIE_ROPE.disabled
|
||||
and can_use_fused_rope_rotate_half(x, cos_, sin_)
|
||||
and (verified or _ERNIE_ROPE.can_attempt_once())
|
||||
):
|
||||
try:
|
||||
out = fused_rope_rotate_half_bitexact(x, cos_, sin_)
|
||||
except Exception as exc:
|
||||
_ERNIE_ROPE.on_exception(exc, logger=logger)
|
||||
else:
|
||||
if verified:
|
||||
return out
|
||||
return _ERNIE_ROPE.accept_or_fallback(
|
||||
out,
|
||||
_apply_rotary_bshd_eager(x, cos_, sin_),
|
||||
logger=logger,
|
||||
mismatch_msg=(
|
||||
"ERNIE fused RoPE fast path is not bit-exact on this "
|
||||
"platform; falling back to eager"
|
||||
),
|
||||
)
|
||||
return _apply_rotary_bshd_eager(x, cos_, sin_)
|
||||
|
||||
|
||||
def _eager_geglu(gate_up: torch.Tensor) -> torch.Tensor:
|
||||
gate, up = gate_up.chunk(2, dim=-1)
|
||||
return up * F.gelu(gate)
|
||||
|
||||
|
||||
def _ernie_geglu(gate_up: torch.Tensor) -> torch.Tensor:
|
||||
"""``up * gelu(gate)`` in one kernel, bit-exact vs the eager pair.
|
||||
|
||||
Uses the activation kernel's rounding variant, which rounds the erf-GELU
|
||||
to bf16 before the multiply exactly like the eager two-step; first call
|
||||
self-verifies like :func:`_ernie_rope`.
|
||||
"""
|
||||
verified = _ERNIE_GEGLU.verified
|
||||
if (
|
||||
not _ERNIE_GEGLU.disabled
|
||||
and gate_up.dtype in (torch.bfloat16, torch.float16)
|
||||
and gate_up.is_cuda
|
||||
and gate_up.is_contiguous()
|
||||
and gate_up.shape[-1] % 2 == 0
|
||||
and (verified or _ERNIE_GEGLU.can_attempt_once())
|
||||
):
|
||||
try:
|
||||
out = gelu_and_mul_with_activation_rounding(gate_up)
|
||||
except Exception as exc:
|
||||
_ERNIE_GEGLU.on_exception(exc, logger=logger)
|
||||
else:
|
||||
if verified:
|
||||
return out
|
||||
return _ERNIE_GEGLU.accept_or_fallback(
|
||||
out,
|
||||
_eager_geglu(gate_up),
|
||||
logger=logger,
|
||||
mismatch_msg=(
|
||||
"ERNIE fused GELU-mul fast path is not bit-exact on this "
|
||||
"platform; falling back to eager"
|
||||
),
|
||||
)
|
||||
return _eager_geglu(gate_up)
|
||||
|
||||
|
||||
class ErnieImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
||||
"""ErnieImage DiT: Single-stream transformer with Shared AdaLN."""
|
||||
|
||||
@@ -575,6 +682,7 @@ class ErnieImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin)
|
||||
|
||||
all_ids = torch.cat([image_ids, text_ids], dim=1)
|
||||
rotary_pos_emb = self.pos_embed(all_ids)
|
||||
rope_cos, rope_sin = _precompute_rope_cos_sin(rotary_pos_emb, dtype)
|
||||
|
||||
attn_mask = attn_mask_meta = None
|
||||
if encoder_hidden_states_mask is not None:
|
||||
@@ -601,7 +709,8 @@ class ErnieImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin)
|
||||
for layer in self.layers:
|
||||
x = layer(
|
||||
x,
|
||||
rotary_pos_emb,
|
||||
rope_cos,
|
||||
rope_sin,
|
||||
shift_msa,
|
||||
scale_msa,
|
||||
gate_msa,
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
import unittest
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.multimodal_gen.runtime.models.dits.ernie_image import (
|
||||
_ernie_geglu,
|
||||
_ernie_rope,
|
||||
_precompute_rope_cos_sin,
|
||||
)
|
||||
|
||||
|
||||
def _reference_rotary_bshd(x: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor:
|
||||
"""The pre-fusion eager chain (per-layer cos/sin) verbatim."""
|
||||
freqs = freqs.permute(1, 0, 2, 3)
|
||||
rot_dim = freqs.shape[-1]
|
||||
x_rot, x_pass = x[..., :rot_dim], x[..., rot_dim:]
|
||||
|
||||
cos_ = torch.cos(freqs).to(x.dtype)
|
||||
sin_ = torch.sin(freqs).to(x.dtype)
|
||||
|
||||
x1, x2 = x_rot.chunk(2, dim=-1)
|
||||
x_rotated = torch.cat((-x2, x1), dim=-1)
|
||||
|
||||
x_rot = x_rot * cos_ + x_rotated * sin_
|
||||
return torch.cat((x_rot, x_pass), dim=-1)
|
||||
|
||||
|
||||
def _make_freqs(batch: int, seq: int, rot: int, device) -> torch.Tensor:
|
||||
# EmbedND3 layout: (S, B, 1, rot), interleave-duplicated frequencies.
|
||||
uniq = torch.randn(seq, batch, 1, rot // 2, device=device) * 3.0
|
||||
return torch.stack([uniq, uniq], dim=-1).reshape(seq, batch, 1, rot)
|
||||
|
||||
|
||||
class TestErnieRopeFusion(unittest.TestCase):
|
||||
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
|
||||
def test_fused_rope_matches_prefusion_chain(self):
|
||||
torch.manual_seed(0)
|
||||
device = torch.device("cuda")
|
||||
for batch, seq, heads, head_dim, rot in [
|
||||
(2, 257, 16, 128, 64),
|
||||
(1, 64, 3, 128, 128),
|
||||
(2, 33, 8, 64, 56),
|
||||
]:
|
||||
x = torch.randn(
|
||||
batch, seq, heads, head_dim, device=device, dtype=torch.bfloat16
|
||||
)
|
||||
freqs = _make_freqs(batch, seq, rot, device)
|
||||
reference = _reference_rotary_bshd(x, freqs)
|
||||
|
||||
cos_, sin_ = _precompute_rope_cos_sin(freqs, torch.bfloat16)
|
||||
fused = _ernie_rope(x, cos_, sin_)
|
||||
self.assertTrue(
|
||||
torch.equal(reference, fused),
|
||||
f"rope mismatch at {(batch, seq, heads, head_dim, rot)}",
|
||||
)
|
||||
|
||||
def test_eager_fallback_matches_prefusion_chain_cpu(self):
|
||||
torch.manual_seed(1)
|
||||
x = torch.randn(2, 17, 4, 32, dtype=torch.float32)
|
||||
freqs = _make_freqs(2, 17, 16, x.device)
|
||||
reference = _reference_rotary_bshd(x, freqs)
|
||||
cos_, sin_ = _precompute_rope_cos_sin(freqs, torch.float32)
|
||||
self.assertTrue(torch.equal(reference, _ernie_rope(x, cos_, sin_)))
|
||||
|
||||
|
||||
class TestErnieGegluFusion(unittest.TestCase):
|
||||
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA")
|
||||
def test_fused_geglu_matches_eager(self):
|
||||
torch.manual_seed(0)
|
||||
gate_up = torch.randn(2, 129, 2 * 3584, device="cuda", dtype=torch.bfloat16)
|
||||
gate, up = gate_up.chunk(2, dim=-1)
|
||||
reference = up * F.gelu(gate)
|
||||
self.assertTrue(torch.equal(reference, _ernie_geglu(gate_up)))
|
||||
|
||||
def test_eager_fallback_cpu(self):
|
||||
torch.manual_seed(1)
|
||||
gate_up = torch.randn(3, 8, dtype=torch.float32)
|
||||
gate, up = gate_up.chunk(2, dim=-1)
|
||||
self.assertTrue(torch.equal(up * F.gelu(gate), _ernie_geglu(gate_up)))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user