Use KMeans clustering to output a dict of cluster labels with their TFIDF-vectorized values

Viewed 29

I have a dataframe looking like this:

       Date         Query                      IsImplicitIntent Country PopularityScore
0      2020-04-01   coronavirus scotland gov        False   United Kingdom  1
1      2020-04-01   origin of coronavirus in china  False   United States   1
2      2020-04-01   wv coronavirus updates          False   United States   1
3      2020-04-01   covid 19 deaths in usa  F.       alse   United States   1
4      2020-04-01   coronavirus karte verbreitung   False   Germany 1
... ... ... ... ... ...
742708  2020-04-30  coronavirus chicago tribune False   United States   1
742709  2020-04-30  corona virus worldometer    False   Australia   1
742710  2020-04-30  coronavirus john hopkins    False   Italy   1
742711  2020-04-30  gov.uk coronavirus update   False   United Kingdom  1
742712  2020-04-30  nys coronavirus update  False   United States   1

I need to use KMeans clustering to output a dict that looks like this:

{
    'Cluster 1': ['India', 'Australia', ..., 'China'],
    'Cluster 0': ['United Kingdom', 'United States', ..., 'France'],
    'Cluster 2': ['Mexico', 'Argentina']
}

So far I have this code

df = pd.read_csv(path, sep='\t')
df = df[df['PopularityScore'] > min_pop_score - 1]
df = df.groupby("Country").filter(lambda x: len(x) >= min_num_qs)

vectorizer = TfidfVectorizer()
query_vectors = vectorizer.fit(df[['Query']])
s = vectorizer.transform(df[df['Country'] == country]['Query'])

num_clusters = 3
kmeans = KMeans(n_clusters=num_clusters, random_state=42).fit(s)

But I'm not sure how to take it to the finish line to get the dict I want. I tried this out:

centroids_partitions = {}
for centr in kmeans.cluster_centers_:
    centroid_label = kmeans.predict([centr])
    partition = []
    for k, v in zip(df['Country'], kmeans.labels_):
        if v == centroid_label:
            partition.append(k)

    centroids_partitions[centroid_label[0]] = partition

print(centroids_partitions)

But the output doesn't match desired:

