batisay-ko-base / sanity_check.py
hero775's picture
Add sanity_check.py for environment verification
4459384 verified
Raw History Blame Contribute Delete
9.94 kB
#!/usr/bin/env python3
"""BatiSay-ko sanity check — 환경별 자동 검증.
사용법:
python sanity_check.py [--format auto|safetensors|ct2|mlx|ggml]
python sanity_check.py --audio /path/to/your/audio.wav
목적:
1. 모델 다운로드 후 환경에서 정상 동작 확인
2. Mac brew whisper.cpp 같은 토 issue 즉시 검출
3. 사용자에게 명확한 pass / fail report
"""
import sys, os, time, json, argparse, platform, urllib.request
from pathlib import Path
# 기본 sample audio + expected substring (10초 한국어, 자체 제공)
DEFAULT_SAMPLE_URL = "https://huggingface.co/batiai/batisay-ko-base/resolve/main/sanity_sample.wav"
EXPECTED_SUBSTRING = "안녕하십니까" # 거의 모든 어닝콜 / 회의 시작 표현
EXPECTED_MIN_CHARS = 30 # 10초 audio 면 최소 30자 transcribe 기대
def detect_platform():
"""Mac M-series / Intel / Linux 자동 감지."""
system = platform.system()
machine = platform.machine()
if system == 'Darwin':
if machine == 'arm64':
return 'mac-apple-silicon'
else:
return 'mac-intel'
elif system == 'Linux':
return 'linux'
else:
return f'{system}-{machine}'
def recommended_format(plat):
"""플랫폼별 권장 format."""
if plat == 'mac-apple-silicon':
return ['coreml', 'mlx', 'ct2', 'safetensors'] # 우선 순위
elif plat == 'mac-intel':
return ['ggml', 'ct2', 'safetensors']
elif plat == 'linux':
return ['ct2', 'safetensors', 'ggml']
else:
return ['safetensors']
def try_safetensors(model_path, audio_path):
"""HF transformers 로 검증."""
try:
import torch
from transformers import pipeline
pipe = pipeline('automatic-speech-recognition', model=model_path,
device='cuda:0' if torch.cuda.is_available() else 'cpu',
chunk_length_s=30)
if pipe.tokenizer.pad_token_id is None:
pipe.tokenizer.pad_token_id = pipe.model.config.eos_token_id
t0 = time.time()
out = pipe(str(audio_path), generate_kwargs={'language': 'korean', 'task': 'transcribe'})
elapsed = time.time() - t0
text = out.get('text', '').strip()
return {'success': True, 'text': text, 'elapsed_s': round(elapsed, 2)}
except Exception as e:
return {'success': False, 'error': str(e)[:200]}
def try_ct2(model_path, audio_path):
"""faster-whisper (CTranslate2) 로 검증."""
try:
from faster_whisper import WhisperModel
ct2_dir = Path(model_path) / 'ct2'
if not ct2_dir.exists():
ct2_dir = Path(model_path) # local download case
model = WhisperModel(str(ct2_dir), compute_type='int8')
t0 = time.time()
segments, _ = model.transcribe(str(audio_path), language='ko')
text = ''.join(s.text for s in segments).strip()
elapsed = time.time() - t0
return {'success': True, 'text': text, 'elapsed_s': round(elapsed, 2)}
except Exception as e:
return {'success': False, 'error': str(e)[:200]}
def try_mlx(model_path, audio_path):
"""mlx-whisper (Apple Silicon) 로 검증."""
try:
import mlx_whisper
mlx_dir = Path(model_path) / 'mlx'
if not mlx_dir.exists():
mlx_dir = Path(model_path) # already mlx repo
t0 = time.time()
result = mlx_whisper.transcribe(str(audio_path),
path_or_hf_repo=str(mlx_dir),
language='ko')
text = result.get('text', '').strip()
elapsed = time.time() - t0
return {'success': True, 'text': text, 'elapsed_s': round(elapsed, 2)}
except ImportError:
return {'success': False, 'error': 'mlx-whisper not installed (pip install mlx-whisper)'}
except Exception as e:
return {'success': False, 'error': str(e)[:200]}
def try_ggml(model_path, audio_path):
"""whisper.cpp ggml 로 검증."""
import subprocess, shutil
# whisper-cli 찾기
cli = shutil.which('whisper-cli') or shutil.which('whisper-cpp') or shutil.which('main')
if not cli:
return {'success': False, 'error': 'whisper-cli not found in PATH (install whisper.cpp)'}
ggml_dir = Path(model_path) / 'ggml'
if not ggml_dir.exists():
ggml_dir = Path(model_path)
ggml_files = list(ggml_dir.glob('*-q5_0.bin')) or list(ggml_dir.glob('*-fp16.bin'))
if not ggml_files:
return {'success': False, 'error': 'No ggml binary found in model dir'}
try:
t0 = time.time()
out = subprocess.run([cli, '-m', str(ggml_files[0]), '-f', str(audio_path),
'-l', 'ko', '--no-prints', '--output-txt',
'--beam-size', '1', '--best-of', '1'],
capture_output=True, text=True, timeout=120)
elapsed = time.time() - t0
# 결과는 .txt 파일에
txt_path = Path(str(audio_path)).with_suffix('.txt')
text = txt_path.read_text(encoding='utf-8').strip() if txt_path.exists() else out.stdout.strip()
# Mac brew 1.8.4 known issue 패턴 검출
if 'Q?' in text and text.count('Q?') > 3:
return {'success': False, 'text': text,
'error': 'Hallucination detected (Q?Q?Q? pattern). See Known Issues — use master build.'}
return {'success': True, 'text': text, 'elapsed_s': round(elapsed, 2)}
except subprocess.TimeoutExpired:
return {'success': False, 'error': 'whisper-cli timeout (120s)'}
except Exception as e:
return {'success': False, 'error': str(e)[:200]}
def verify_result(result, audio_seconds):
"""결과 검증 — expected substring + 최소 글자 + RTF."""
if not result.get('success'):
return ('FAIL', result.get('error', 'unknown'))
text = result.get('text', '')
issues = []
if EXPECTED_SUBSTRING not in text:
issues.append(f'expected substring "{EXPECTED_SUBSTRING}" missing')
if len(text) < EXPECTED_MIN_CHARS:
issues.append(f'output too short ({len(text)} < {EXPECTED_MIN_CHARS} chars)')
rtf = result.get('elapsed_s', 0) / audio_seconds if audio_seconds > 0 else 0
if rtf > 5.0:
issues.append(f'very slow (RTF {rtf:.1f}, expected <2)')
if issues:
return ('WARN', '; '.join(issues))
return ('PASS', f'RTF {rtf:.2f}, {len(text)} chars')
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--model', default='.', help='Model directory (current dir if downloaded)')
ap.add_argument('--format', default='auto', help='auto|safetensors|ct2|mlx|ggml|all')
ap.add_argument('--audio', help='Custom audio file (default: download sanity sample)')
args = ap.parse_args()
# Audio 준비
if args.audio:
audio = Path(args.audio)
if not audio.exists():
print(f"❌ audio file not found: {audio}", file=sys.stderr); sys.exit(1)
else:
audio = Path('/tmp/batisay_sanity_sample.wav')
if not audio.exists():
print(f" downloading sample audio...", file=sys.stderr)
try:
urllib.request.urlretrieve(DEFAULT_SAMPLE_URL, str(audio))
except Exception as e:
print(f" ⚠ download failed: {e}", file=sys.stderr)
print(f" Please provide --audio path manually", file=sys.stderr)
sys.exit(1)
# Audio 길이
try:
import soundfile as sf
info = sf.info(str(audio))
audio_seconds = info.duration
except Exception:
audio_seconds = 10.0 # default
print(f"\n=== BatiSay-ko Sanity Check ===")
print(f" audio: {audio} ({audio_seconds:.1f}s)")
print(f" expected: '{EXPECTED_SUBSTRING}...' (Korean)")
# Platform 감지
plat = detect_platform()
print(f" platform: {plat}")
# Format 결정
if args.format == 'auto':
formats = recommended_format(plat)
elif args.format == 'all':
formats = ['safetensors', 'ct2', 'mlx', 'ggml']
else:
formats = [args.format]
print(f" formats to test: {formats}\n")
# 각 format 시도
results = {}
for fmt in formats:
print(f"--- Testing {fmt} ---", file=sys.stderr)
if fmt == 'safetensors':
r = try_safetensors(args.model, audio)
elif fmt == 'ct2':
r = try_ct2(args.model, audio)
elif fmt == 'mlx':
r = try_mlx(args.model, audio)
elif fmt == 'ggml':
r = try_ggml(args.model, audio)
else:
print(f" unknown format: {fmt}"); continue
status, msg = verify_result(r, audio_seconds)
symbol = {'PASS': '✅', 'WARN': '⚠️', 'FAIL': '❌'}.get(status, '?')
text_preview = (r.get('text','')[:80] + '...') if r.get('text') else r.get('error', '')[:80]
print(f"{symbol} {fmt:>12s}: {status:5s} {msg}")
print(f" output: {text_preview}")
results[fmt] = {'status': status, 'message': msg, 'result': r}
# 종합
print(f"\n=== Summary ===")
pass_count = sum(1 for v in results.values() if v['status'] == 'PASS')
total = len(results)
print(f" {pass_count} / {total} formats verified")
if pass_count == 0:
print(f" ❌ No format worked. Check Known Issues in README.")
sys.exit(1)
elif pass_count < total:
print(f" ⚠️ Partial success — see warnings above")
else:
print(f" ✅ All formats working in your environment!")
# JSON 결과 저장 (옵션)
Path('/tmp/batisay_sanity_report.json').write_text(
json.dumps({'platform': plat, 'audio_seconds': audio_seconds,
'formats': results}, ensure_ascii=False, indent=2, default=str))
print(f" detailed report: /tmp/batisay_sanity_report.json")
if __name__ == '__main__':
main()