Download sanity_check.py from batiai/batisay-ko-base: direct link, hf CLI and curl.
- Browser
- Download file 9.94 kB
-
https://huggingface.co/batiai/batisay-ko-base/resolve/main/sanity_check.py
- Command line
-
hf download hf://batiai/batisay-ko-base/sanity_check.py
-
curl -L -o sanity_check.py https://huggingface.co/batiai/batisay-ko-base/resolve/main/sanity_check.py
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() | |