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
Original file line number Diff line number Diff line change
@@ -0,0 +1,70 @@
#!/usr/bin/env python3
"""Numerically compare two last-token-logit files produced by the DSV4 probe."""

from __future__ import annotations

import argparse
import json
from pathlib import Path

import torch


def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument('reference', type=Path)
parser.add_argument('candidate', type=Path)
parser.add_argument('--rtol', type=float, default=1e-2)
parser.add_argument('--atol', type=float, default=2e-2)
parser.add_argument('--output', type=Path)
return parser.parse_args()


def load_logits(path: Path) -> tuple[str, list[int], torch.Tensor]:
payload = torch.load(path, map_location='cpu', weights_only=True)
return str(payload.get('mode', path.stem)), list(payload['input_ids']), payload['last_logits'].float()


def main() -> None:
args = parse_args()
reference_mode, reference_ids, reference = load_logits(args.reference)
candidate_mode, candidate_ids, candidate = load_logits(args.candidate)
if reference_ids != candidate_ids:
raise RuntimeError(f'Input IDs differ: {reference_ids} != {candidate_ids}')
if reference.shape != candidate.shape:
raise RuntimeError(f'Logit shapes differ: {tuple(reference.shape)} != {tuple(candidate.shape)}')

difference = (candidate - reference).abs()
close = torch.isclose(candidate, reference, rtol=args.rtol, atol=args.atol)
reference_top = torch.topk(reference, k=min(20, reference.numel())).indices.tolist()
candidate_top = torch.topk(candidate, k=min(20, candidate.numel())).indices.tolist()
report = {
'reference': str(args.reference),
'reference_mode': reference_mode,
'candidate': str(args.candidate),
'candidate_mode': candidate_mode,
'input_ids': reference_ids,
'shape': list(reference.shape),
'rtol': args.rtol,
'atol': args.atol,
'allclose': bool(close.all().item()),
'close_fraction': close.float().mean().item(),
'max_abs_diff': difference.max().item(),
'mean_abs_diff': difference.mean().item(),
'reference_top20': reference_top,
'candidate_top20': candidate_top,
'top20_overlap': len(set(reference_top) & set(candidate_top)),
}

output = args.output or args.candidate.with_name(
f'compare_{reference_mode}_vs_{candidate_mode}.json')
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding='utf-8')
print(json.dumps(report, ensure_ascii=False, indent=2))
print(f'Report saved to: {output.resolve()}')
if not report['allclose']:
raise SystemExit(1)


if __name__ == '__main__':
main()
Original file line number Diff line number Diff line change
@@ -0,0 +1,146 @@
#!/usr/bin/env python3
"""Save last-token logits from a running four-layer Twinkle server."""

from __future__ import annotations

import argparse
import hashlib
import json
from pathlib import Path
from typing import Any

import torch
from peft import LoraConfig

from twinkle_client import init_twinkle_client
from twinkle_client.model import MultiLoraTransformersModel


DEFAULT_INPUT_IDS = [0, 128803, 2788, 6573, 70979, 36005, 320, 128804, 128821]
TARGET_PARAMETERS = ['mlp.experts.gate_up_proj', 'mlp.experts.down_proj']


def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument('--mode', required=True, choices=('no_ep', 'ep_loop', 'ep_gmm'))
parser.add_argument('--server-url', default='http://127.0.0.1:8000')
parser.add_argument('--server-token', default='EMPTY_TOKEN')
parser.add_argument('--served-model', default='deepseek-v4-0731-local')
parser.add_argument('--output-dir', type=Path, default=Path('output/dsv4_ep_diag'))
parser.add_argument('--input-ids', default=','.join(str(item) for item in DEFAULT_INPUT_IDS))
return parser.parse_args()


def _extract_logits(result: Any) -> torch.Tensor:
if hasattr(result, 'model_dump'):
result = result.model_dump()
if isinstance(result, list) and len(result) == 1 and isinstance(result[0], dict):
result = result[0]
if not isinstance(result, dict) or result.get('logits') is None:
raise RuntimeError(f'forward_only did not return logits; result type={type(result).__name__}')