{{0: ['India', 'Australia', 'Indonesia', 'United Kingdom', 'United Kingdom', 'United States', 'Germany', 'United States', 'Australia', 'Australia', 'Australia', 'United Kingdom', 'Canada', 'Canada', 'Canada', 'Brazil', 'United Kingdom', 'Puerto Rico', 'Canada', 'Australia', 'Puerto Rico', 'Brazil', 'Puerto Rico', 'Puerto Rico', 'Puerto Rico', 'Puerto Rico', 'Australia', 'Puerto Rico', 'Puerto Rico', 'Australia', 'Germany', 'Austria', 'United Kingdom', 'Australia', 'Germany', 'Australia', 'United Kingdom', 'United Kingdom', 'United Kingdom', 'United Kingdom', 'Canada', 'Brazil', 'Italy', 'Germany', 'Italy', 'United Kingdom', 'France', 'Indonesia', 'Austria', 'Canada', 'Mexico', 'France', 'Germany', 'Mexico', 'Mexico', 'Indonesia', 'Germany', 'United States', 'France', 'Mexico', 'Italy', 'Brazil', 'Italy', 'Germany', 'Canada', 'Indonesia', 'Italy', 'United States', 'United States', 'United States', 'Canada', 'Italy', 'Indonesia', 'Austria', 'France', 'United States', 'Canada', 'Germany', 'Australia', 'France', 'United Kingdom', 'Brazil', 'Austria', 'Australia', 'France', 'France', 'Germany', 'Germany', 'Germany', 'Canada', 'Germany', 'United Kingdom', 'Canada', 'United States', 'Brazil', 'Brazil', 'Germany', 'Brazil', 'United Kingdom', 'Italy', 'India', 'Italy', 'Austria', 'United Kingdom', 'Canada', 'United Kingdom', 'Italy', 'United States', 'Germany', 'Germany', 'United Kingdom', 'Australia', 'United Kingdom', 'Canada', 'Germany', 'Indonesia', 'Canada', 'Austria', 'India', 'Brazil', 'India', 'Canada', 'Canada', 'Australia', 'United Kingdom', 'Canada', 'Canada', 'Australia', 'Australia', 'Australia', 'Australia', 'Australia', 'South Africa', 'Canada', 'United Kingdom', 'Canada', 'Italy', 'Germany', 'Brazil', 'Canada', 'Mexico', 'Australia', 'Italy', 'United Kingdom', 'Brazil', 'Argentina', 'France', 'Brazil', 'Italy', 'South Africa', 'Australia', 'Argentina', 'Australia', 'Australia', 'Indonesia', 'Australia', 'India', 'United Kingdom', 'Italy', 'Canada', 'Australia', 'France', 'Brazil', 'United Kingdom', 'India', 'Austria', 'Austria', 'Austria', 'South Africa', 'Argentina', 'Italy', 'France', 'United States', 'Italy', 'South Africa', 'Italy', 'United States', 'Australia', 'Mexico', 'India', 'Canada', 'India', 'Canada', 'South Africa', 'United Kingdom', 'Germany', 'United States', 'Austria', 'South Africa', 'Australia', 'Italy', 'United Kingdom', 'India', 'Australia', 'Canada', 'Argentina', 'Malaysia', 'Malaysia', 'Argentina', 'Malaysia', 'Malaysia', 'Malaysia', 'Malaysia', 'Italy', 'Malaysia', 'France', 'India', 'Italy', 'Malaysia', 'France', 'South Africa', 'Canada', 'Brazil', 'United States', 'France', 'South Africa', 'Canada', 'Italy', 'United Kingdom', 'United Kingdom', 'Canada', 'United Kingdom', 'Italy', 'Canada', 'Canada', 'France', 'China', 'India', 'Canada', 'Australia', 'Australia', 'India', 'Australia', 'Indonesia', 'India', 'Italy', 'France', 'Australia', 'China', 'China', 'France', 'China', 'Canada', 'Australia', 'Australia', 'United States', 'France', 'China', 'France', 'United Kingdom', 'Canada', 'Australia', 'United States', 'France', 'India', 'Austria', 'Argentina', 'Canada', 'Mexico', 'Canada', 'United States', 'France', 'United Kingdom', 'Argentina', 'Argentina', 'United Kingdom', 'United Kingdom', 'France', 'Indonesia', 'United Kingdom', 'France', 'United States', 'India', 'Argentina', 'Australia', 'United Kingdom', 'Australia', 'Austria', 'Brazil', 'Austria', 'Brazil', 'Australia', 'Australia', 'Australia', 'Brazil', 'Canada', 'Argentina', 'Brazil', 'Indonesia', 'Australia', 'Italy', 'Italy', 'Australia', 'France', 'France', 'Australia', 'France', 'United States', 'Canada', 'Canada', 'Australia', 'Canada', 'United Kingdom', 'France', 'Canada', 'Australia', 'United States', 'United States', 'United States', 'France', 'United States', 'Australia', 'Australia', 'Australia', 'United Kingdom', 'Italy', 'India', 'Australia', 'Australia', 'Australia', 'United Kingdom', 'India', 'Italy', 'Italy', 'Italy', 'Australia', 'Australia', 'United Kingdom', 'United States', 'Australia', 'Australia', 'Australia', 'Australia', 'United States', 'Australia', 'Australia', 'United Kingdom', 'Australia', 'Germany', 'India', 'United States', 'India', 'Germany', 'Germany', 'United Kingdom', 'United States', 'United States', 'Germany', 'Germany', 'United States', 'United States', 'China', 'United States', 'Germany', 'Germany', 'Germany', 'Germany', 'Germany', 'Germany', 'Germany', 'China', 'United Kingdom', 'United Kingdom', 'United States', 'China', 'China', 'China', 'China', 'China', 'China', 'Germany', 'United Kingdom', 'United States', 'Argentina', 'Argentina', 'Argentina', 'Argentina', 'United Kingdom', 'Argentina', 'Argentina', 'United Kingdom', 'United Kingdom', 'Argentina', 'Argentina', 'India', 'Argentina', 'Argentina', 'Argentina', 'India', 'Argentina', 'Argentina', 'Argentina', 'United States', 'United Kingdom', 'India', 'United Kingdom', 'United Kingdom', 'United States', 'France', 'United States', 'India', 'United States', 'United Kingdom', 'France', 'India', 'United States', 'United Kingdom', 'United States', 'United States', 'United States', 'United Kingdom', 'India', 'France', 'Canada', 'Canada', 'United States', 'United States', 'United Kingdom', 'France', 'Canada']}`
0 Answers
Related