Skip to content

Commit 52075f5

Browse files
authored
Add fallback for data downloading (omegalabsinc#64)
* Add fallback for o1 data downloading * Fallback for V1 competition * Fix print logs
1 parent 5622017 commit 52075f5

2 files changed

Lines changed: 91 additions & 6 deletions

File tree

neurons/docker_inference_v2v.py

Lines changed: 59 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,56 @@ def pull_latest_diarization_dataset() -> Optional[Dataset]:
6464

6565
return Dataset.from_dict(overall_dataset)
6666

67+
68+
def pull_latest_diarization_dataset_fallback() -> Optional[Dataset]:
69+
"""Pull latest dataset from HuggingFace."""
70+
try:
71+
os.system("rm -rf ./data_cache/*")
72+
omega_ds_files = huggingface_hub.repo_info(repo_id=HF_DATASET, repo_type="dataset").siblings
73+
recent_files = [
74+
f.rfilename
75+
for f in omega_ds_files if
76+
f.rfilename.startswith(DATA_FILES_PREFIX)
77+
][:MAX_DS_FILES]
78+
79+
if len(recent_files) == 0:
80+
return None
81+
82+
download_config = DownloadConfig(download_desc="Downloading Omega Voice Dataset")
83+
84+
with TemporaryDirectory(dir='./data_cache') as temp_dir:
85+
86+
omega_dataset = load_dataset(HF_DATASET, data_files=recent_files, cache_dir=temp_dir, download_config=download_config)["train"]
87+
omega_dataset.cast_column("audio", Audio(sampling_rate=16000))
88+
89+
omega_dataset = next(omega_dataset.shuffle().iter(batch_size=64))
90+
91+
overall_dataset = {k: [] for k in omega_dataset.keys()}
92+
93+
for i in range(len(omega_dataset['audio'])):
94+
audio_array = omega_dataset['audio'][i]
95+
diar_timestamps_start = np.array(omega_dataset['diar_timestamps_start'][i])
96+
diar_speakers = np.array(omega_dataset['diar_speakers'][i])
97+
98+
if len(set(diar_speakers)) == 1:
99+
continue
100+
101+
for k in omega_dataset.keys():
102+
value = audio_array if k == 'audio' else omega_dataset[k][i]
103+
overall_dataset[k].append(value)
104+
105+
if len(overall_dataset['audio']) >= 8:
106+
break
107+
108+
if len(overall_dataset['audio']) < 1:
109+
return None
110+
111+
return Dataset.from_dict(overall_dataset)
112+
113+
except Exception as e:
114+
bt.logging.error(f"Error pulling dataset: {str(e)}")
115+
return None
116+
67117
def compute_s2s_metrics(hf_repo_id: str, local_dir: str, mini_batch: Dataset, hotkey: str, block, model_tracker, device: str='cuda'):
68118
cleanup_gpu_memory()
69119
log_gpu_memory('before container start')
@@ -201,6 +251,13 @@ def run_v2v_scoring(hf_repo_id: str, hotkey: str, block: int, model_tracker: str
201251
diar_time = time.time()
202252

203253
mini_batch = pull_latest_diarization_dataset()
254+
if mini_batch is None:
255+
bt.logging.info(f"Pulling fallback dataset.")
256+
mini_batch = pull_latest_diarization_dataset_fallback()
257+
if mini_batch is None:
258+
bt.logging.error(f"No diarization dataset found")
259+
return 0
260+
204261
bt.logging.info(f"Time taken for diarization dataset: {time.time() - diar_time:.2f} seconds")
205262

206263
vals = compute_s2s_metrics(
@@ -220,7 +277,7 @@ def run_v2v_scoring(hf_repo_id: str, hotkey: str, block: int, model_tracker: str
220277

221278
if __name__ == "__main__":
222279
for epoch in range(2):
223-
for hf_repo_id in ["shinthet/v1_model", "tezuesh/moshi_general"]:
224-
vals = run_v2v_scoring(hf_repo_id, hotkey=None, block=0, model_tracker=None, local_dir="./model_cache")
280+
for hf_repo_id in ["eggmoo/omega_gQdQiVq", "tezuesh/moshi_general"]:
281+
vals = run_v2v_scoring(hf_repo_id, hotkey=None, block=5268488, model_tracker=None, local_dir="./model_cache")
225282
print(vals)
226283
exit(0)

neurons/docker_model_scoring.py

Lines changed: 32 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -65,6 +65,33 @@ def pull_latest_dataset() -> Optional[Dataset]:
6565
bt.logging.error(f"Error pulling dataset: {str(e)}")
6666
return None
6767

68+
69+
def pull_latest_dataset_fallback() -> Optional[Dataset]:
70+
"""Pull latest dataset from HuggingFace."""
71+
try:
72+
os.system("rm -rf ./data_cache/*")
73+
omega_ds_files = huggingface_hub.repo_info(repo_id=HF_DATASET, repo_type="dataset").siblings
74+
recent_files = [
75+
f.rfilename
76+
for f in omega_ds_files if
77+
f.rfilename.startswith(DATA_FILES_PREFIX)
78+
][:MAX_DS_FILES]
79+
80+
if len(recent_files) == 0:
81+
return None
82+
83+
download_config = DownloadConfig(download_desc="Downloading Omega Multimodal Dataset")
84+
85+
with TemporaryDirectory(dir='./data_cache') as temp_dir:
86+
print("temp_dir", temp_dir)
87+
omega_dataset = load_dataset(HF_DATASET, data_files=recent_files, cache_dir=temp_dir, download_config=download_config)["train"]
88+
omega_dataset = next(omega_dataset.shuffle().iter(batch_size=64))
89+
return omega_dataset
90+
91+
except Exception as e:
92+
bt.logging.error(f"Error pulling dataset: {str(e)}")
93+
return None
94+
6895
def verify_hotkey(hf_repo_id: str, local_dir: str, hotkey: str) -> bool:
6996
"""Verify hotkey matches the one in the repository."""
7097
try:
@@ -209,8 +236,8 @@ def run_o1_scoring(hf_repo_id: str, hotkey: str, block: int, model_tracker: str,
209236
mini_batch = pull_latest_dataset()
210237

211238
if mini_batch is None:
212-
bt.logging.error("Failed to pull dataset")
213-
return
239+
bt.logging.error("Failed to pull latest dataset, trying fallback")
240+
mini_batch = pull_latest_dataset_fallback()
214241

215242
start_time = time.time()
216243
score = compute_model_score(
@@ -227,6 +254,7 @@ def run_o1_scoring(hf_repo_id: str, hotkey: str, block: int, model_tracker: str,
227254
bt.logging.info(f"Score: {score}")
228255
return score
229256
if __name__ == "__main__":
257+
230258
# Example usage
231-
score = run_o1_scoring(hf_repo_id="kiwikiw/o1_2", hotkey=None, block=1, model_tracker=None, local_dir="./model_cache")
232-
print("score", score)
259+
score = run_o1_scoring(hf_repo_id="TFOCUS/mfm_8", hotkey=None, block=1, model_tracker=None, local_dir="./model_cache")
260+
print("score", score)

0 commit comments

Comments
 (0)