logits = torch.as_tensor(result['logits'], dtype=torch.float32)
original_shape = tuple(logits.shape)
while logits.ndim > 3:
logits = logits[0]
if logits.ndim == 3:
logits = logits[0, -1]
elif logits.ndim == 2:
logits = logits[-1]
elif logits.ndim != 1:
raise RuntimeError(f'Unsupported logits shape: {original_shape}')
if logits.numel() < 1000:
raise RuntimeError(f'Last-token logits look too small: shape={tuple(logits.shape)}')
return logits.contiguous()


def _tensor_sha256(tensor: torch.Tensor) -> str:
values = tensor.detach().cpu().contiguous().numpy().tobytes()
return hashlib.sha256(values).hexdigest()


def main() -> None:
args = parse_args()
input_ids = [int(item.strip()) for item in args.input_ids.split(',') if item.strip()]
if not input_ids:
raise SystemExit('--input-ids must contain at least one token ID')

client = init_twinkle_client(
base_url=args.server_url,
api_key=args.server_token,
session_heartbeat_interval=10,
)
try:
capacity = client.get_capacity_info()
if capacity.free_loras < 1:
raise RuntimeError('No free LoRA slot. Restart the diagnostic server before running the probe.')

model = MultiLoraTransformersModel(model_id=args.served_model)
model.add_adapter_to_model(
f'dsv4_diag_{args.mode}',
LoraConfig(
r=8,
lora_alpha=32,
lora_dropout=0.0,
target_modules=None,
target_parameters=TARGET_PARAMETERS,
bias='none',
),
gradient_accumulation_steps=1,
)
model.set_processor('InputProcessor', padding_side='left', padding_free=False)

raw_input = {
'input_ids': input_ids,
'attention_mask': [1] * len(input_ids),
'position_ids': list(range(len(input_ids))),
}
response = model.forward_only(
inputs=[raw_input],
disable_lora=True,
return_logits=True,
)
last_logits = _extract_logits(response.result)
finally:
client.close()

finite = torch.isfinite(last_logits)
top_values, top_indices = torch.topk(last_logits, k=min(20, last_logits.numel()))
report = {
'mode': args.mode,
'server_url': args.server_url,
'served_model': args.served_model,
'input_ids': input_ids,
'last_logits_shape': list(last_logits.shape),
'dtype_saved': str(last_logits.dtype),
'sha256': _tensor_sha256(last_logits),
'finite': bool(finite.all().item()),
'nan_count': int(torch.isnan(last_logits).sum().item()),
'inf_count': int(torch.isinf(last_logits).sum().item()),
'sum': last_logits.sum().item(),
'abs_sum': last_logits.abs().sum().item(),
'min': last_logits.min().item(),
'max': last_logits.max().item(),
'top_token_ids': top_indices.tolist(),
'top_logits': top_values.tolist(),
}

args.output_dir.mkdir(parents=True, exist_ok=True)
tensor_path = args.output_dir / f'{args.mode}_last_logits.pt'
json_path = args.output_dir / f'{args.mode}_last_logits.json'
torch.save(
{
'mode': args.mode,
'input_ids': input_ids,
'last_logits': last_logits,
},
tensor_path,
)
json_path.write_text(json.dumps(report, ensure_ascii=False, indent=2), encoding='utf-8')
print(json.dumps(report, ensure_ascii=False, indent=2))
print(f'Logits saved to: {tensor_path.resolve()}')
print(f'Report saved to: {json_path.resolve()}')


if __name__ == '__main__':
main()
Original file line number Diff line number Diff line change
@@ -0,0 +1,131 @@
#!/usr/bin/env python3
"""Compare the DeepSeek-V4 square expert layout on NPU GMM against F.linear."""

from __future__ import annotations

import argparse
import json
from pathlib import Path

import torch
import torch.nn.functional as F
from torch import nn

from twinkle.kernel.ops.moe.npu import GmmFunction, _normalize_packed_expert_weights


class PackedExperts(nn.Module):

def __init__(self, gate_up_proj: torch.Tensor, down_proj: torch.Tensor):
super().__init__()
self.gate_up_proj = nn.Parameter(gate_up_proj, requires_grad=False)
self.down_proj = nn.Parameter(down_proj, requires_grad=False)


