Views
No views yet
1import torch
2from transformers import AutoModelForCausalLM, AutoTokenizer
3
4repo = "CodeSoft/MetaDiffusion-600M-ChatBase"
5
6m = AutoModelForCausalLM.from_pretrained(
7 repo,
8 trust_remote_code=True,
9 dtype=torch.bfloat16,
10).to("cuda")
11
12tok = AutoTokenizer.from_pretrained(
13 repo,
14 subfolder="tokenizer",
15 trust_remote_code=True,
16)
17
18prompt = tok.apply_chat_template(
19 [{"role": "user", "content": "hi"}],
20 tokenize=False,
21 add_generation_prompt=True,
22)
23
24inputs = tok(prompt, return_tensors="pt").to("cuda")
25
26with torch.inference_mode():
27 out = m.generate(
28 **inputs,
29 max_new_tokens=100,
30 )
31
32print(tok.decode(out[0], skip_special_tokens=True))
331python chat.py \
2 --model-path model.safetensors \
3 --tokenizer ./tokenizer \
4 --im-end-bias 2.0 --im-end-bias-t 0.3 --watch1# 1. Init: convert the AR model to a diffusion init
2python convert.py --source Qwen/Qwen3-0.6B \
3 --output init/metadiffusion-600M-instruct.pt \
4 --tokenizer-out data/tokenizer
5
6# 2. Corpus: smol, opc, math and no_robots, or a local --jsonl of {"messages": [...]} rows.
7# --val-fraction holds out a disjoint val set for early stopping.
8python prepare_data.py --datasets smol,math --out data \
9 --val-fraction 0.05
10
11# 3. Train (defaults: lr 5e-5, bf16, seq 512, batch auto-detected)
12python train.py --init-checkpoint init/metadiffusion-600M-instruct.pt \
13 --data-dir data --output-dir checkpoints --max-steps 30000
14
15# 4. Continue a run: checkpoints carry model + optimizer + scheduler
16# state, so --resume-from picks up LR position and momentum exactly
17python train.py --init-checkpoint init/metadiffusion-600M-instruct.pt \
18 --data-dir data --output-dir checkpoints \
19 --resume-from checkpoints_p2/step_20000.pt --max-steps 16000
20
21# 5. Test, then ship
22python chat.py --model-path checkpoints_/step_30000.pt \
23 --tokenizer data/tokenizer --watch
24python export_hf.py --checkpoint checkpoints/step_30000.pt \
25 --tokenizer data/tokenizer --output MetaDiffusion-600M-ChatBase