Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
Original file line number Diff line number Diff line change
Expand Up @@ -135,17 +135,19 @@ numpy==1.26.4
oauthlib==3.2.2
onnxruntime==1.19.2
openai==1.47.0
opentelemetry-api==1.27.0
opentelemetry-exporter-otlp-proto-common==1.27.0
opentelemetry-exporter-otlp-proto-grpc==1.27.0
opentelemetry-exporter-otlp-proto-http==1.27.0
opentelemetry-instrumentation==0.48b0
opentelemetry-instrumentation-asgi==0.48b0
opentelemetry-instrumentation-fastapi==0.48b0
opentelemetry-proto==1.27.0
opentelemetry-sdk==1.27.0
opentelemetry-semantic-conventions==0.48b0
opentelemetry-util-http==0.48b0
# 1.28 is the first line that admits protobuf 5.x (1.27 requires protobuf<5.0).
# Keep the 1.28.2 / 0.49b2 set together; mixing minors fails to resolve.
opentelemetry-api==1.28.2
opentelemetry-exporter-otlp-proto-common==1.28.2
opentelemetry-exporter-otlp-proto-grpc==1.28.2
opentelemetry-exporter-otlp-proto-http==1.28.2
opentelemetry-instrumentation==0.49b2
opentelemetry-instrumentation-asgi==0.49b2
opentelemetry-instrumentation-fastapi==0.49b2
opentelemetry-proto==1.28.2
opentelemetry-sdk==1.28.2
opentelemetry-semantic-conventions==0.49b2
opentelemetry-util-http==0.49b2
orderly-set==5.2.2
orjson==3.10.7
overrides==7.7.0
Expand All @@ -161,7 +163,9 @@ posthog==3.6.6
prettytable==3.11.0
prompt_toolkit==3.0.47
proto-plus==1.24.0
protobuf==4.25.5
# GHSA-7GCM-G887-7QV7 (CVE-2026-0994): 4.25.5 -> first patched 5.x release.
# streamlit==1.38.0 caps protobuf <6, so stay on the 5.x line.
protobuf==5.29.6
psutil==6.0.0
ptyprocess==0.7.0
pulsar-client==3.5.0
Expand Down
28 changes: 16 additions & 12 deletions archived/learn/rag/project_simple-rag-with-chroma/requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -135,17 +135,19 @@ numpy==1.26.4
oauthlib==3.2.2
onnxruntime==1.19.2
openai==1.47.0
opentelemetry-api==1.27.0
opentelemetry-exporter-otlp-proto-common==1.27.0
opentelemetry-exporter-otlp-proto-grpc==1.27.0
opentelemetry-exporter-otlp-proto-http==1.27.0
opentelemetry-instrumentation==0.48b0
opentelemetry-instrumentation-asgi==0.48b0
opentelemetry-instrumentation-fastapi==0.48b0
opentelemetry-proto==1.27.0
opentelemetry-sdk==1.27.0
opentelemetry-semantic-conventions==0.48b0
opentelemetry-util-http==0.48b0
# 1.28 is the first line that admits protobuf 5.x (1.27 requires protobuf<5.0).
# Keep the 1.28.2 / 0.49b2 set together; mixing minors fails to resolve.
opentelemetry-api==1.28.2
opentelemetry-exporter-otlp-proto-common==1.28.2
opentelemetry-exporter-otlp-proto-grpc==1.28.2
opentelemetry-exporter-otlp-proto-http==1.28.2
opentelemetry-instrumentation==0.49b2
opentelemetry-instrumentation-asgi==0.49b2
opentelemetry-instrumentation-fastapi==0.49b2
opentelemetry-proto==1.28.2
opentelemetry-sdk==1.28.2
opentelemetry-semantic-conventions==0.49b2
opentelemetry-util-http==0.49b2
orderly-set==5.2.2
orjson==3.10.7
overrides==7.7.0
Expand All @@ -161,7 +163,9 @@ posthog==3.6.6
prettytable==3.11.0
prompt_toolkit==3.0.47
proto-plus==1.24.0
protobuf==4.25.5
# GHSA-7GCM-G887-7QV7 (CVE-2026-0994): 4.25.5 -> first patched 5.x release.
# streamlit==1.38.0 caps protobuf <6, so stay on the 5.x line.
protobuf==5.29.6
psutil==6.0.0
ptyprocess==0.7.0
pulsar-client==3.5.0
Expand Down
8 changes: 8 additions & 0 deletions skills/fireworks-training/references/rl-async.md
Original file line number Diff line number Diff line change
Expand Up @@ -94,6 +94,14 @@ run. Each trajectory is then packed as a separate trainer datum. Thus a group
with `R` surviving runs has `R` rewards and advantages, while its trainer datum
count is `sum(len(run.segments) for run in runs)` and may exceed `R`.

