ftb-sciworld-repro / tcod /tests /common /experience_test.py
SeanWang0027's picture
Upload folder using huggingface_hub
8c9ba62 verified
Raw
History Blame Contribute Delete
17.5 kB
# -*- coding: utf-8 -*-
"""Test cases for Storage modules."""
import os
import unittest
import torch
from trinity.buffer.schema.sql_schema import ExperienceModel
from trinity.common.experience import EID, CustomField, Experience, Experiences
db_url = os.path.join(os.path.dirname(__file__), "tmp", "test.db")
dataset_path = os.path.join(os.path.dirname(__file__), "data")
class TestEID(unittest.TestCase):
def test_eid_properties(self):
# test properties
eid = EID(batch=1, task=2, run=3, step=4, suffix="abc123")
self.assertEqual(eid.uid, "1/2/3/4/abc123")
self.assertEqual(eid.sid, "1/2/4")
self.assertEqual(eid.rid, "1/2/3")
self.assertEqual(eid.tid, "1/2")
self.assertEqual(str(eid), "1/2/3/4/abc123")
self.assertIn("EID(batch=1, task=2, run=3, step=4, uuid=abc123)", repr(eid))
# test unique
eid1 = EID(batch=1, task=2, run=3, step=4)
eid2 = EID(batch=1, task=2, run=3, step=4)
self.assertNotEqual(eid1.suffix, eid2.suffix)
self.assertNotEqual(eid1.uid, eid2.uid)
# test default
eid = EID()
eid2 = EID()
self.assertIsInstance(eid.suffix, str)
self.assertEqual(eid.batch, "")
self.assertEqual(eid.task, "")
self.assertEqual(eid.run, 0)
self.assertEqual(eid.step, 0)
self.assertNotEqual(eid.uid, eid2.uid)
class TestExperience(unittest.TestCase):
def test_single_turn_experience(self):
tokens = torch.tensor([10, 11, 12], dtype=torch.int32)
logprobs = torch.tensor([0.2, 0.3], dtype=torch.float32)
exp = Experience(tokens=tokens, logprobs=logprobs, reward=1.0, prompt_length=1)
self.assertEqual(exp.experience_type, "single_turn")
self.assertTrue(torch.equal(exp.tokens, tokens))
self.assertTrue(torch.equal(exp.logprobs, logprobs))
self.assertEqual(exp.reward, 1.0)
self.assertEqual(exp.prompt_length, 1)
self.assertTrue(torch.equal(exp.action_mask, torch.tensor([1, 1], dtype=torch.bool)))
def test_multi_turn_experience(self):
tokens = torch.tensor([1, 2, 3, 4])
logprobs = torch.tensor([0.1, 0.2, 0.3, 0.4])
action_mask = torch.tensor([1, 0, 1, 0], dtype=torch.bool)
exp = Experience(tokens=tokens, logprobs=logprobs, reward=2.0, action_mask=action_mask)
self.assertEqual(exp.experience_type, "multi_turn")
self.assertTrue(torch.equal(exp.action_mask, action_mask))
self.assertEqual(exp.prompt_length, 1)
def test_dpo_experience(self):
tokens = torch.tensor([1, 2])
chosen = torch.tensor([3, 4])
rejected = torch.tensor([5, 6])
exp = Experience(tokens=tokens, chosen=chosen, rejected=rejected, reward=0.5)
self.assertEqual(exp.experience_type, "dpo")
self.assertTrue(torch.equal(exp.chosen, chosen))
self.assertTrue(torch.equal(exp.rejected, rejected))
self.assertEqual(exp.prompt_length, 2)
def test_serialize_deserialize(self):
tokens = torch.tensor([1, 2, 3])
exp = Experience(tokens=tokens, reward=1.23, prompt_length=1)
data = exp.serialize()
exp2 = Experience.deserialize(data)
self.assertTrue(torch.equal(exp.tokens, exp2.tokens))
self.assertEqual(exp.reward, exp2.reward)
self.assertEqual(exp.prompt_length, exp2.prompt_length)
self.assertEqual(exp.experience_type, exp2.experience_type)
def test_to_dict(self):
tokens = torch.tensor([1, 2, 3])
exp = Experience(
tokens=tokens, reward=2.5, prompt_length=1, prompt_text="hi", response_text="yo"
)
d = exp.to_dict()
self.assertIn("eid", d)
self.assertIn("type", d)
self.assertIn("reward", d)
self.assertEqual(d["prompt_text"], "hi")
self.assertEqual(d["response_text"], "yo")
self.assertEqual(d["reward"], 2.5)
def test_gather(self):
# test empty gathering
batch = Experiences.gather_experiences([])
self.assertEqual(batch.tokens.numel(), 0)
self.assertEqual(batch.rewards.numel(), 0)
self.assertEqual(batch.eids, [])
# test single experience gathering
exp = Experience(tokens=torch.tensor([1, 2, 3]), reward=1.0, prompt_length=1)
batch = Experiences.gather_experiences([exp])
self.assertEqual(batch.batch_size, 1)
self.assertTrue(
torch.equal(batch.tokens[0], torch.tensor([0, 1, 2, 3], dtype=torch.int64)[-3:])
)
self.assertEqual(batch.prompt_length, 1)
self.assertEqual(batch.rewards[0], 1.0)
# test multiple experiences gathering
exps = [
Experience(tokens=torch.tensor([1, 2]), reward=0.1, prompt_length=1),
Experience(tokens=torch.tensor([3, 4, 5]), reward=0.2, prompt_length=2),
]
batch = Experiences.gather_experiences(exps)
self.assertEqual(batch.batch_size, 2)
self.assertEqual(batch.prompt_length, 2)
self.assertEqual(batch.tokens.shape[1], 3)
self.assertEqual(batch.rewards[0], 0.1)
self.assertEqual(batch.rewards[1], 0.2)
def test_gather_with_token_level_reward(self):
# test empty gathering
batch = Experiences.gather_experiences([])
self.assertEqual(batch.tokens.numel(), 0)
self.assertEqual(batch.rewards.numel(), 0)
self.assertEqual(batch.token_level_rewards.numel(), 0)
self.assertEqual(batch.eids, [])
# test single experience gathering
exp = Experience(
tokens=torch.tensor([1, 2, 3]),
token_level_reward=torch.tensor([0, 1.0]),
prompt_length=1,
)
batch = Experiences.gather_experiences([exp])
self.assertEqual(batch.batch_size, 1)
self.assertTrue(
torch.equal(batch.tokens[0], torch.tensor([0, 1, 2, 3], dtype=torch.int64)[-3:])
)
self.assertEqual(batch.prompt_length, 1)
self.assertIsNone(batch.rewards)
self.assertTrue(torch.equal(batch.token_level_rewards[0], torch.tensor([0, 1.0])))
# test multiple experiences gathering
exps = [
Experience(
tokens=torch.tensor([1, 2]), token_level_reward=torch.tensor([0.1]), prompt_length=1
),
Experience(
tokens=torch.tensor([3, 4, 5]),
token_level_reward=torch.tensor([0.2]),
prompt_length=2,
),
]
batch = Experiences.gather_experiences(exps)
self.assertEqual(batch.batch_size, 2)
self.assertEqual(batch.prompt_length, 2)
self.assertEqual(batch.tokens.shape[1], 3)
self.assertIsNone(batch.rewards)
self.assertTrue(torch.equal(batch.token_level_rewards[0], torch.tensor([0.1])))
self.assertTrue(torch.equal(batch.token_level_rewards[1], torch.tensor([0.2])))
def test_action_mask_and_logprobs_type(self):
exp = Experience(tokens=[1, 2, 3], logprobs=[0.1, 0.2, 0.3], prompt_length=1)
self.assertIsInstance(exp.tokens, torch.Tensor)
self.assertIsInstance(exp.logprobs, torch.Tensor)
self.assertIsInstance(exp.action_mask, torch.Tensor)
def test_assertions(self):
# prompt_length must be > 0
with self.assertRaises(AssertionError):
Experience(tokens=[1, 2, 3], prompt_length=0)
# tokens must be larger than prompt_length for single-turn
with self.assertRaises(AssertionError):
Experience(tokens=[1, 2], prompt_length=2)
# DPO: tokens must match prompt_length
exp = Experience(tokens=[1, 2], chosen=[3], rejected=[4], prompt_length=1)
exp.prompt_length = 2 # should automatically adjust
def test_hf_datasets_conversion(self):
import torch
from trinity.common.experience import (
Experience,
from_hf_datasets,
to_hf_datasets,
)
exps = [
Experience(
eid=EID(batch=1, task=2, run=3, step=4),
tokens=torch.tensor([1, 2, 3]),
reward=1.0,
logprobs=torch.tensor([0.2, 0.3]),
prompt_length=1,
advantages=None,
returns=None,
info={"key": "value"},
metrics={"accuracy": 0.9},
action_mask=torch.tensor([1, 1]),
),
Experience(
eid=EID(batch=1, task=5, run=6, step=7),
tokens=torch.tensor([4, 5, 6]),
reward=2.0,
logprobs=torch.tensor([0.5]),
prompt_length=2,
advantages=torch.tensor([0.9]),
returns=torch.tensor([0.9]),
info={"key": "value"},
metrics={"accuracy": 0.95},
action_mask=torch.tensor([1]),
),
]
ds = to_hf_datasets(exps)
self.assertEqual(len(ds), 2)
exps2 = from_hf_datasets(ds)
self.assertEqual(len(exps2), 2)
for e1, e2 in zip(exps, exps2):
self.assertTrue(torch.equal(e1.tokens, e2.tokens))
self.assertEqual(e1.reward, e2.reward)
self.assertTrue(torch.equal(e1.logprobs, e2.logprobs))
self.assertEqual(e1.prompt_length, e2.prompt_length)
self.assertTrue(torch.equal(e1.action_mask, e2.action_mask))
self.assertEqual(e1.eid.uid, e2.eid.uid)
self.assertEqual(e1.info, e2.info)
self.assertEqual(e1.metrics, e2.metrics)
if e1.advantages is not None:
self.assertTrue(torch.equal(e1.advantages, e2.advantages))
if e1.returns is not None:
self.assertTrue(torch.equal(e1.returns, e2.returns))
class TestExperienceConversion(unittest.TestCase):
"""Test cases for ExperienceModel"""
def test_experience_model_experience_conversion(self):
"""Test the conversion between Experience and ExperienceModel"""
tokens = torch.tensor([1, 2, 3], dtype=torch.int32)
reward = 0.6
prompt_length = 2
logprobs = torch.tensor([0, 0, 0.1], dtype=torch.float32)
experience = Experience(
tokens=tokens,
reward=reward,
prompt_length=prompt_length,
logprobs=logprobs,
info={"model_version": 0},
)
model = ExperienceModel.from_experience(experience)
new_experience = model.to_experience()
self.assertTrue(torch.equal(new_experience.tokens, tokens))
self.assertEqual(new_experience.prompt_length, prompt_length)
self.assertEqual(new_experience.reward, reward)
self.assertTrue(torch.equal(new_experience.logprobs, logprobs))
self.assertTrue(torch.equal(new_experience.action_mask, experience.action_mask))
def test_batch_conversion(self):
exps = [
Experience(
tokens=torch.tensor([1, 2]),
prompt_length=1,
reward=float(0.1),
logprobs=torch.tensor([0.1]),
advantages=torch.tensor([0.1]),
returns=torch.tensor([0.4]),
),
Experience(
tokens=torch.tensor([1, 2, 3]),
prompt_length=2,
reward=float(0.2),
logprobs=torch.tensor([0.1]),
advantages=torch.tensor([0.3]),
returns=torch.tensor([0.2]),
),
]
batch = Experiences.gather_experiences(exps)
self.assertEqual(batch.batch_size, 2)
self.assertEqual(batch.prompt_length, 2)
prompt_length = batch.prompt_length
for i in range(batch.batch_size):
self.assertEqual(batch.rewards[i], exps[i].reward)
self.assertTrue(
torch.all(
batch.tokens[i][
prompt_length
- exps[i].prompt_length : prompt_length
- exps[i].prompt_length
+ exps[i].tokens.size(0)
]
== exps[i].tokens
)
)
self.assertTrue(
torch.all(
batch.logprobs[i][: exps[i].tokens.size(0) - exps[i].prompt_length]
== exps[i].logprobs
)
)
self.assertTrue(
torch.all(
batch.action_masks[i][: exps[i].tokens.size(0) - exps[i].prompt_length]
== exps[i].action_mask
)
)
self.assertTrue(
torch.all(
batch.advantages[i][: exps[i].tokens.size(0) - exps[i].prompt_length]
== exps[i].advantages
)
)
self.assertTrue(
torch.all(
batch.returns[i][: exps[i].tokens.size(0) - exps[i].prompt_length]
== exps[i].returns
)
)
def test_multiturn_experience_batch_converstion(self):
exps = [
Experience(
tokens=torch.tensor([1, 2, 3, 4, 5, 6]),
reward=float(0.3),
logprobs=torch.tensor([0, 0.1, 0.2, 0.3]),
prompt_length=2,
action_mask=torch.tensor([1, 0, 1, 1]),
advantages=torch.tensor([0.1, 0, 0.2, 0.3]),
returns=torch.tensor([0.5, 0, 0.7, 0.8]),
),
Experience(
tokens=torch.tensor([1, 2, 3, 4]),
reward=float(0.4),
logprobs=torch.tensor([0, 0.1]),
prompt_length=2,
action_mask=torch.tensor([1, 1]),
advantages=torch.tensor([0.2, 0.3]),
returns=torch.tensor([0.6, 0.9]),
),
]
batch = Experiences.gather_experiences(exps)
self.assertEqual(batch.batch_size, 2)
self.assertEqual(batch.prompt_length, 2)
prompt_length = batch.prompt_length
for i in range(batch.batch_size):
self.assertEqual(batch.rewards[i], exps[i].reward)
self.assertTrue(
torch.all(
batch.tokens[i][
prompt_length
- exps[i].prompt_length : prompt_length
- exps[i].prompt_length
+ exps[i].tokens.size(0)
]
== exps[i].tokens
)
)
self.assertTrue(
torch.all(
batch.logprobs[i][: exps[i].tokens.size(0) - exps[i].prompt_length]
== exps[i].logprobs
)
)
self.assertTrue(
torch.all(
batch.action_masks[i][: exps[i].tokens.size(0) - exps[i].prompt_length]
== exps[i].action_mask
)
)
self.assertTrue(
torch.all(
batch.advantages[i][: exps[i].tokens.size(0) - exps[i].prompt_length]
== exps[i].advantages
)
)
self.assertTrue(
torch.all(
batch.returns[i][: exps[i].tokens.size(0) - exps[i].prompt_length]
== exps[i].returns
)
)
def test_dpo_experience_batch_conversion(self):
exps = [
Experience(
tokens=torch.tensor([1, 2]),
chosen=torch.tensor([3, 4]),
rejected=torch.tensor([5, 6]),
),
Experience(
tokens=torch.tensor([7, 8, 9]),
chosen=torch.tensor([10, 11]),
rejected=torch.tensor([12, 13]),
),
]
batch = Experiences.gather_experiences(exps)
self.assertEqual(batch.batch_size, 4)
self.assertEqual(batch.prompt_length, 3)
prompt_length = batch.prompt_length
for i in range(batch.batch_size):
j = i // 2
self.assertTrue(
torch.all(
batch.tokens[i][
prompt_length
- exps[j].prompt_length : prompt_length
- exps[j].prompt_length
+ exps[j].tokens.size(0)
]
== exps[j].tokens
)
)
def test_gather_experiences_with_custom_fields(self):
# test multiple experiences gathering
exps = [
Experience(
tokens=torch.tensor([1, 2]), reward=0.1, prompt_length=1, info={"a": 1.0, "b": 3}
),
Experience(
tokens=torch.tensor([3, 4, 5]), reward=0.2, prompt_length=2, info={"a": 2, "c": 4}
),
]
batch = Experiences.gather_experiences(
exps, custom_fields=[CustomField("a", "a", torch.float32)]
)
self.assertEqual(batch.batch_size, 2)
self.assertEqual(batch.prompt_length, 2)
self.assertEqual(batch.tokens.shape[1], 3)
self.assertEqual(batch.rewards[0], 0.1)
self.assertEqual(batch.rewards[1], 0.2)
self.assertIn("a", batch.custom_fields)
self.assertEqual(batch.a[0], 1.0)
self.assertEqual(batch.a[1], 2.0)
if __name__ == "__main__":
unittest.main()