Jianshu001 commited on
Commit
d81a4e0
·
verified ·
1 Parent(s): 643bd25

Upload pipeline_v2.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. pipeline_v2.py +31 -3
pipeline_v2.py CHANGED
@@ -106,13 +106,17 @@ class PhoneScorerHead(nn.Module):
106
  # ============================================================
107
  # Main Pipeline
108
  # ============================================================
 
 
 
 
109
  class PronunciationAssessorV2:
110
  """Phoneme-level pronunciation assessment using fine-tuned WavLM backbone."""
111
 
112
  def __init__(self, checkpoint_path=None, device=None, pherr_threshold=DEFAULT_PHERR_THRESHOLD):
113
  self.device = device or torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
114
  self.pherr_threshold = pherr_threshold
115
- self._checkpoint_path = checkpoint_path or self._default_checkpoint()
116
  self._backbone = None
117
  self._scorer = None
118
  self._fe_backbone = None
@@ -123,8 +127,32 @@ class PronunciationAssessorV2:
123
  self._g2p = None
124
 
125
  @staticmethod
126
- def _default_checkpoint():
127
- return str(Path(__file__).parent / "wavlm_finetuned.pt")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
128
 
129
  def _load_models(self):
130
  if self._backbone is not None:
 
106
  # ============================================================
107
  # Main Pipeline
108
  # ============================================================
109
+ HF_REPO_ID = "Jianshu001/wavlm-phoneme-scorer"
110
+ HF_CHECKPOINT_FILENAME = "wavlm_finetuned.pt"
111
+
112
+
113
  class PronunciationAssessorV2:
114
  """Phoneme-level pronunciation assessment using fine-tuned WavLM backbone."""
115
 
116
  def __init__(self, checkpoint_path=None, device=None, pherr_threshold=DEFAULT_PHERR_THRESHOLD):
117
  self.device = device or torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
118
  self.pherr_threshold = pherr_threshold
119
+ self._checkpoint_path = checkpoint_path or self._resolve_checkpoint()
120
  self._backbone = None
121
  self._scorer = None
122
  self._fe_backbone = None
 
127
  self._g2p = None
128
 
129
  @staticmethod
130
+ def _resolve_checkpoint():
131
+ """Try local file first, then download from HuggingFace."""
132
+ local_path = Path(__file__).parent / HF_CHECKPOINT_FILENAME
133
+ if local_path.exists():
134
+ return str(local_path)
135
+ print(f"Downloading model from huggingface.co/{HF_REPO_ID}...", file=sys.stderr)
136
+ return hf_hub_download(repo_id=HF_REPO_ID, filename=HF_CHECKPOINT_FILENAME)
137
+
138
+ @classmethod
139
+ def from_pretrained(cls, repo_id=None, device=None, pherr_threshold=DEFAULT_PHERR_THRESHOLD):
140
+ """
141
+ Load model from HuggingFace Hub.
142
+
143
+ Args:
144
+ repo_id: HuggingFace repo (default: Jianshu001/wavlm-phoneme-scorer)
145
+ device: torch device
146
+ pherr_threshold: error detection threshold
147
+
148
+ Example:
149
+ assessor = PronunciationAssessorV2.from_pretrained()
150
+ result = assessor.assess("audio.mp3", "Hello, Peter.")
151
+ """
152
+ repo = repo_id or HF_REPO_ID
153
+ print(f"Downloading model from huggingface.co/{repo}...", file=sys.stderr)
154
+ checkpoint_path = hf_hub_download(repo_id=repo, filename=HF_CHECKPOINT_FILENAME)
155
+ return cls(checkpoint_path=checkpoint_path, device=device, pherr_threshold=pherr_threshold)
156
 
157
  def _load_models(self):
158
  if self._backbone is not None: