Skip to content
Closed
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
42 changes: 38 additions & 4 deletions docs/features/additional-histories.mdx
Original file line number Diff line number Diff line change
Expand Up @@ -160,7 +160,7 @@ for history in histories:

Use `await art.tokenize_sampled(trajectories)` (or trajectory groups) when you
need complete native Chat Completions output with model-bound STOP flags. This
opt-in API first performs ordinary `multi_history=True` tokenization without
default `representation="rendered"` mode first performs ordinary `multi_history=True` tokenization without
renderer overrides or text reconciliation. After that succeeds, it resolves the
exact history model's tokenizer configuration, including its revision, to certify
each selected sampled source's nonempty original conditioning, output IDs and
Expand All @@ -187,15 +187,49 @@ a different `base_model` does not supply authority.

Incomplete native evidence, changed conditioning, unsupported sampled protocols,
and extra or incorrect STOP flags raise an error, even if ordinary tokenization
succeeded. Missing or null logprob carriers are not recorded NaNs; explicitly
recorded raw NaNs remain valid evidence. Unsupported finish reasons such as
`content_filter` are not certified as an absence of STOP. Nonsampled histories are retained. This API does not recover failed
succeeded. Nonsampled histories are retained. This default mode does not recover failed
rendering, split or join histories, or establish SFT equivalence.

Generic `art.tokenize` remains unchanged: a native-only path can avoid loading a
tokenizer, so missing STOP flags there do not prove that a terminating suffix is
absent. Use the explicit API when that distinction matters.

### Using recorded native conditioning instead of rendering

For a loss defined on **SAMPLED first occurrences, selected before filtering
nonfinite float32 logprobs**, you can explicitly select a native representation:

```python
tokenized = await art.tokenize_sampled(
trajectories, model="my-policy", representation="native"
)
```

This mode never renders messages or catches an ordinary tokenizer failure. It
requires complete original Chat Completions prompt/output tokens and logprobs,
validates the entire selected source inventory, and uses each source model's
resolved STOP authority. Content, literal reasoning delimiters, structured
reasoning and tool calls remain in the original captured objects. Missing
evidence is an error; a different renderer is not substituted.

Consecutive generations share a history only when each complete earlier native
prompt and output exactly prefixes the later request and their captured message
views agree. Other generations retain separate histories. Source encounters
remain in their original canonical order, including repeats across histories;
there is no global source deduplication. Native request gaps are exact,
**nonsampled context**, without guessed assistant/OUTPUT roles or synthetic
renderer STOP tails. IDs, conditioning, logprobs, sampled flags and STOP ownership
are checked again on constructed output.

This preserves the ordered sampled objective, not generic OUTPUT/SFT loss,
arbitrary per-history weighting, packing, or floating-point execution order.
Mixed protocols, legacy/additional histories, incomplete source messages and
edited projections are unsupported and raise errors. This includes assistant
turns echoed without their original structured reasoning and response IDs reused
within one model. Original trajectory/group
metadata and source objects are retained. The same current-model-configuration
limitation on historical STOP authority applies to both representations.

### Data Structure

The legacy `LegacyHistory` payload structure:
Expand Down
24 changes: 21 additions & 3 deletions src/art/trajectories/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -1613,6 +1613,7 @@ async def tokenize_sampled(
*,
model: str | None = None,
base_model: str | None = None,
representation: Literal["rendered", "native"] = "rendered",
) -> list[TokenizedMultiHistoryTrajectory]: ...


Expand All @@ -1622,6 +1623,7 @@ async def tokenize_sampled(
*,
model: str | None = None,
base_model: str | None = None,
representation: Literal["rendered", "native"] = "rendered",
) -> list[TokenizedTrajectoryGroup[TokenizedMultiHistoryTrajectory]]: ...


Expand All @@ -1630,14 +1632,15 @@ async def tokenize_sampled(
*,
model: str | None = None,
base_model: str | None = None,
representation: Literal["rendered", "native"] = "rendered",
) -> (
list[TokenizedMultiHistoryTrajectory]
| list[TokenizedTrajectoryGroup[TokenizedMultiHistoryTrajectory]]
):
"""Tokenize ordinary histories, then certify native sampled output and STOP.
"""Certify sampled output using ordinary rendering or explicit native sources.

