Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion python/freetoken/engine/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,7 @@ class EngineConfig:
# ratio default above. A runtime cache rebuild sets this (num_swa_pages) to pin the window
# regardless of the full anchor; the ratio is the startup default and the fallback.
swa_num_pages_override: int | None = None
distributed_timeout: float = 60.0
distributed_timeout: float = 1800.0 # ranks reach the first collective minutes apart on a 100+ GiB offload load
use_dummy_weight: bool = False
use_pynccl: bool = True
max_seq_len_override: int | None = None
Expand Down
8 changes: 6 additions & 2 deletions python/freetoken/layers/linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,10 +59,14 @@ def __init__(
input_size: int,
output_sizes: List[int],
has_bias: bool,
local_output_sizes: List[int] | None = None,
):
# check that all output sizes are divisible by tp_size
# check that all output sizes are divisible by tp_size (a caller that replicates
# GQA kv heads across ranks passes the per-rank sizes explicitly)
tp_info = get_tp_info()
tp_output_sizes = [div_even(size, tp_info.size) for size in output_sizes]
if local_output_sizes is None:
local_output_sizes = [div_even(size, tp_info.size) for size in output_sizes]
tp_output_sizes = local_output_sizes
output_size = sum(output_sizes)
tp_output_size = sum(tp_output_sizes)
super().__init__(input_size, output_size, input_size, tp_output_size, has_bias)
Expand Down
123 changes: 51 additions & 72 deletions python/freetoken/models/nvfp4_banks.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,8 @@

import safetensors
import torch
from freetoken.utils import download_hf_weight
from freetoken.distributed import get_tp_info
from freetoken.utils import div_even, download_hf_weight
from tqdm import tqdm

LayerToBank = Callable[[int, object], int | None]
Expand Down Expand Up @@ -78,6 +79,41 @@ def _alloc_nvfp4_host_banks(num_layers: int, E: int, H: int, I: int):
}, num_layers)


def _tp_slice(inter: int) -> tuple[int, int]:
"""``(i_local, i_lo)``: this rank's slice of the intermediate axis. TP shards every expert
along I (gate/up rows, down columns), the ``stream_moe_expert_sources`` convention, so the
routed output is a partial sum the MoE layer all-reduces."""
tp = get_tp_info()
i_local = div_even(inter, tp.size)
assert i_local % 16 == 0, f"NVFP4 TP shard {i_local} must cover whole 16-wide scale blocks"
return i_local, tp.rank * i_local


class _Placer:
"""Writes one checkpoint expert tensor into its bank slot (this rank's I slice only)."""

def __init__(self, banks: dict, inter: int):
self.b = banks
self.i_local, self.i_lo = _tp_slice(inter)

