hku_diffusion_dllm / experiments /dllm_timestep_alignment /compute_alignment_streaming.py
Ouzhang's picture
Add files using upload-large-folder tool
d91766b verified
Raw
History Blame Contribute Delete
9.54 kB
#!/usr/bin/env python3
"""Streaming adjacent-step dLLM mask-ratio gradient alignment.
Use this for very fine ratio grids such as 0.001. It avoids storing hundreds or
thousands of full gradient vectors by computing adjacent cosine online and only
keeping a configurable downsampled subset for a small pairwise heatmap.
"""
from __future__ import annotations
import argparse
import json
import sys
from datetime import datetime, timezone
from pathlib import Path
import torch
from compute_alignment import (
cosine_matrix,
deterministic_mask,
ensure_dir,
flatten_grads,
load_text_samples,
parse_ratios,
select_dflash_params,
write_csv,
write_report,
write_svg,
)
def grad_for_ratio(
*,
ratio: float,
ratio_idx: int,
args: argparse.Namespace,
samples,
tokenizer,
target,
draft,
selected_params,
mask_id: int,
device: torch.device,
) -> tuple[torch.Tensor, float]:
from dflash.model import extract_context_feature
import torch.nn.functional as F
draft.zero_grad(set_to_none=True)
total_loss = 0.0
for sample_idx, sample in enumerate(samples):
ids = torch.tensor(sample.input_ids, device=device, dtype=torch.long).unsqueeze(0)
prefix = ids[:, : args.prefix_tokens]
clean_block = ids[:, args.prefix_tokens : args.prefix_tokens + args.block_size]
mask = deterministic_mask(args.block_size, ratio, args.seed + ratio_idx * 1009 + sample_idx, device)
noisy_block = clean_block.clone()
noisy_block[:, mask] = mask_id
with torch.no_grad():
target_out = target(prefix, output_hidden_states=True, use_cache=False)
target_hidden = extract_context_feature(target_out.hidden_states, draft.target_layer_ids)
noise_embedding = target.model.embed_tokens(noisy_block)
position_ids = torch.arange(args.prefix_tokens + args.block_size, device=device).unsqueeze(0)
hidden = draft(
target_hidden=target_hidden,
noise_embedding=noise_embedding,
position_ids=position_ids,
use_cache=False,
is_causal=False,
)
logits = target.lm_head(hidden)
loss = F.cross_entropy(
logits[:, mask, :].float().reshape(-1, logits.shape[-1]),
clean_block[:, mask].reshape(-1),
)
(loss / len(samples)).backward()
total_loss += float(loss.detach().item())
del ids, prefix, clean_block, noisy_block, target_out, target_hidden, noise_embedding, hidden, logits, loss
vec = flatten_grads(selected_params, device=torch.device("cpu"))
if torch.cuda.is_available():
torch.cuda.empty_cache()
return vec, total_loss / len(samples)
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--output-dir", type=Path, required=True)
parser.add_argument("--ratios", default="0.001:0.001:0.999")
parser.add_argument("--num-samples", type=int, default=8)
parser.add_argument("--block-size", type=int, default=64)
parser.add_argument("--seed", type=int, default=1234)
parser.add_argument("--device", default="cuda")
parser.add_argument("--model", default="/home/l/liyj/shiying/hku_diffusion_dllm/models/qwen3-4b-target")
parser.add_argument("--draft-model", default="/home/l/liyj/shiying/hku_diffusion_dllm/models/qwen3-4b-dflash-b16")
parser.add_argument("--dataset", choices=["gsm8k", "math500", "mbpp"], default="gsm8k")
parser.add_argument("--prefix-tokens", type=int, default=128)
parser.add_argument("--grad-scope", choices=["fc", "last_layer", "norms", "all_draft"], default="fc")
parser.add_argument("--dtype", choices=["bf16", "fp16", "fp32"], default="bf16")
parser.add_argument("--attn-implementation", default="sdpa")
parser.add_argument("--downsample-every", type=int, default=10)
args = parser.parse_args()
from transformers import AutoModelForCausalLM, AutoTokenizer
from dflash.model import DFlashDraftModel
ratios = parse_ratios(args.ratios)
out_dir = args.output_dir
ensure_dir(out_dir)
device = torch.device(args.device)
dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16 if args.dtype == "fp16" else torch.float32
torch.manual_seed(args.seed)
tokenizer = AutoTokenizer.from_pretrained(args.model)
target = AutoModelForCausalLM.from_pretrained(
args.model,
attn_implementation=args.attn_implementation,
dtype=dtype,
).to(device).eval()
draft = DFlashDraftModel.from_pretrained(
args.draft_model,
attn_implementation=args.attn_implementation,
dtype=dtype,
).to(device).train()
for param in target.parameters():
param.requires_grad_(False)
selected_params = select_dflash_params(draft, args.grad_scope)
mask_id = draft.mask_token_id
if mask_id is None:
raise RuntimeError("DFlash draft model did not expose mask_token_id")
samples = load_text_samples(args, tokenizer)
adjacent = []
losses_by_ratio = []
downsample_ratios = []
downsample_vectors = []
prev_ratio = None
prev_vec = None
for idx, ratio in enumerate(ratios):
vec, avg_loss = grad_for_ratio(
ratio=ratio,
ratio_idx=idx,
args=args,
samples=samples,
tokenizer=tokenizer,
target=target,
draft=draft,
selected_params=selected_params,
mask_id=mask_id,
device=device,
)
losses_by_ratio.append(avg_loss)
if prev_vec is not None:
sim = cosine_matrix([prev_vec, vec])[0][1]
adjacent.append({"left": prev_ratio, "right": ratio, "mid": (prev_ratio + ratio) / 2, "cosine": sim})
if idx % args.downsample_every == 0 or idx == len(ratios) - 1:
downsample_ratios.append(ratio)
downsample_vectors.append(vec)
prev_ratio = ratio
prev_vec = vec
downsample_matrix = cosine_matrix(downsample_vectors)
selected_param_count = sum(p.numel() for p in selected_params)
candidate = min(adjacent, key=lambda x: x["cosine"]) if adjacent else None
payload = {
"updated_at": datetime.now(timezone.utc).isoformat(),
"command": " ".join(sys.argv),
"mode": "dflash_streaming_adjacent",
"model": args.model,
"draft_model": args.draft_model,
"dataset": args.dataset,
"num_samples": args.num_samples,
"block_size": args.block_size,
"prefix_tokens": args.prefix_tokens,
"grad_scope": args.grad_scope,
"selected_param_count": selected_param_count,
"ratios": ratios,
"loss_by_ratio": losses_by_ratio,
"adjacent_cosines": adjacent,
"candidate_boundary": None if candidate is None else f"{candidate['left']:g} -> {candidate['right']:g}",
"downsample_every": args.downsample_every,
"downsampled_ratios": downsample_ratios,
"downsampled_cosine_similarity": downsample_matrix,
"output_dir": str(out_dir),
}
(out_dir / "adjacent_alignment.json").write_text(json.dumps(payload, indent=2, ensure_ascii=False) + "\n", encoding="utf-8")
with (out_dir / "adjacent_alignment.csv").open("w", encoding="utf-8") as f:
f.write("left,right,mid,cosine\n")
for row in adjacent:
f.write(f"{row['left']},{row['right']},{row['mid']},{row['cosine']}\n")
heat_payload = {
"mode": payload["mode"],
"dataset": payload["dataset"],
"num_samples": payload["num_samples"],
"block_size": payload["block_size"],
"grad_scope": payload["grad_scope"],
"ratios": downsample_ratios,
"cosine_similarity": downsample_matrix,
"output_dir": str(out_dir),
"updated_at": payload["updated_at"],
}
(out_dir / "downsampled_alignment.json").write_text(json.dumps(heat_payload, indent=2, ensure_ascii=False) + "\n", encoding="utf-8")
write_csv(out_dir / "downsampled_alignment.csv", downsample_ratios, downsample_matrix)
write_svg(out_dir / "downsampled_heatmap.svg", downsample_ratios, downsample_matrix, f"downsampled {args.block_size} adjacent run")
write_report(out_dir / "downsampled_report.md", heat_payload)
lines = [
"# Streaming Adjacent dLLM Alignment",
"",
f"- Updated: `{payload['updated_at']}`",
f"- Dataset: `{args.dataset}`",
f"- Samples: `{args.num_samples}`",
f"- Block size: `{args.block_size}`",
f"- Ratio grid: `{args.ratios}`",
f"- Gradient scope: `{args.grad_scope}`",
f"- Candidate boundary: `{payload['candidate_boundary']}`",
"",
"## Adjacent Cosines",
"",
"| Step pair | Cosine similarity |",
"| --- | ---: |",
]
for row in adjacent:
lines.append(f"| {row['left']:g} -> {row['right']:g} | {row['cosine']:.3f} |")
(out_dir / "adjacent_report.md").write_text("\n".join(lines) + "\n", encoding="utf-8")
print(json.dumps({
"status": "ok",
"adjacent_json": str(out_dir / "adjacent_alignment.json"),
"adjacent_report": str(out_dir / "adjacent_report.md"),
"downsampled_json": str(out_dir / "downsampled_alignment.json"),
"downsampled_heatmap": str(out_dir / "downsampled_heatmap.svg"),
"candidate_boundary": payload["candidate_boundary"],
}, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())