SupritiVijay's picture
Add track-codes (training scripts + eval)
ae6dcfe verified
Raw
History Blame Contribute Delete
7.32 kB
"""
Reformat convergent_final_sft_format into a more compact XML layout.
All information retained — just denser nesting, fewer lines.
Original (verbose):
<stakeholder>Name</stakeholder>
<harms>
<action>
<action_name>Cat > Sub > Risk</action_name>
<effects>
<effect>
<effect_name>Effect Name</effect_name>
<immediacy>True</immediacy>
<extent>Significant</extent>
<likelihood>Medium</likelihood>
<effect_score>0.0596</effect_score>
</effect>
</effects>
</action>
</harms>
<harm_score>...</harm_score>
Compact:
<s name="Name">
<h>
<a name="Cat > Sub > Risk">
<e name="Effect Name" imm="T" ext="Significant" lik="Medium" score="0.0596"/>
</a>
</h>
<hs>...</hs>
<b/>
<bs>0</bs>
</s>
This should be ~40% fewer tokens while keeping all data.
"""
import re
import pandas as pd
from transformers import AutoTokenizer
def compact_effects(effects_block):
"""Convert <effects>...<effect>...</effect>...</effects> to inline <e .../> tags."""
effects = re.findall(
r'<effect>\s*'
r'<effect_name>(.*?)</effect_name>\s*'
r'<immediacy>(.*?)</immediacy>\s*'
r'<extent>(.*?)</extent>\s*'
r'<likelihood>(.*?)</likelihood>\s*'
r'<effect_score>(.*?)</effect_score>\s*'
r'</effect>',
effects_block, re.DOTALL
)
parts = []
for name, imm, ext, lik, score in effects:
imm_short = "T" if imm.strip() == "True" else "F"
parts.append(f'<e name="{name.strip()}" imm="{imm_short}" ext="{ext.strip()}" lik="{lik.strip()}" score="{score.strip()}"/>')
return parts
def compact_actions(section_content):
"""Convert <action>...<action_name>...<effects>... to <a name="...">...</a>."""
actions = re.findall(
r'<action>\s*<action_name>(.*?)</action_name>\s*<effects>(.*?)</effects>\s*</action>',
section_content, re.DOTALL
)
parts = []
for action_name, effects_block in actions:
effects = compact_effects(effects_block)
effects_str = "".join(effects)
parts.append(f'<a name="{action_name.strip()}">{effects_str}</a>')
return parts
def compact_stakeholder_block(block):
"""Convert a full stakeholder section into compact form."""
# Extract stakeholder name
name_match = re.search(r'<stakeholder>(.*?)</stakeholder>', block)
if not name_match:
return block
name = name_match.group(1).strip()
# Extract harms
harms_match = re.search(r'<harms>(.*?)</harms>', block, re.DOTALL)
harms_content = harms_match.group(1).strip() if harms_match else ""
harm_actions = compact_actions(harms_content) if harms_content else []
# Extract harm_score
hs_match = re.search(r'<harm_score>(.*?)</harm_score>', block)
harm_score = hs_match.group(1).strip() if hs_match else "0"
# Extract benefits
benefits_match = re.search(r'<benefits>(.*?)</benefits>', block, re.DOTALL)
benefits_content = benefits_match.group(1).strip() if benefits_match else ""
benefit_actions = compact_actions(benefits_content) if benefits_content else []
# Extract benefit_score
bs_match = re.search(r'<benefit_score>(.*?)</benefit_score>', block)
benefit_score = bs_match.group(1).strip() if bs_match else "0"
# Build compact
lines = [f'<s name="{name}">']
if harm_actions:
lines.append(f'<h>{"".join(harm_actions)}</h>')
else:
lines.append('<h/>')
lines.append(f'<hs>{harm_score}</hs>')
if benefit_actions:
lines.append(f'<b>{"".join(benefit_actions)}</b>')
else:
lines.append('<b/>')
lines.append(f'<bs>{benefit_score}</bs>')
lines.append('</s>')
return "".join(lines)
def compact_safety_check_score(score_block):
"""Convert the <safety_check_score> block to compact form."""
ht_match = re.search(r'<harms_total>(.*?)</harms_total>', score_block)
bt_match = re.search(r'<benefits_total>(.*?)</benefits_total>', score_block)
rs_match = re.search(r'<raw_score>(.*?)</raw_score>', score_block)
fs_match = re.search(r'<final_score>(.*?)</final_score>', score_block)
lb_match = re.search(r'<label>(.*?)</label>', score_block)
ht = ht_match.group(1).strip() if ht_match else "0"
bt = bt_match.group(1).strip() if bt_match else "0"
rs = rs_match.group(1).strip() if rs_match else "0"
fs = fs_match.group(1).strip() if fs_match else "0"
lb = lb_match.group(1).strip() if lb_match else "unknown"
return f'<score ht="{ht}" bt="{bt}" raw="{rs}" final="{fs}" label="{lb}"/>'
def reformat_to_compact(text):
"""Convert full convergent_final_sft_format to compact XML."""
# Split into safety_check content and the rest (think + response)
sc_match = re.search(r'<safety_check>(.*?)</safety_check>', text, re.DOTALL)
if not sc_match:
return text
sc_content = sc_match.group(1)
after_sc = text[sc_match.end():]
# Extract and compact the score section first
score_match = re.search(r'<safety_check_score>(.*?)</safety_check_score>', sc_content, re.DOTALL)
score_compact = ""
if score_match:
score_compact = compact_safety_check_score(score_match.group(1))
sc_content = sc_content[:score_match.start()] + sc_content[score_match.end():]
# Split into stakeholder blocks
# Each block starts with <stakeholder> and ends before the next <stakeholder> or end
stakeholder_splits = re.split(r'(?=<stakeholder>)', sc_content)
compact_blocks = []
for block in stakeholder_splits:
block = block.strip()
if not block or '<stakeholder>' not in block:
continue
compact_blocks.append(compact_stakeholder_block(block))
# Reconstruct
compact_sc = "<safety_check>" + "".join(compact_blocks) + score_compact + "</safety_check>"
return compact_sc + after_sc
def main():
print("Loading data...")
df = pd.read_parquet("data/sft-training-data/convergent_data_10k.parquet")
print("Reformatting to compact XML...")
df["convergent_final_sft_format"] = df["convergent_final_sft_format"].apply(reformat_to_compact)
# Check token savings
tokenizer = AutoTokenizer.from_pretrained("models/starting-checkpoint")
original_df = pd.read_parquet("data/sft-training-data/convergent_data_10k.parquet")
orig_tokens = []
compact_tokens = []
for orig, compact in zip(original_df["convergent_final_sft_format"], df["convergent_final_sft_format"]):
orig_tokens.append(len(tokenizer.encode(orig)))
compact_tokens.append(len(tokenizer.encode(compact)))
import numpy as np
orig_arr = np.array(orig_tokens)
comp_arr = np.array(compact_tokens)
savings = 1 - comp_arr.mean() / orig_arr.mean()
print(f"Original: mean={orig_arr.mean():.0f}, max={orig_arr.max()}")
print(f"Compact: mean={comp_arr.mean():.0f}, max={comp_arr.max()}")
print(f"Savings: {savings:.1%} fewer tokens")
# Verify no data loss on a sample
print("\nSample compact output:")
print(df["convergent_final_sft_format"].iloc[0][:1500])
output_path = "data/sft-training-data/convergent_data_10k_compact.parquet"
df.to_parquet(output_path)
print(f"\nSaved to: {output_path}")
if __name__ == "__main__":
main()