`main(..., advantage_fn=...)` can override the per-group reward-to-advantage
calculation. The default `compute_advantages` subtracts the group mean and
divides by the group standard deviation. For mean-only REINFORCE, pass a
function that subtracts the mean without that division. Reward scaling belongs
to the caller; the recipe does not automatically convert binary rewards to
`-1/+1`. This choice is independent of token-level score centering and gradient
accumulation normalization.

`RolloutSetup` contains the recipe-owned sampler, tokenizer, tokenizer ID,
sampling kwargs, inference base URL, API key, deployment model, group size, and
caller-provided `extras`. Treat `setup.sampler` as borrowed: reuse it for every
Expand Down
10 changes: 6 additions & 4 deletions training/examples/multihop_qa/train_multihop_qa_igpo.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,7 @@
wandb_finish,
log_metrics_json,
build_service_client,
make_weight_sync,
read_api_extra_headers_env,
validate_config,
load_jsonl_dataset,
Expand Down Expand Up @@ -382,6 +383,9 @@ def _signal_handler(signum, frame):
job_id=service.trainer_job_id,
service=service,
)
publish_weights = make_weight_sync(
policy, service, cfg.deployment, extended=True
)
policy_job_id = service.trainer_job_id
sampler = service.create_deployment_sampler()
reference = None
Expand Down Expand Up @@ -416,8 +420,7 @@ def _signal_handler(signum, frame):
name = (
f"resume-{step_offset}-base" if step_offset > 0 else "step-0-base"
)
saved = policy.save_weights_for_sampler_ext(name, checkpoint_type="base")
service.hotload_sampler_snapshot(saved.snapshot_name)
publish_weights(name, checkpoint_type="base")

