Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
041ad2d
feat: add omni rollout contract and audio encode for sglang-omni RL
Hayden727 Jun 19, 2026
db08da2
fix(omni): harden OmniGenerateFn budget/audio handling + audio_data w…
Hayden727 Jun 19, 2026
e183512
fix(omni): partial-rollout loss_mask, call_processor audio route, uni…
Hayden727 Jun 19, 2026
f450d5b
feat(omni): thinker-submodule checkpoint extractor for FSDP training
Hayden727 Jun 19, 2026
0d82cfc
feat(omni): text-only math reward for GATE-A thinker RL smoke
Hayden727 Jun 19, 2026
1230255
feat(omni): GATE-A closed-loop GRPO LoRA smoke + math dataset
Hayden727 Jun 19, 2026
6e6626d
feat(omni): full on-policy GATE-A with NCCL weight-sync to served thi…
Hayden727 Jun 19, 2026
2a72ca2
fix(omni): bind CUDA device before NCCL collectives in full GATE-A
Hayden727 Jun 19, 2026
9a7f7bf
feat(omni): make weight-sync GROUP_NAME env-configurable
Hayden727 Jun 19, 2026
059c28d
docs(omni): record NCCL_P2P_DISABLE fix that makes GATE-A weight-sync…
Hayden727 Jun 20, 2026
76d75bc
feat(omni): composite TTS reward (ASR CER + audio guards) for GATE-B
Hayden727 Jun 20, 2026
074b5e7
feat(omni): GATE-B rollout->reward->advantage loop on real Higgs TTS
Hayden727 Jun 22, 2026
285340b
feat(omni): trainer-side Higgs TTS actor + logprob-parity gate
Hayden727 Jun 25, 2026
edf87ea
feat(omni): GATE-B full on-policy loop (GRPO + NCCL tts_engine weight…
Hayden727 Jun 25, 2026
e5ae79a
feat(omni): parametrize GATE-B rollout temperature/max_new_tokens
Hayden727 Jun 25, 2026
1d3a32a
chore(omni): add corrected higgs codec-logprob patch (supersedes brok…
Hayden727 Jun 25, 2026
7a358d5
chore(omni): remove stale broken higgs codec-logprob patch
Hayden727 Jun 26, 2026
8b0521c
feat(omni): non-saturating eval sets + MAX_NEW knob for GATE-A/B lear…
Hayden727 Jun 30, 2026
69764df
refactor(omni): rename gate_a/gate_b examples to function-descriptive…
Hayden727 Jun 30, 2026
dd5d2ad
refactor(omni): scrub GATE-A/GATE-B terminology + fix internal path refs
Hayden727 Jun 30, 2026
d55adfc
feat(rollout): wire --rollout-external to use an external sglang server
Hayden727 Jun 30, 2026
3186a5a
refactor(omni): consume Higgs rollout contract
Hayden727 Jul 8, 2026
aae8700
fix(omni): preserve external rollout metadata paths
Hayden727 Jul 8, 2026
7b8fb14
fix(omni): train all Higgs codebook actions
Hayden727 Jul 11, 2026
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
106 changes: 106 additions & 0 deletions examples/higgs_tts_rl/logprob_parity_probe.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,106 @@
"""Logprob-parity check for the Higgs TTS trainable actor.

Right after load the trainer-side actor and the served model are the same policy,
so the actor's recomputed codebook-0 log-probs must match the rollout's
`output_token_logprobs`. This is the make-or-break correctness check before any
GRPO update.

Run (container, miles venv; Higgs server serving on SERVER):
SERVER=http://localhost:8010 HIGGS_CKPT='<snapshot glob>' CUDA_VISIBLE_DEVICES=4 \
PYTHONPATH=/root/rl-omni/sglang-omni:/root/rl-omni/miles \
python examples/higgs_tts_rl/logprob_parity_probe.py
"""

from __future__ import annotations

import glob
import json
import os
import urllib.request

from miles_plugins.omni.rollout_contract import (
build_generate_payload,
parse_generate_response,
)

SERVER = os.environ.get("SERVER", "http://localhost:8010")
# Gate on mean|Δ|: the residual is the served model's bf16 + sglang-kernel numeric
# floor (an fp32 trainer gives the SAME ~0.05 residual), so per-token max|Δ| of ~0.2
# is irreducible cross-implementation noise, not a reconstruction error. exp(0.2)≈1.22
# sits at the GRPO clip boundary and only biases the first ratio after each sync.
TOL = float(os.environ.get("PARITY_TOL", "0.10"))


def _rollout(prompt_ids: list[int], seed: int) -> dict:
req = build_generate_payload(
prompt_ids,
{"temperature": 0.8, "top_p": 0.95, "max_new_tokens": 256, "seed": seed},
output_modalities=["audio"],
)
resp = json.loads(
urllib.request.urlopen(
urllib.request.Request(
SERVER + "/generate",
data=json.dumps(req).encode(),
headers={"Content-Type": "application/json"},
),
timeout=180,
).read()
)
result = parse_generate_response(resp)
return {
"old_logprobs": result.response_log_probs,
"cb0_tokens": result.response_tokens,
"codebook_tokens": result.output_codebook_tokens,
}


def main() -> None:
import torch
from tokenizers import Tokenizer
from transformers import PreTrainedTokenizerFast

from sglang_omni.models.higgs_tts.text_tokenizer import HiggsTokenizerAdapter

from miles_plugins.omni.higgs_actor import HiggsTtsActor

ckpt = glob.glob(os.environ["HIGGS_CKPT"])[0] if "*" in os.environ["HIGGS_CKPT"] else os.environ["HIGGS_CKPT"]
tok = PreTrainedTokenizerFast(tokenizer_object=Tokenizer.from_file(os.path.join(ckpt, "tokenizer.json")))
adapter = HiggsTokenizerAdapter(tok)
device = os.environ.get("ACTOR_DEVICE", "cuda:0")

dtype = torch.float32 if os.environ.get("ACTOR_DTYPE") == "fp32" else torch.bfloat16
actor = HiggsTtsActor(ckpt, device=device, dtype=dtype)
print("actor loaded; backbone dtype", actor.dtype)

texts = ["Hello world.", "The quick brown fox."]
worst = 0.0
worst_mean = 0.0
for i, text in enumerate(texts):
pid = list(map(int, adapter.build_prompt(text, num_ref_tokens=0)))
r = _rollout(pid, seed=1000 + i)
codes = r["codebook_tokens"]
old = r["old_logprobs"]
if not codes or not old:
print(f"[{text!r}] no codes/logprobs returned -> SKIP (server missing Step-1 fix?)")
continue
assert all(row[0] == t for row, t in zip(codes, r["cb0_tokens"])), "cb0 mismatch codes vs logprob tokens"

with torch.no_grad():
new = actor.codebook0_logprobs(pid, codes).tolist()
n = min(len(new), len(old))
diffs = [abs(new[j] - old[j]) for j in range(n)]
max_d = max(diffs)
mean_d = sum(diffs) / n
worst = max(worst, max_d)
worst_mean = max(worst_mean, mean_d)
print(f"[{text!r}] T={n} max|Δ|={max_d:.4f} mean|Δ|={mean_d:.4f}")
print(f" old[:5]={[round(x,3) for x in old[:5]]}")
print(f" new[:5]={[round(x,3) for x in new[:5]]}")

print(f"WORST_MEAN_ABS_DIFF: {worst_mean:.4f} (tol={TOL}) WORST_MAX_ABS_DIFF: {worst:.4f}")
print(f"PARITY_OK: {worst_mean < TOL}")


if __name__ == "__main__":
main()
236 changes: 236 additions & 0 deletions examples/higgs_tts_rl/onpolicy_grpo_weight_sync.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,236 @@
"""On-policy Higgs TTS GRPO with per-step SGLang-Omni weight sync.

Set ``TRAIN_MODE=lora`` (default) for the low-memory smoke path or ``full`` to
train and sync the complete backbone plus tied codebook embedding/head.
"""

from __future__ import annotations

import glob
import json
import os
import threading
import urllib.request

import torch
from peft import LoraConfig, get_peft_model

from miles_plugins.omni.rollout_contract import (
build_generate_payload,
parse_generate_response,
parse_omni_action_stream,
)

SERVER = os.environ.get("SERVER", "http://localhost:8010")
HIGGS_CKPT = os.environ["HIGGS_CKPT"]
DATA = os.environ.get("DATA", "examples/higgs_tts_rl/tts_smoke.jsonl")
STEPS = int(os.environ.get("STEPS", "3"))
GROUP = int(os.environ.get("GROUP", "4"))
PROMPTS = int(os.environ.get("PROMPTS", "4"))
MASTER_PORT = int(os.environ.get("MASTER_PORT", "29641"))
GROUP_NAME = os.environ.get("GROUP_NAME", "higgs_tts_wsync")
TEMP = float(os.environ.get("TEMP", "0.8"))
MAX_NEW = int(os.environ.get("MAX_NEW", "256"))
TOP_K = int(os.environ["TOP_K"]) if os.environ.get("TOP_K") else None
TRAIN_MODE = os.environ.get("TRAIN_MODE", "lora").lower()
LR = float(os.environ.get("LR", "2e-5" if TRAIN_MODE == "lora" else "1e-6"))
EPS = 0.2


def post(path: str, body: dict, timeout: int = 300):
req = urllib.request.Request(
SERVER + path, data=json.dumps(body).encode(), headers={"Content-Type": "application/json"}
)
return json.loads(urllib.request.urlopen(req, timeout=timeout).read())


def rollout(input_ids: list[int], seed: int) -> dict:
sampling_params = {
"temperature": TEMP,
"top_p": 0.95,
"max_new_tokens": MAX_NEW,
"seed": seed,
}
if TOP_K is not None:
sampling_params["top_k"] = TOP_K
resp = post(
"/generate",
build_generate_payload(
input_ids,
sampling_params,
output_modalities=["audio"],
return_omni_rollout=True,
),
timeout=180,
)
result = parse_generate_response(resp)
stream = parse_omni_action_stream(result.omni_rollout, "higgs_codes")
if result.output_codebook_tokens != stream.actions:
raise ValueError("Higgs output_codebook_tokens do not match omni_rollout actions")
return {
"old": stream.logprobs,
"mask": stream.action_mask,
"codes": stream.actions,
"audio": (result.audio or {}).get("data"),
}


def main() -> None:
torch.cuda.set_device(0) # bind this process to its visible GPU for NCCL collectives

from sglang_omni.models.higgs_tts.text_tokenizer import HiggsTokenizerAdapter
from tokenizers import Tokenizer
from transformers import PreTrainedTokenizerFast

from miles_plugins.omni.higgs_actor import HiggsTtsActor, clipped_grpo_loss
from miles_plugins.omni.tts_reward import TtsCompositeReward

ckpt = glob.glob(HIGGS_CKPT)[0] if "*" in HIGGS_CKPT else HIGGS_CKPT
tok = PreTrainedTokenizerFast(tokenizer_object=Tokenizer.from_file(os.path.join(ckpt, "tokenizer.json")))
adapter = HiggsTokenizerAdapter(tok)
reward_fn = TtsCompositeReward()

actor = HiggsTtsActor(ckpt, device="cuda:0")
if TRAIN_MODE == "lora":
actor.fused_embed.requires_grad_(False)
actor.backbone = get_peft_model(
actor.backbone,
LoraConfig(r=8, lora_alpha=16, target_modules=["q_proj", "v_proj"], task_type=None),
)
elif TRAIN_MODE != "full":
raise ValueError(f"TRAIN_MODE must be 'lora' or 'full', got {TRAIN_MODE!r}")
actor.train()
trainable_params = [param for param in actor.parameters() if param.requires_grad]
opt = torch.optim.AdamW(trainable_params, lr=LR)

try:
from sglang.srt.utils import init_custom_process_group
except Exception:
from sglang.srt.utils.common import init_custom_process_group

# Rendezvous a 2-rank NCCL group with the served tts_engine stage (server = rank 1).
init_err: list = []

def _init_server():
try:
post(
"/init_weights_update_group",
{
"master_address": "localhost",
"master_port": MASTER_PORT,
"rank_offset": 1,
"world_size": 2,
"group_name": GROUP_NAME,
"backend": "nccl",
"stages": ["tts_engine"],
},
timeout=180,
)
except Exception as exc: # noqa: BLE001
init_err.append(exc)

th = threading.Thread(target=_init_server)
th.start()
pg = init_custom_process_group(
backend="nccl", init_method=f"tcp://localhost:{MASTER_PORT}", world_size=2, rank=0, group_name=GROUP_NAME
)
th.join()
torch.cuda.synchronize()
if init_err:
raise init_err[0]
print("WEIGHT_UPDATE_GROUP_READY", flush=True)

@torch.no_grad()
def merged_lora_weights() -> dict[str, torch.Tensor]:
out: dict[str, torch.Tensor] = {}
for name, mod in actor.backbone.named_modules():
if hasattr(mod, "lora_A") and hasattr(mod, "base_layer"):
a = mod.lora_A["default"].weight
b = mod.lora_B["default"].weight
scaling = mod.scaling["default"]
w = mod.base_layer.weight.data + scaling * (b @ a)
# peft module name base_model.model.layers.N... -> ckpt body.layers.N...
hf = name.replace("base_model.model.", "")
out["body." + hf + ".weight"] = w.to(torch.bfloat16).contiguous()
return out

@torch.no_grad()
def weights_to_sync() -> dict[str, torch.Tensor]:
if TRAIN_MODE == "lora":
return merged_lora_weights()
return {
name: tensor.detach().to(torch.bfloat16).contiguous()
for name, tensor in actor.full_server_weights().items()
}

def sync_to_server() -> int:
wd = weights_to_sync()
names = sorted(wd)
spec = {
"names": names,
"dtypes": [str(wd[n].dtype).replace("torch.", "") for n in names],
"shapes": [list(wd[n].shape) for n in names],
"group_name": GROUP_NAME,
"stages": ["tts_engine"],
}
err: list = []

def _update():
try:
post("/update_weights_from_distributed", spec, timeout=300)
except Exception as exc: # noqa: BLE001
err.append(exc)

t = threading.Thread(target=_update)
t.start()
for n in names:
torch.distributed.broadcast(wd[n], src=0, group=pg)
torch.cuda.synchronize()
t.join()
if err:
raise err[0]
return len(names)

data = [json.loads(line) for line in open(DATA)]
print("step | mean_reward | mean_cer | avg_loss | synced_params")
for step in range(STEPS):
opt.zero_grad()
step_reward, step_cer, n_cer, step_loss, n = 0.0, 0.0, 0, 0.0, 0
for ex in data[:PROMPTS]:
pid = list(map(int, adapter.build_prompt(ex["text"], num_ref_tokens=0)))
samples = [rollout(pid, step * 1000 + g) for g in range(GROUP)]
comps = [reward_fn.score(s["audio"], ex["label"]) for s in samples]
rewards = [c.reward for c in comps]
mean_r = sum(rewards) / len(rewards)
step_reward += mean_r
for c in comps:
if c.cer is not None:
step_cer += c.cer
n_cer += 1
for s, adv in zip(samples, [r - mean_r for r in rewards], strict=True):
codes = s["codes"]
if not codes or adv == 0.0:
continue
new = actor.codebook_logprobs(pid, codes, temperature=TEMP, top_k=TOP_K)
old = torch.tensor(s["old"], dtype=new.dtype, device="cuda:0")
mask = torch.tensor(s["mask"], dtype=torch.bool, device="cuda:0")
loss = clipped_grpo_loss(new, old, mask, advantage=adv, clip_eps=EPS)
loss = loss / (GROUP * PROMPTS)
loss.backward()
step_loss += loss.item() * (GROUP * PROMPTS)
n += 1
torch.nn.utils.clip_grad_norm_(trainable_params, 1.0)
opt.step()
synced = sync_to_server() # next step's rollouts are on-policy
mean_cer = step_cer / n_cer if n_cer else float("nan")
print(
f"{step:4d} | {step_reward / PROMPTS:11.3f} | {mean_cer:8.3f} | "
f"{step_loss / max(n, 1):8.4f} | {synced}",
flush=True,
)

print("Higgs TTS on-policy loop complete (per-step NCCL weight-sync to served tts_engine)")


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