def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument('--device', default='npu:0')
parser.add_argument('--dtype', choices=('float16', 'bfloat16'), default='bfloat16')
parser.add_argument('--atol', type=float, default=2e-2)
parser.add_argument('--rtol', type=float, default=1e-2)
parser.add_argument('--output', type=Path, default=Path('output/dsv4_ep_diag/npu_gmm_layout.json'))
return parser.parse_args()


def reference_forward(
inputs: torch.Tensor,
counts: list[int],
gate_up_proj: torch.Tensor,
down_proj: torch.Tensor,
) -> torch.Tensor:
outputs = []
start = 0
for expert, count in enumerate(counts):
expert_input = inputs[start:start + count]
gate_up = F.linear(expert_input, gate_up_proj[expert])
gate, up = gate_up.chunk(2, dim=-1)
outputs.append(F.linear(F.silu(gate) * up, down_proj[expert]))
start += count
return torch.cat(outputs, dim=0)


def gmm_forward(
inputs: torch.Tensor,
counts: torch.Tensor,
gate_up_weight: torch.Tensor,
down_weight: torch.Tensor,
) -> torch.Tensor:
import torch_npu

gate_up = GmmFunction.apply(inputs, counts, gate_up_weight)
activated = torch_npu.npu_swiglu(gate_up, dim=-1)
return GmmFunction.apply(activated, counts, down_weight)


def main() -> None:
args = parse_args()
try:
import torch_npu # noqa: F401
except ImportError as exc:
raise SystemExit('torch_npu is required; run this script in the Ascend container.') from exc

if not torch.npu.is_available():
raise SystemExit('torch.npu.is_available() is False')

device = torch.device(args.device)
dtype = getattr(torch, args.dtype)
torch.npu.set_device(device.index or 0)
torch.manual_seed(20260901)
torch.npu.manual_seed_all(20260901)

# Preserve the DeepSeek-V4 relation hidden == 2 * intermediate. The gate/up
# matrix is deliberately non-symmetric so an omitted transpose is visible.
experts = 2
hidden = 64
intermediate = 32
token_counts = [8, 8]
inputs = torch.randn(sum(token_counts), hidden, device=device, dtype=dtype) * 0.1
gate_up_proj = torch.randn(experts, 2 * intermediate, hidden, device=device, dtype=dtype) * 0.02
down_proj = torch.randn(experts, hidden, intermediate, device=device, dtype=dtype) * 0.02
module = PackedExperts(gate_up_proj, down_proj).to(device)

normalized_gate_up, normalized_down = _normalize_packed_expert_weights(module, dtype, hidden)
counts = torch.tensor(token_counts, device=device, dtype=torch.int64)

with torch.no_grad():
expected = reference_forward(inputs, token_counts, gate_up_proj, down_proj)
actual = gmm_forward(inputs, counts, normalized_gate_up, normalized_down)
# Reproduce the old DeepSeek-V4 bug: the square gate/up tensor was not transposed.
old_bug = gmm_forward(inputs, counts, gate_up_proj, down_proj.transpose(1, 2))
torch.npu.synchronize()

difference = (actual.float() - expected.float()).abs()
old_difference = (old_bug.float() - expected.float()).abs()
passed = torch.allclose(actual.float(), expected.float(), rtol=args.rtol, atol=args.atol)
report = {
'device': str(device),
'dtype': str(dtype),
'input_shape': list(inputs.shape),
'gate_up_shape_transformers': list(gate_up_proj.shape),
'down_shape_transformers': list(down_proj.shape),
'gate_up_shape_gmm': list(normalized_gate_up.shape),
'down_shape_gmm': list(normalized_down.shape),
'rtol': args.rtol,
'atol': args.atol,
'max_abs_diff': difference.max().item(),
'mean_abs_diff': difference.mean().item(),
'old_bug_max_abs_diff': old_difference.max().item(),
'old_bug_mean_abs_diff': old_difference.mean().item(),
'passed': bool(passed),
}

args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(report, indent=2), encoding='utf-8')
print(json.dumps(report, indent=2))
print(f'Report saved to: {args.output.resolve()}')
if not passed:
raise SystemExit(1)


if __name__ == '__main__':
main()
Loading
Loading