Text Generation
PEFT
Safetensors
lora
trl
grpo
gdpo
dpo
divpo
rlhf
diversity
creative-writing
mode-collapse
Instructions to use Mercity/creative-writing-llm with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use Mercity/creative-writing-llm with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
| """ | |
| Build the scored generation pool that feeds E3 (multi-positive diverse DPO) and | |
| E4 (faithful DivPO), and doubles as base-model analysis data. | |
| Per prompt: N=16 samples at temperature 1.0 from the BASE policy, each carrying | |
| - text, token count, cumulative + length-normalized logprob (E4 divpo-prob) | |
| - programmatic gate result | |
| - judge quality / novelty (gate-passing stories only) | |
| - embedding, per-group deviation d_i and marginal contribution m_i | |
| Both DPO arms consume this identical artifact, which is the point: E4-vs-E3 is | |
| then a comparison of PAIR CONSTRUCTION and LOSS, with the data held fixed. | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import sys | |
| import time | |
| from pathlib import Path | |
| import numpy as np | |
| ROOT = Path(__file__).resolve().parent.parent | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--model", default="Qwen/Qwen3-4B-Instruct-2507") | |
| ap.add_argument("--tag", default="4b") | |
| ap.add_argument("--n", type=int, default=16) | |
| ap.add_argument("--temperature", type=float, default=1.0) | |
| ap.add_argument("--top-p", type=float, default=1.0) | |
| ap.add_argument("--limit", type=int, default=None, help="prompt subset, for smoke") | |
| ap.add_argument("--split", default="train") | |
| ap.add_argument("--seed", type=int, default=1234) | |
| ap.add_argument("--gpu-mem", type=float, default=0.85) | |
| args = ap.parse_args() | |
| from transformers import AutoTokenizer | |
| import gates | |
| import logbook | |
| from data import load_prompts | |
| from diversity import (l2_normalize, logdet_volume, marginal_contributions, | |
| pairwise_deviation) | |
| from generate import build_llm, generate | |
| from judge import build_judge | |
| out_dir = ROOT / "outputs" / f"pool_{args.tag}" | |
| out_dir.mkdir(parents=True, exist_ok=True) | |
| pool_path = out_dir / f"pool_{args.split}.jsonl" | |
| prompts = load_prompts(args.split, ROOT / "data") | |
| if args.limit: | |
| prompts = prompts[: args.limit] | |
| print(f"[pool] {len(prompts)} prompts x N={args.n} = {len(prompts)*args.n} stories") | |
| # ---- 1. generate ----------------------------------------------------- | |
| t0 = time.time() | |
| tok = AutoTokenizer.from_pretrained(args.model) | |
| llm = build_llm(args.model, gpu_mem_util=args.gpu_mem, seed=args.seed) | |
| gens = generate(llm, tok, prompts, n=args.n, temperature=args.temperature, | |
| top_p=args.top_p, seed=args.seed) | |
| t_gen = time.time() - t0 | |
| ntok = sum(g.n_tokens for g in gens) | |
| print(f"[gen] {len(gens)} stories, {ntok} tok in {t_gen/60:.1f} min " | |
| f"({ntok/t_gen:.0f} tok/s)") | |
| # free the GPU before loading the embedder | |
| del llm | |
| import gc, torch | |
| gc.collect(); torch.cuda.empty_cache() | |
| # ---- 2. gates -------------------------------------------------------- | |
| grs = [gates.check(g.text, finish_reason=g.finish_reason) for g in gens] | |
| n_pass = sum(r.passed for r in grs) | |
| print(f"[gates] pass {n_pass}/{len(grs)} ({100*n_pass/len(grs):.1f}%)") | |
| # ---- 3. judge (gate-passers only) ------------------------------------ | |
| judge = build_judge(cache_path=str(ROOT / "cache" / "judge.sqlite"), | |
| concurrency=24) | |
| idx = [i for i in range(len(gens)) if grs[i].passed] | |
| t1 = time.time() | |
| scores = judge.score_many_sync([(gens[i].prompt, gens[i].text) for i in idx]) | |
| print(f"[judge] {len(idx)} scored in {(time.time()-t1)/60:.1f} min | " | |
| f"health={judge.health()}") | |
| judge.assert_healthy() | |
| quality = np.zeros(len(gens)); novelty = np.zeros(len(gens)) | |
| for i, s in zip(idx, scores): | |
| quality[i] = s.quality; novelty[i] = s.novelty | |
| # ---- 4. embeddings + per-group diversity ----------------------------- | |
| from sentence_transformers import SentenceTransformer | |
| enc = SentenceTransformer("BAAI/bge-base-en-v1.5", device="cuda") | |
| t2 = time.time() | |
| E = enc.encode([g.text for g in gens], normalize_embeddings=True, | |
| batch_size=64, show_progress_bar=False, convert_to_numpy=True) | |
| E = l2_normalize(np.asarray(E, dtype=np.float64)) | |
| print(f"[embed] {E.shape} in {time.time()-t2:.0f}s") | |
| by_prompt: dict[str, list[int]] = {} | |
| for i, g in enumerate(gens): | |
| by_prompt.setdefault(g.prompt_id, []).append(i) | |
| dev = np.zeros(len(gens)); marg = np.zeros(len(gens)) | |
| group_logdet: dict[str, float] = {} | |
| for pid, ids in by_prompt.items(): | |
| sub = E[ids] | |
| d = pairwise_deviation(sub); m = marginal_contributions(sub) | |
| for k, i in enumerate(ids): | |
| dev[i] = d[k]; marg[i] = m[k] | |
| group_logdet[pid] = logdet_volume(sub) | |
| # ---- 5. write -------------------------------------------------------- | |
| with open(pool_path, "w") as f: | |
| for i, g in enumerate(gens): | |
| f.write(json.dumps({ | |
| **g.as_dict(), | |
| "gate_passed": bool(grs[i].passed), | |
| "gate_reasons": grs[i].reasons, | |
| "ends_cleanly": grs[i].completeness, | |
| "n_words": grs[i].n_words, | |
| "quality": float(quality[i]), | |
| "novelty": float(novelty[i]), | |
| "deviation": float(dev[i]), | |
| "marginal": float(marg[i]), | |
| "group_logdet": float(group_logdet[g.prompt_id]), | |
| }) + "\n") | |
| np.save(out_dir / f"emb_{args.split}.npy", E.astype(np.float32)) | |
| # ---- 6. summary ------------------------------------------------------ | |
| q_pass = quality[[i for i in idx]] | |
| summary = { | |
| "tag": args.tag, "model": args.model, "split": args.split, | |
| "n_prompts": len(prompts), "n_per_prompt": args.n, "n_stories": len(gens), | |
| "temperature": args.temperature, "seed": args.seed, | |
| "gate_pass_rate": float(n_pass / len(grs)), | |
| "ends_cleanly_rate": float(np.mean([r.completeness for r in grs])), | |
| "median_words": float(np.median([r.n_words for r in grs])), | |
| "quality_mean": float(q_pass.mean()) if len(q_pass) else 0.0, | |
| "quality_sd": float(q_pass.std()) if len(q_pass) else 0.0, | |
| "novelty_mean": float(novelty[idx].mean()) if len(idx) else 0.0, | |
| "deviation_mean": float(dev.mean()), | |
| "logdet_mean": float(np.mean(list(group_logdet.values()))), | |
| "gen_minutes": t_gen / 60, | |
| "judge_cost": judge.cost_estimate(0.140, 0.280), | |
| "judge_health": judge.health(), | |
| } | |
| json.dump(summary, open(out_dir / f"summary_{args.split}.json", "w"), indent=2) | |
| print("\n" + json.dumps(summary, indent=1)) | |
| logbook.note(f"pool built: {args.tag}/{args.split}", | |
| f"```json\n{json.dumps(summary, indent=1)}\n```") | |
| logbook.checkpoint(f"pool_{args.tag}_{args.split}") | |
| return 0 | |
| if __name__ == "__main__": | |
| sys.exit(main()) | |