Skip to content

Commit ebe62ed

Browse files
authored
Make reverse-role formatter pickle safe (#16123)
Signed-off-by: Edresson Casanova <ecasanova@nvidia.com>
1 parent 0d4ac61 commit ebe62ed

2 files changed

Lines changed: 135 additions & 65 deletions

File tree

nemo/collections/common/data/lhotse/cutset.py

Lines changed: 73 additions & 65 deletions
Original file line numberDiff line numberDiff line change
@@ -1206,6 +1206,69 @@ def filter_target_speaker_fn(cut: Cut) -> bool:
12061206
return cuts, is_tarred
12071207

12081208

1209+
def s2s_duplex_reverse_role_for_one_speaker(
1210+
speaker: str | None,
1211+
agent_roles: tuple[str, ...],
1212+
user_roles: tuple[str, ...],
1213+
target_agent_name: str,
1214+
target_user_name: str,
1215+
) -> str | None:
1216+
"""Swap one speaker label for the reverse-role duplex view."""
1217+
if speaker is None:
1218+
return speaker
1219+
1220+
speaker_l = speaker.lower()
1221+
if speaker_l in user_roles:
1222+
return target_agent_name
1223+
if speaker_l in agent_roles:
1224+
return target_user_name
1225+
return speaker
1226+
1227+
1228+
def s2s_duplex_reverse_role_for_one_cut(
1229+
cut: Cut,
1230+
agent_roles: tuple[str, ...],
1231+
user_roles: tuple[str, ...],
1232+
target_agent_name: str,
1233+
target_user_name: str,
1234+
) -> Cut:
1235+
"""Swap speaker roles and source/target audio streams for one duplex cut."""
1236+
new_cut = deepcopy(cut)
1237+
1238+
if getattr(new_cut, "supervisions", None):
1239+
new_sups = []
1240+
for supervision in new_cut.supervisions:
1241+
swapped_supervision = deepcopy(supervision)
1242+
swapped_supervision.speaker = s2s_duplex_reverse_role_for_one_speaker(
1243+
getattr(swapped_supervision, "speaker", None),
1244+
agent_roles=agent_roles,
1245+
user_roles=user_roles,
1246+
target_agent_name=target_agent_name,
1247+
target_user_name=target_user_name,
1248+
)
1249+
new_sups.append(swapped_supervision)
1250+
new_cut.supervisions = new_sups
1251+
1252+
old_recording = new_cut.recording
1253+
old_target_audio = new_cut.target_audio
1254+
old_rec_id = old_recording.id
1255+
old_tar_id = old_target_audio.id
1256+
1257+
new_cut.recording = old_target_audio
1258+
new_cut.target_audio = old_recording
1259+
1260+
if hasattr(new_cut, "duration"):
1261+
new_cut.duration = new_cut.recording.duration
1262+
1263+
assert new_cut.target_audio.id == old_rec_id, f"{new_cut.id}: recording swap failed"
1264+
assert new_cut.recording.id == old_tar_id, f"{new_cut.id}: target_audio swap failed"
1265+
assert new_cut.recording is old_target_audio, f"{new_cut.id}: recording object not swapped"
1266+
assert new_cut.target_audio is old_recording, f"{new_cut.id}: target_audio object not swapped"
1267+
1268+
new_cut.task = "s2s_duplex_reverse_role"
1269+
return new_cut
1270+
1271+
12091272
@data_type_parser(["s2s_duplex_reverse_role"])
12101273
def read_s2s_duplex_reverse_role(config) -> Tuple[CutSet, bool]:
12111274
"""
@@ -1233,74 +1296,19 @@ def read_s2s_duplex_reverse_role(config) -> Tuple[CutSet, bool]:
12331296
"""
12341297
cuts, is_tarred = read_cutset_from_config(config)
12351298

1236-
# Roles coming from config
1237-
agent_roles = config.get("agent_roles", ["agent", "Agent", "Assistant", "assistant"])
1238-
user_roles = config.get("user_roles", ["user", "User"])
1239-
1240-
# Normalize for robust matching
1241-
agent_roles_set = {r.lower() for r in agent_roles}
1242-
user_roles_set = {r.lower() for r in user_roles}
1243-
1244-
# Canonical names you want after swapping
1299+
agent_roles = tuple(r.lower() for r in config.get("agent_roles", ["agent", "Agent", "Assistant", "assistant"]))
1300+
user_roles = tuple(r.lower() for r in config.get("user_roles", ["user", "User"]))
12451301
target_agent_name = config.get("target_agent_name", "agent")
12461302
target_user_name = config.get("target_user_name", "user")
12471303

1248-
def swap_speaker(role: str) -> str:
1249-
"""Swap a given role based on the configured user/agent sets."""
1250-
if role is None:
1251-
return role
1252-
1253-
role_l = role.lower()
1254-
1255-
# user -> agent
1256-
if role_l in user_roles_set:
1257-
return target_agent_name
1258-
1259-
# agent -> user
1260-
if role_l in agent_roles_set:
1261-
return target_user_name
1262-
1263-
# untouched roles (e.g., narrator, system, etc.)
1264-
return role
1265-
1266-
def convert_cut_fn(cut: Cut) -> Cut:
1267-
"""Convert a single cut by swapping supervisions and audio streams."""
1268-
new_cut = deepcopy(cut)
1269-
1270-
# swap supervisions
1271-
if getattr(new_cut, "supervisions", None):
1272-
new_sups = []
1273-
for s in new_cut.supervisions:
1274-
s2 = deepcopy(s)
1275-
s2.speaker = swap_speaker(getattr(s2, "speaker", None))
1276-
new_sups.append(s2)
1277-
new_cut.supervisions = new_sups
1278-
1279-
# swap audio streams
1280-
old_recording = new_cut.recording
1281-
old_target_audio = new_cut.target_audio
1282-
old_rec_id = old_recording.id
1283-
old_tar_id = old_target_audio.id
1284-
1285-
new_cut.recording = old_target_audio
1286-
new_cut.target_audio = old_recording
1287-
1288-
# keep duration consistent
1289-
if hasattr(new_cut, "duration"):
1290-
new_cut.duration = new_cut.recording.duration
1291-
1292-
# Debug assertions
1293-
assert new_cut.target_audio.id == old_rec_id, f"{new_cut.id}: recording swap failed"
1294-
assert new_cut.recording.id == old_tar_id, f"{new_cut.id}: target_audio swap failed"
1295-
1296-
# Optional stronger assertions (object identity)
1297-
assert new_cut.recording is old_target_audio, f"{new_cut.id}: recording object not swapped"
1298-
assert new_cut.target_audio is old_recording, f"{new_cut.id}: target_audio object not swapped"
1299-
1300-
new_cut.task = "s2s_duplex_reverse_role"
1301-
return new_cut
1302-
1303-
cuts = cuts.map(convert_cut_fn)
1304+
convert_fn = partial(
1305+
s2s_duplex_reverse_role_for_one_cut,
1306+
agent_roles=agent_roles,
1307+
user_roles=user_roles,
1308+
target_agent_name=target_agent_name,
1309+
target_user_name=target_user_name,
1310+
)
1311+
cuts = cuts.map(convert_fn)
13041312
return cuts, is_tarred
13051313

13061314

tests/collections/common/test_lhotse_dataloading_duplex.py

Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -292,3 +292,65 @@ def test_data_input_cfg_reverse_role(regular_duplex_s2s_format):
292292
# Ensure the recording streams were swapped
293293
assert cut.recording.id.startswith("rr_target")
294294
assert cut.target_audio.id.startswith("rr_main")
295+
296+
297+
def test_data_input_cfg_reverse_role_multi_config_with_workers(regular_duplex_s2s_format):
298+
config = OmegaConf.create(
299+
{
300+
"multi_config": True,
301+
"sampler_fusion": "randomized_round_robin",
302+
"sampler_weights": {"reverse_a": 0.5, "reverse_b": 0.5},
303+
"seed": 0,
304+
"shard_seed": 0,
305+
"shuffle": True,
306+
"num_workers": 2,
307+
"reverse_a": {
308+
"input_cfg": [
309+
{
310+
"type": "s2s_duplex_reverse_role",
311+
"shar_path": str(regular_duplex_s2s_format),
312+
"weight": 1.0,
313+
"target_agent_name": "swapped_agent",
314+
"target_user_name": "swapped_user",
315+
"tags": {
316+
"dataset_name": "ReverseRoleDataA",
317+
},
318+
},
319+
],
320+
"batch_size": 2,
321+
},
322+
"reverse_b": {
323+
"input_cfg": [
324+
{
325+
"type": "s2s_duplex_reverse_role",
326+
"shar_path": str(regular_duplex_s2s_format),
327+
"weight": 1.0,
328+
"target_agent_name": "swapped_agent",
329+
"target_user_name": "swapped_user",
330+
"tags": {
331+
"dataset_name": "ReverseRoleDataB",
332+
},
333+
},
334+
],
335+
"batch_size": 2,
336+
},
337+
}
338+
)
339+
340+
dl = get_lhotse_dataloader_from_config(config=config, global_rank=0, world_size=1, dataset=Identity())
341+
iterator = iter(dl)
342+
try:
343+
batch = next(iterator)
344+
finally:
345+
if hasattr(iterator, "_shutdown_workers"):
346+
iterator._shutdown_workers()
347+
348+
assert isinstance(batch, lhotse.CutSet)
349+
assert len(batch) > 0
350+
351+
for cut in batch:
352+
assert cut.task == "s2s_duplex_reverse_role"
353+
354+
sups = sorted(cut.supervisions, key=lambda s: s.start)
355+
assert sups[0].speaker == "swapped_agent"
356+
assert sups[1].speaker == "swapped_user"

0 commit comments

Comments
 (0)