# Rollout processor
rollout_base_url = sampler.base_url.rstrip("/") + (
Expand Down Expand Up @@ -717,8 +720,7 @@ def train_step(
and step % WEIGHT_SYNC_INTERVAL == 0
):
with timer("weight_sync"):
saved = policy.save_weights_for_sampler_ext(f"step-{step}")
service.hotload_sampler_snapshot(saved.snapshot_name)
publish_weights(f"step-{step}")

if (
DCP_SAVE_INTERVAL > 0
Expand Down
10 changes: 6 additions & 4 deletions training/examples/rl/frozen_lake/train_frozen_lake.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@
wandb_finish,
log_metrics_json,
build_service_client,
make_weight_sync,
read_api_extra_headers_env,
build_training_datum_from_token_mask,
validate_config,
Expand Down Expand Up @@ -526,6 +527,9 @@ def _signal_handler(signum, frame):
job_id=service.trainer_job_id,
service=service,
)
publish_weights = make_weight_sync(
policy, service, cfg.deployment, extended=True
)
policy_job_id = service.trainer_job_id
sampler = service.create_deployment_sampler()
reference = None
Expand Down Expand Up @@ -555,8 +559,7 @@ def _signal_handler(signum, frame):
step_offset = resume_info.step if resume_info else 0
prior_rows_consumed = resume_info.data_consumed if resume_info else 0
name = f"resume-{step_offset}-base" if step_offset > 0 else "step-0-base"
saved = policy.save_weights_for_sampler_ext(name, checkpoint_type="base")
service.hotload_sampler_snapshot(saved.snapshot_name)
publish_weights(name, checkpoint_type="base")

# -- Build rollout processor ----------------------------------------
rollout_base_url = sampler.base_url.rstrip("/") + (
Expand Down Expand Up @@ -751,8 +754,7 @@ def save_intermediate_checkpoint(step: int, data_consumed: int) -> None:
def _weight_sync(step: int) -> float:
logger.info("[step %d] weight_sync: saving + loading...", step)
with wall_timer() as span:
saved = policy.save_weights_for_sampler_ext(f"step-{step}")
service.hotload_sampler_snapshot(saved.snapshot_name)
publish_weights(f"step-{step}")
logger.info(
"[step %d] weight_sync: done (%.1fs)",
step,
Expand Down
2 changes: 1 addition & 1 deletion training/examples/rl/harbor/recipes/textworld/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -168,7 +168,7 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
choices=("binary_tv", "binary_kl"),
default="binary_tv",
)
parser.add_argument("--dppo-threshold", type=float, default=0.15)
parser.add_argument("--dppo-threshold", type=float, default=None)
parser.add_argument("--dppo-ratio-log-cap", type=float, default=20.0)
parser.add_argument(
"--score-centering-top-k",
Expand Down
2 changes: 1 addition & 1 deletion training/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@ name = "fireworks-training-cookbook"
version = "0.1.0"
requires-python = ">=3.11"
dependencies = [
"fireworks-ai[training]>=1.2.14,<2",
"fireworks-ai[training]>=1.2.15,<2",
# Preserve the runtime floors previously supplied by tinker-cookbook.
"aiohttp>=3.9.0",
"pydantic>=2.0.0",
Expand Down
24 changes: 16 additions & 8 deletions training/recipes/async_rl_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@
WandBConfig,
ReconnectableClient,
build_service_client,
make_weight_sync,
flush_phase_trace,
log_metrics,
load_deployment_tokenizer,
Expand All @@ -67,6 +68,7 @@
)
from training.utils.checkpoints import TrainingCheckpoints, validate_warm_start_config
from training.utils.dataloader import CursorDataLoader
from training.utils.data import compute_advantages
from training.utils.logging import ASYNC_RL_WANDB_METRIC_STEPS
from training.utils.rl import PromptGroup
from training.utils.rl.async_rl import (
Expand Down Expand Up @@ -584,6 +586,7 @@ def main(
config: Config,
*,
rollout_fn_factory: RolloutFnFactory,
advantage_fn: Callable[[list[float]], list[float]] = compute_advantages,
dynamic_filter_fn: DynamicFilterFn | None = None,
evaluation_fn: RolloutEvaluationFn | None = None,
evaluation_interval: int = 1,
Expand All @@ -598,6 +601,10 @@ def main(
``completions_per_prompt`` times per dataset row (each invocation is
one trajectory draw against the inference deployment).

``advantage_fn`` maps one prompt group's rewards to per-rollout advantages.
The default subtracts the group mean and divides by its standard deviation.
Supply a mean-only function for REINFORCE with group-centered rewards.

Remote trainer and sampler setup and lifecycle are owned by the SDK-managed
Tinker path.
"""
Expand Down Expand Up @@ -809,6 +816,7 @@ def _signal_handler(signum, _):
job_id=service.trainer_job_id,
service=service,
)
publish_weights = make_weight_sync(policy, service, cfg.deployment)
reference = None
if cfg.kl_beta > 0:
reference_training_client = service.create_reference_client(
Expand Down Expand Up @@ -844,11 +852,7 @@ def _signal_handler(signum, _):
)

with elapsed_timer("weight_sync") as span:
saved = policy.save_weights_for_sampler(
f"step-{step_offset}",
checkpoint_type="base",
)
service.hotload_sampler_snapshot(saved.path)
publish_weights(f"step-{step_offset}", checkpoint_type="base")
logger.info(
"[step %d] initial weight sync (%.1fs)",
step_offset,
Expand Down Expand Up @@ -1096,7 +1100,11 @@ def fwd_bwd_batch(
)
elif cfg.policy_loss == "dppo":
loss_fn = make_dppo_loss_fn(
**common_loss_kwargs,
advantages=adv,
ref_logprobs=ref_lp,
inf_logprobs=inf_lp,
prompt_len=prompt_lens,
old_policy_logprobs=old_policy_logprobs,
dppo_config=cfg.dppo,
)
else:
Expand Down Expand Up @@ -1202,8 +1210,7 @@ def optimizer_step(step: int) -> dict[str, Any]:

def sync_weights(step: int) -> float:
with wall_timer() as span:
saved = policy.save_weights_for_sampler(f"step-{step}")
service.hotload_sampler_snapshot(saved.path)
publish_weights(f"step-{step}")
return span.elapsed

async def run_training() -> tuple[int, dict[str, Any]]:
Expand All @@ -1222,6 +1229,7 @@ def _step_metrics(metrics: dict[str, Any], step: int) -> None:
)
coordinator = AsyncRLCoordinator(
rows=make_row_requests(),
advantage_fn=advantage_fn,
completions_per_prompt=cfg.completions_per_prompt,
prompt_groups_per_step=cfg.prompt_groups_per_step,
training_chunks_per_step=cfg.pipeline_chunks_per_step,
Expand Down
10 changes: 6 additions & 4 deletions training/recipes/distillation_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@
TrainerConfig,
WandBConfig,
build_service_client,
make_weight_sync,
load_jsonl_dataset,
load_deployment_tokenizer,
log_metrics_json,
Expand Down Expand Up @@ -549,6 +550,9 @@ def _signal_handler(signum, frame):
default_timeout=cfg.step_timeout or 3600,
service=service,
)
publish_weights = make_weight_sync(
policy, service, cfg.deployment, extended=True
)
tokenizer = load_deployment_tokenizer(cfg.deployment)
max_seq_len = service.max_context_length
deployment_id = service.deployment_id
Expand Down Expand Up @@ -693,8 +697,7 @@ async def _score_topk_routed(

if cfg.weight_sync_before_training:
name = f"resume-{step_offset}-base" if step_offset > 0 else "step-0-base"
saved = policy.save_weights_for_sampler_ext(name, checkpoint_type="base")
service.hotload_sampler_snapshot(saved.snapshot_name)
publish_weights(name, checkpoint_type="base")

# -- Prepare sampling and training --------------------------------------

Expand Down Expand Up @@ -989,8 +992,7 @@ def train_step(
logger.info("[step %d] weight_sync: saving + loading...", step)
t0 = _time.time()
with timer("weight_sync"):
saved = policy.save_weights_for_sampler_ext(f"step-{step}")
service.hotload_sampler_snapshot(saved.snapshot_name)
publish_weights(f"step-{step}")
logger.info("[step %d] weight_sync: done (%.1fs)", step, _time.time() - t0)
if (
cfg.step_eval is not None
Expand Down
13 changes: 6 additions & 7 deletions training/recipes/igpo_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@
RawRowCursor,
RLPromptDataset,
build_service_client,
make_weight_sync,
wandb_log,
setup_wandb,
wandb_finish,
Expand Down Expand Up @@ -335,6 +336,9 @@ def _signal_handler(signum, frame):
default_timeout=_timeout,
service=service,
)
publish_weights = make_weight_sync(
policy, service, cfg.deployment, extended=True
)
# KL reference (optional in iGPO): the SDK owns the shared-vs-separate
# decision. LoRA without an explicit reference shape reuses the policy
# session; full-param (or an explicit reference_training_shape_id)
Expand Down Expand Up @@ -377,11 +381,7 @@ def _signal_handler(signum, frame):

if cfg.weight_sync_before_training:
with timer("weight_sync"):
saved = policy.save_weights_for_sampler_ext(
f"step-{step_offset}",
checkpoint_type="base",
)
service.hotload_sampler_snapshot(saved.snapshot_name)
publish_weights(f"step-{step_offset}", checkpoint_type="base")

# Dataset
raw_dataset = load_jsonl_dataset(cfg.dataset, cfg.max_rows)
Expand Down Expand Up @@ -688,8 +688,7 @@ def train_step(
# 5. Sync weights
if cfg.weight_sync_interval > 0 and step % cfg.weight_sync_interval == 0:
with timer("weight_sync"):
saved = policy.save_weights_for_sampler_ext(f"step-{step}")
service.hotload_sampler_snapshot(saved.snapshot_name)
publish_weights(f"step-{step}")
if cfg.dcp_save_interval > 0 and step % cfg.dcp_save_interval == 0:
ckpt.save(
f"step-{step}",
Expand Down
11 changes: 4 additions & 7 deletions training/recipes/rl_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,7 @@
WandBConfig,
build_renderer,
build_service_client,
make_weight_sync,
load_deployment_tokenizer,
load_jsonl_dataset,
log_metrics,
Expand Down Expand Up @@ -303,6 +304,7 @@ def _signal_handler(signum, _):
job_id=service.trainer_job_id,
service=service,
)
publish_weights = make_weight_sync(policy, service, cfg.deployment)
reference = None
if cfg.kl_beta > 0:
reference = ReconnectableClient.from_training_client(
Expand Down Expand Up @@ -344,11 +346,7 @@ def _signal_handler(signum, _):
# The synchronous recipe is always strict on-policy: initialize the
# sampler from the trainer, then repeat this sync after every update.
with elapsed_timer("weight_sync") as span:
saved = policy.save_weights_for_sampler(
f"step-{step_offset}",
checkpoint_type="base",
)
service.hotload_sampler_snapshot(saved.path)
publish_weights(f"step-{step_offset}", checkpoint_type="base")
logger.info("[step %d] initial weight sync (%.1fs)", step_offset, span.elapsed)
flush_timing()

Expand Down Expand Up @@ -567,8 +565,7 @@ async def run_training() -> int:

# 4. Publish this policy before the next rollout batch.
with elapsed_timer("weight_sync"):
saved = policy.save_weights_for_sampler(f"step-{step}")
service.hotload_sampler_snapshot(saved.path)
publish_weights(f"step-{step}")

for index in row_indices:
row_loader.mark_resolved(index)
Expand Down
Loading
Loading