This opt-in API uses ``multi_history=True`` without renderer overrides or
text reconciliation. It requires complete Chat Completions source messages,
The default ``representation="rendered"`` uses ``multi_history=True`` without
renderer overrides or text reconciliation. It requires complete Chat Completions source messages,
nonempty original conditioning, output IDs and logprobs for every sampled span.
Unsupported or incomplete sampled histories refuse; nonsampled histories
are retained. Ordinary tokenization failures propagate without recovery.
Expand All @@ -1651,9 +1654,23 @@ async def tokenize_sampled(
rendering, provide SFT equivalence or repartition histories. Generic
:func:`tokenize` retains its native-only, no-load behavior; absent STOP flags
there do not imply that a terminating suffix is known to be absent.

``representation="native"`` instead constructs complete original Chat
Completions sources without rendering. Consecutive sources join only when
every earlier prompt/output is an exact prefix of the later native request;
otherwise they remain separate histories. Repeated encounters across
canonical histories remain in order. This mode preserves the objective only
for SAMPLED first-occurrence ownership chosen BEFORE float32 finite-logprob
filtering. It does not preserve OUTPUT/SFT loss, per-history weighting,
layout, packing or floating-point execution order. Native request gaps are
exact nonsampled context; synthetic renderer STOP tails are not invented.
Mixed protocols, additional/legacy histories, incomplete or edited sources
refuse. STOP authority still follows each resolved source-model configuration.
"""
from ._parallel import transform

if representation not in {"rendered", "native"}:
raise ValueError("Unknown sampled representation")
return cast(
Any,
await transform(
Expand All @@ -1667,6 +1684,7 @@ async def tokenize_sampled(
chat_template=None,
chat_template_kwargs=None,
_sampled=True,
_native_sampled=representation == "native",
),
)

Expand Down
58 changes: 43 additions & 15 deletions src/art/trajectories/_parallel.py
Original file line number Diff line number Diff line change
Expand Up @@ -536,6 +536,7 @@ class _ProcessOptions:
chat_template: str | None
chat_template_kwargs: Mapping[str, object] | None
sampled: bool = False
native_sampled: bool = False


class _ProcessTransferError(RuntimeError):
Expand All @@ -555,21 +556,40 @@ def _tokenize_process_payload(payload: bytes) -> bytes:
raise _ProcessTransferError(
f"could not deserialize process input: {type(error).__name__}: {error}"
) from None
tokenized = trajectory.tokenize(
multi_history=options.multi_history,
reconcile_text_equivalent_tokenizations=(
options.reconcile_text_equivalent_tokenizations
),
model=options.model,
base_model=options.base_model,
tokenizer=None,
chat_template=options.chat_template,
chat_template_kwargs=options.chat_template_kwargs,
)
if options.sampled:
from ._sampled import reconcile_sampled_stops
if options.native_sampled:
if (
not options.sampled
or not options.multi_history
or options.reconcile_text_equivalent_tokenizations
or options.chat_template is not None
or options.chat_template_kwargs is not None
):
raise ValueError(
"Native representation requires unmodified sampled options"
)
from ._sampled_native import tokenize_native

tokenized = tokenize_native(
trajectory, model=options.model, base_model=options.base_model
)
else:
tokenized = trajectory.tokenize(
multi_history=options.multi_history,
reconcile_text_equivalent_tokenizations=(
options.reconcile_text_equivalent_tokenizations
),
model=options.model,
base_model=options.base_model,
tokenizer=None,
chat_template=options.chat_template,
chat_template_kwargs=options.chat_template_kwargs,
)
if options.sampled:
from ._sampled import reconcile_sampled_stops

tokenized = reconcile_sampled_stops(tokenized, base_model=options.base_model)
tokenized = reconcile_sampled_stops(
tokenized, base_model=options.base_model
)
try:
return pickle.dumps(tokenized, protocol=pickle.HIGHEST_PROTOCOL)
except Exception as error:
Expand Down Expand Up @@ -724,7 +744,10 @@ async def transform(
chat_template_kwargs: Mapping[str, object] | None,
device: Any = None,
_sampled: bool = False,
_native_sampled: bool = False,
) -> list[object]:
if _native_sampled and not _sampled:
raise ValueError("Native representation requires sampled tokenization")
if _sampled and (
operation != "tokenize"
or not multi_history
Expand All @@ -745,6 +768,10 @@ async def transform(
)

def convert(trajectory: Trajectory) -> object:
if _native_sampled:
from ._sampled_native import tokenize_native

return tokenize_native(trajectory, model=model, base_model=base_model)
tokenized = trajectory.tokenize(
multi_history=multi_history,
reconcile_text_equivalent_tokenizations=reconcile_text_equivalent_tokenizations,
Expand Down Expand Up @@ -775,7 +802,7 @@ def convert(trajectory: Trajectory) -> object:
capacity=capacity,
)
if _sampled:
key = (*key, "sampled_stops")
key = (*key, "sampled_native" if _native_sampled else "sampled_stops")
use_processes = _supports_processes(
capacity=capacity, size=len(leaves), tokenizer=tokenizer
) and _processes_enabled(key)
Expand All @@ -790,6 +817,7 @@ def convert(trajectory: Trajectory) -> object:
chat_template=chat_template,
chat_template_kwargs=chat_template_kwargs,
sampled=_sampled,
native_sampled=_native_sampled,
)
try:
workers = _process_workers(key, capacity=capacity, size=len(leaves))
Expand Down
Loading
Loading