def put(self, layer: int, expert: int, role: str, kind: str, tensor, global_scale=None):
n, lo = self.i_local, self.i_lo
b = self.b
if role == "down": # [H, I/2] codes, [H, I/16] scales, [H] global
if kind == "weight":
b["down_packed"][layer][expert] = tensor[:, lo // 2 : (lo + n) // 2]
else:
b["down_scale"][layer][expert] = tensor[:, lo // 16 : (lo + n) // 16]
b["down_global"][layer][expert] = global_scale
return
rows = slice(0, n) if role == "gate" else slice(n, 2 * n) # gate | up on the row axis
if kind == "weight":
b["gate_up_packed"][layer][expert, rows] = tensor[lo : lo + n]
else:
b["gate_up_scale"][layer][expert, rows] = tensor[lo : lo + n]
b["gate_up_global"][layer][expert, rows] = global_scale # per-tensor scalar


def load_nvfp4_expert_source_banks(
model_path: str,
config,
Expand Down Expand Up @@ -149,13 +185,9 @@ def load_nvfp4_expert_source_banks(
globals_map[key] = _ingest_global(spec, f.get_tensor(name))
drop_page_cache(path)

_hb = _alloc_nvfp4_host_banks(num_layers, E, H, I) # unpinned; pinned after fill
gate_up_packed = [b.tensor for b in _hb["gate_up_packed"]]
gate_up_scale = [b.tensor for b in _hb["gate_up_scale"]]
gate_up_global = [b.tensor for b in _hb["gate_up_global"]]
down_packed = [b.tensor for b in _hb["down_packed"]]
down_scale = [b.tensor for b in _hb["down_scale"]]
down_global = [b.tensor for b in _hb["down_global"]]
_hb = _alloc_nvfp4_host_banks(num_layers, E, H, _tp_slice(I)[0]) # unpinned; pinned after fill
banks = {name: [b.tensor for b in layers] for name, layers in _hb.items()}
place = _Placer(banks, I)

from freetoken.moe.host_banks import LayerCompletionTracker, PinPipeline

Expand All @@ -170,30 +202,11 @@ def _load(sink) -> int:
expert = int(match.group("expert"))
proj = match.group("proj")
role = spec.proj_to_role[proj]
if role not in ("gate", "up", "down"):
raise ValueError(f"{spec.desc}: unknown projection role {role!r}")
kind = _canon_kind(spec, match.group("kind"))
tensor = f.get_tensor(name)
if kind == "weight":
if role == "gate":
gate_up_packed[bank_layer_id][expert, :I] = tensor
elif role == "up":
gate_up_packed[bank_layer_id][expert, I:] = tensor
elif role == "down":
down_packed[bank_layer_id][expert] = tensor
else:
raise ValueError(f"{spec.desc}: unknown projection role {role!r}")
else:
global_scale = globals_map[(layer, expert, proj)]
if role == "gate":
gate_up_scale[bank_layer_id][expert, :I] = tensor
gate_up_global[bank_layer_id][expert, :I] = global_scale
elif role == "up":
gate_up_scale[bank_layer_id][expert, I:] = tensor
gate_up_global[bank_layer_id][expert, I:] = global_scale
elif role == "down":
down_scale[bank_layer_id][expert] = tensor
down_global[bank_layer_id][expert] = global_scale
else:
raise ValueError(f"{spec.desc}: unknown projection role {role!r}")
place.put(bank_layer_id, expert, role, kind, f.get_tensor(name),
None if kind == "weight" else globals_map[(layer, expert, proj)])
tracker.note(bank_layer_id)
placed += 1
drop_page_cache(path)
Expand All @@ -207,14 +220,7 @@ def _load(sink) -> int:

expected = num_layers * E * 6
assert placed == expected, f"{spec.desc}: loaded {placed} expert tensors, expected {expected}"
return {
"gate_up_packed": gate_up_packed,
"gate_up_scale": gate_up_scale,
"gate_up_global": gate_up_global,
"down_packed": down_packed,
"down_scale": down_scale,
"down_global": down_global,
}
return banks


def load_nvfp4_expert_source_banks_parallel(
Expand Down Expand Up @@ -273,13 +279,9 @@ def load_nvfp4_expert_source_banks_parallel(
)
drop_page_cache(path)

_hb = _alloc_nvfp4_host_banks(num_layers, E, H, I) # unpinned; pinned after fill
gate_up_packed = [b.tensor for b in _hb["gate_up_packed"]]
gate_up_scale = [b.tensor for b in _hb["gate_up_scale"]]
gate_up_global = [b.tensor for b in _hb["gate_up_global"]]
down_packed = [b.tensor for b in _hb["down_packed"]]
down_scale = [b.tensor for b in _hb["down_scale"]]
down_global = [b.tensor for b in _hb["down_global"]]
_hb = _alloc_nvfp4_host_banks(num_layers, E, H, _tp_slice(I)[0]) # unpinned; pinned after fill
banks = {name: [b.tensor for b in layers] for name, layers in _hb.items()}
place = _Placer(banks, I)

from freetoken.moe.host_banks import LayerCompletionTracker, PinPipeline

Expand All @@ -296,24 +298,8 @@ def _load(sink) -> int:
proj = match.group("proj")
role = spec.proj_to_role[proj]
kind = _canon_kind(spec, match.group("kind"))
if kind == "weight":
if role == "gate":
gate_up_packed[bank_layer_id][expert, :I] = tensor
elif role == "up":
gate_up_packed[bank_layer_id][expert, I:] = tensor
else:
down_packed[bank_layer_id][expert] = tensor
else:
g = globals_map[(layer, expert, proj)]
if role == "gate":
gate_up_scale[bank_layer_id][expert, :I] = tensor
gate_up_global[bank_layer_id][expert, :I] = g
elif role == "up":
gate_up_scale[bank_layer_id][expert, I:] = tensor
gate_up_global[bank_layer_id][expert, I:] = g
else:
down_scale[bank_layer_id][expert] = tensor
down_global[bank_layer_id][expert] = g
place.put(bank_layer_id, expert, role, kind, tensor,
None if kind == "weight" else globals_map[(layer, expert, proj)])
tracker.note(bank_layer_id)
placed += 1
return placed
Expand All @@ -326,14 +312,7 @@ def _load(sink) -> int:

expected = num_layers * E * 6
assert placed == expected, f"{spec.desc}: loaded {placed} expert tensors, expected {expected}"
return {
"gate_up_packed": gate_up_packed,
"gate_up_scale": gate_up_scale,
"gate_up_global": gate_up_global,
"down_packed": down_packed,
"down_scale": down_scale,
"down_global": down_global,
}
return banks


__all__ = [
Expand Down
39 changes: 28 additions & 11 deletions python/freetoken/models/qwen4_exp/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,9 +19,16 @@

import torch
from freetoken.core import get_global_ctx
from freetoken.layers import BaseOP, GemmaPlusOneRMSNorm, LinearColParallelMerged, LinearReplicated
from freetoken.distributed import get_tp_info
from freetoken.layers import (
BaseOP,
GemmaPlusOneRMSNorm,
LinearColParallelMerged,
LinearOProj,
LinearReplicated,
)
from freetoken.layers.rotary import get_rope
from freetoken.utils import nvtx_annotate
from freetoken.utils import div_even, nvtx_annotate

if TYPE_CHECKING:
from freetoken.core import Batch
Expand Down Expand Up @@ -120,11 +127,21 @@ def __init__(self, config: ModelConfig, layer_id: int) -> None:
self.head_dim = config.head_dim
self.qo_attn_dim = self.num_q * self.head_dim
self.kv_attn_dim = self.num_kv * self.head_dim
self._qkv_split = [self.qo_attn_dim * 2, self.kv_attn_dim, self.kv_attn_dim]
# TP: q heads split across ranks, kv heads split or (num_kv < tp) replicated; the
# indexer stays replicated so every rank selects the same blocks.
tp = get_tp_info()
self._local_num_q = div_even(self.num_q, tp.size)
self._local_num_kv = div_even(self.num_kv, tp.size, allow_replicate=True)
self._local_qo_dim = self._local_num_q * self.head_dim
self._local_kv_dim = self._local_num_kv * self.head_dim
self._qkv_split = [self._local_qo_dim * 2, self._local_kv_dim, self._local_kv_dim]
self.qkv_proj = LinearColParallelMerged(
config.hidden_size, self._qkv_split, has_bias=False
config.hidden_size,
[self.qo_attn_dim * 2, self.kv_attn_dim, self.kv_attn_dim],
has_bias=False,
local_output_sizes=self._qkv_split,
)
self.o_proj = LinearReplicated(self.qo_attn_dim, config.hidden_size, has_bias=False)
self.o_proj = LinearOProj(self.qo_attn_dim, config.hidden_size, has_bias=False)
self.q_norm = GemmaPlusOneRMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.k_norm = GemmaPlusOneRMSNorm(self.head_dim, eps=config.rms_norm_eps)
rotary = config.rotary_config
Expand All @@ -140,21 +157,21 @@ def __init__(self, config: ModelConfig, layer_id: int) -> None:
@nvtx_annotate("QSA")
def forward(self, x: torch.Tensor, batch: Batch) -> torch.Tensor:
qg, k, v = self.qkv_proj.forward(x).split(self._qkv_split, dim=-1)
qg = qg.view(-1, self.num_q, self.head_dim * 2)
qg = qg.view(-1, self._local_num_q, self.head_dim * 2)
q = qg[..., : self.head_dim].contiguous()
gate = qg[..., self.head_dim :].reshape(-1, self.qo_attn_dim)
k = k.contiguous().view(-1, self.num_kv, self.head_dim)
gate = qg[..., self.head_dim :].reshape(-1, self._local_qo_dim)
k = k.contiguous().view(-1, self._local_num_kv, self.head_dim)
v = v.contiguous()
self.q_norm.forward_inplace(q)
self.k_norm.forward_inplace(k)
q, k = self.rotary.forward(
batch.positions, q.view(-1, self.qo_attn_dim), k.view(-1, self.kv_attn_dim)
batch.positions, q.view(-1, self._local_qo_dim), k.view(-1, self._local_kv_dim)
)
index = self.indexer.forward(x)
o = get_global_ctx().attn_backend.qsa_forward(
q.view(-1, self.num_q, self.head_dim), k, v, index, self.layer_id, batch
q.view(-1, self._local_num_q, self.head_dim), k, v, index, self.layer_id, batch
)
gated = o.reshape(-1, self.qo_attn_dim) * torch.sigmoid(gate)
gated = o.reshape(-1, self._local_qo_dim) * torch.sigmoid(gate)
return self.o_proj.forward(gated)


Expand Down
Loading