| |
| """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()) |
|
|