(embeddings, max_clusters=50, random_state=0, rel_tol=1e-3)
| 211 | |
| 212 | |
| 213 | def get_optimal_clusters(embeddings, max_clusters=50, random_state=0, rel_tol=1e-3): |
| 214 | max_clusters = min(len(embeddings), max_clusters) |
| 215 | n_clusters = np.arange(1, max_clusters) |
| 216 | bics = [] |
| 217 | prev_bic = float('inf') |
| 218 | for n in tqdm(n_clusters): |
| 219 | bic = fit_gaussian_mixture(n, embeddings, random_state) |
| 220 | # print(bic) |
| 221 | bics.append(bic) |
| 222 | # early stop |
| 223 | if (abs(prev_bic - bic) / abs(prev_bic)) < rel_tol: |
| 224 | break |
| 225 | prev_bic = bic |
| 226 | optimal_clusters = n_clusters[np.argmin(bics)] |
| 227 | return optimal_clusters |
| 228 | |
| 229 | |
| 230 | def GMM_cluster(embeddings: np.ndarray, threshold: float, random_state: int = 0,cluster_size: int = 20): |
no test coverage detected