openai_format SFT dataset fails at validation with jinja2.exceptions.UndefinedError: list object has no element -1
- Dominant language
- Python
- Stars
- 2k
- Forks
- 561
- Avg merge
- 4d 5h
- Merged PRs (30d)
- 145
Description
**Describe the bug**
I followed this [guide](https://docs.nvidia.com/nemo/rl/latest/guides/sft.html) to set system prompt using the openai format. After setting my data config to:
```
data:
dataset_name: openai_format
train_data_path: "/home/ameliay/RL/nemo_rl/data/datasets/sft_fraud_datasets/nss_to_rl_dataset_20251022_060530_train_openai.jsonl" # Path to training data
val_data_path: "/home/ameliay/RL/nemo_rl/data/datasets/sft_fraud_datasets/nss_to_rl_dataset_20251022_060530_val_openai.jsonl" # Path to validation data
chat_key: "messages" # Key for messages in the data (default: "messages")
system_key: null # Key for system message in the data (optional)
system_prompt: "You are a helpful assistant." # Default system prompt if not in data (optional)
tool_key: "tools" # Key for tools in the data (default: "tools")
use_preserving_dataset: false # Set to true for heterogeneous tool schemas (see below)
add_bos: true
add_eos: true
add_generation_prompt: true
shuffle: true
num_workers: 0
max_input_seq_length: ${policy.max_total_sequence_length}
```
This template error shows up:
```
Starting validation at step 0...
Traceback (most recent call last):
File "/home/ameliay/RL/examples/run_sft.py", line 220, in
main()
File "/home/ameliay/RL/examples/run_sft.py", line 205, in main
sft_train(
File "/home/ameliay/RL/nemo_rl/algorithms/sft.py", line 393, in sft_train
val_metrics, validation_timings = validate(
^^^^^^^^^
File "/home/ameliay/RL/nemo_rl/algorithms/sft.py", line 263, in validate
for batch_idx, val_batch in enumerate(val_dataloader):
^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/torchdata/stateful_dataloader/stateful_dataloader.py", line 450, in __next__
return super().__next__()
^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/torch/utils/data/dataloader.py", line 734, in __next__
data = self._next_data()
^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/torchdata/stateful_dataloader/stateful_dataloader.py", line 491, in _next_data
data = self._dataset_fetcher.fetch(index) # may raise StopIteration
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/torch/utils/data/_utils/fetch.py", line 52, in fetch
data = [self.dataset[idx] for idx in possibly_batched_index]
~~~~~~~~~~~~^^^^^
File "/home/ameliay/RL/nemo_rl/data/datasets/processed_dataset.py", line 107, in __getitem__
datum_spec = task_data_processor(
^^^^^^^^^^^^^^^^^^^^
File "/home/ameliay/RL/examples/run_sft.py", line 69, in sft_preprocessor
message_log = get_formatted_message_log(
^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ameliay/RL/nemo_rl/data/llm_message_utils.py", line 549, in get_formatted_message_log
formatted_message: str = tokenizer.apply_chat_template( # type: ignore
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/transformers/tokenization_utils_base.py", line 1640, in apply_chat_template
rendered_chat, generation_indices = render_jinja_template(
^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/transformers/utils/chat_template_utils.py", line 521, in render_jinja_template
rendered_chat = compiled_template.render(
^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/jinja2/environment.py", line 1295, in render
self.environment.handle_exception()
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/jinja2/environment.py", line 942, in handle_exception
raise rewrite_traceback_stack(source=source)
File "", line 15, in top-level template code
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/jinja2/sandbox.py", line 293, in getitem
return obj[argument]
~~~^^^^^^^^^^
jinja2.exceptions.UndefinedError: list object has no element -1
Traceback (most recent call last):
File "/home/ameliay/RL/examples/run_sft.py", line 220, in
main()
File "/home/ameliay/RL/examples/run_sft.py", line 205, in main
sft_train(
File "/home/ameliay/RL/nemo_rl/algorithms/sft.py", line 393, in sft_train
val_metrics, validation_timings = validate(
^^^^^^^^^
File "/home/ameliay/RL/nemo_rl/algorithms/sft.py", line 263, in validate
for batch_idx, val_batch in enumerate(val_dataloader):
^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/torchdata/stateful_dataloader/stateful_dataloader.py", line 450, in __next__
return super().__next__()
^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/torch/utils/data/dataloader.py", line 734, in __next__
data = self._next_data()
^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/torchdata/stateful_dataloader/stateful_dataloader.py", line 491, in _next_data
data = self._dataset_fetcher.fetch(index) # may raise StopIteration
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/torch/utils/data/_utils/fetch.py", line 52, in fetch
data = [self.dataset[idx] for idx in possibly_batched_index]
~~~~~~~~~~~~^^^^^
File "/home/ameliay/RL/nemo_rl/data/datasets/processed_dataset.py", line 107, in __getitem__
datum_spec = task_data_processor(
^^^^^^^^^^^^^^^^^^^^
File "/home/ameliay/RL/examples/run_sft.py", line 69, in sft_preprocessor
message_log = get_formatted_message_log(
^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/home/ameliay/RL/nemo_rl/data/llm_message_utils.py", line 549, in get_formatted_message_log
formatted_message: str = tokenizer.apply_chat_template( # type: ignore
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/transformers/tokenization_utils_base.py", line 1640, in apply_chat_template
rendered_chat, generation_indices = render_jinja_template(
^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/transformers/utils/chat_template_utils.py", line 521, in render_jinja_template
rendered_chat = compiled_template.render(
^^^^^^^^^^^^^^^^^^^^^^^^^
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/jinja2/environment.py", line 1295, in render
self.environment.handle_exception()
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/jinja2/environment.py", line 942, in handle_exception
raise rewrite_traceback_stack(source=source)
File "", line 15, in top-level template code
File "/opt/nemo_rl_venv/lib/python3.12/site-packages/jinja2/sandbox.py", line 293, in getitem
return obj[argument]
~~~^^^^^^^^^^
jinja2.exceptions.UndefinedError: list object has no element -1
wandb:
wandb: 🚀 View run openmathinstruct-nemorl-1M_train at: https://wandb.ai/amelia123-nvidia/sft-dev/runs/g55nxuqw
(DTensorPolicyWorker pid=1030879) No weights path provided. Starting from scratch (default policy init) [repeated 7x across cluster]
```
**Steps/Code to reproduce bug**
1. Prepare a dataset in OpenAI chat-completions format.
2. Set dataset_name=openai_format and define a system_prompt in the data config.
3. Run SFT:
`uv run examples/run_sft.py --config=examples/configs/sft_openmathinstruct2.yaml
`
4. Validation begins and crashes with the above Jinja error.
**Expected behavior**
The system prompt should be passed into the template seamlessly.
**Additional context**
Contributor guide
Assessment
This issue has not been assessed yet.