Source code for image_suggestions.commons

#!/usr/bin/env python3
# -*- coding: utf-8 -*-

"""Build full and delta datasets of
`weighted tags <https://wikitech.wikimedia.org/wiki/Search/WeightedTags#Querying_the_data>`_
for Commons' search index.

Images can receive tags from 3 sources,
as output by :mod:`image_suggestions.wikidata_and_lead_images`:

- Wikidata `image <https://www.wikidata.org/wiki/Property:P18>`_ property
- Wikidata `Commons category <https://www.wikidata.org/wiki/Property:P373>`_ property
- Wikipedia article's lead image

The full dataset is stored in the :const:`image_suggestions.shared.SEARCH_INDEX_FULL_TABLE`,
and the delta in the :const:`image_suggestions.shared.SEARCH_INDEX_DELTA_TABLE`
`Hive <https://hive.apache.org/>`_ table of Wikimedia Foundation's
`Analytics Data Lake <https://wikitech.wikimedia.org/wiki/Analytics/Data_Lake>`_.
"""

import argparse

from typing import Tuple

from pyspark.sql import DataFrame, SparkSession
from pyspark.sql import functions as F

from image_suggestions import shared


FILE_NAMESPACE_VALUE = 6


[docs] def build_weighted_tags(wd_data: DataFrame, li_data: DataFrame) -> DataFrame: """Build the full state of a Commons search index’s weighted tags dataset. :param wd_data: a dataframe of Wikidata claims as output by :func:`shared._load_wikidata` :param li_data: a dataframe of Wikipedia article lead images as output by :func:`shared.load_lead_images` :return: the dataframe of weighted tags """ return ( wd_data.withColumn('value', F.concat_ws('|', wd_data.item_id, wd_data.score)) .select('page_id', shared.TAG_COLNAME, 'item_id', shared.SCORE_COLNAME, 'value') .union( li_data.withColumn( 'value', F.concat_ws('|', li_data.item_id, li_data.score) ).select( 'page_id', shared.TAG_COLNAME, 'item_id', shared.SCORE_COLNAME, 'value' ) ) .orderBy('page_id', shared.TAG_COLNAME, 'value') .groupBy('page_id', shared.TAG_COLNAME) # `collect_set` doesn't preserve order, so use `collect_list` with `array_distinct` instead .agg(F.array_distinct(F.collect_list('value')).alias(shared.VALUES_COLNAME)) .withColumn(shared.WIKIID_COLNAME, F.lit(shared.COMMONSWIKI_VALUE)) .withColumn(shared.PAGE_NAMESPACE_COLNAME, F.lit(FILE_NAMESPACE_VALUE)) .select( shared.WIKIID_COLNAME, shared.PAGE_NAMESPACE_COLNAME, 'page_id', shared.TAG_COLNAME, shared.VALUES_COLNAME, ) .orderBy( shared.WIKIID_COLNAME, shared.PAGE_NAMESPACE_COLNAME, 'page_id', shared.TAG_COLNAME, ) )
[docs] def write_weighted_tags( spark: SparkSession, tags: DataFrame, hive_db: str, snapshot: str, coalesce: int, delta_threshold: int, ) -> Tuple[DataFrame, DataFrame]: """Write the full and delta datasets to Hive. Don't write the delta if its row count is greater than a given threshold. :param spark: an active Spark session :param tags: a dataframe of weighted tags as returned by :func:`get_commonswiki_file_data` :param hive_db: an output Hive database name :param snapshot: a ``YYYY-MM-DD`` date :param coalesce: an integer to control the amount of files per output partition. A higher value implies more files but a faster and lighter execution :param delta_threshold: an integer row count threshold that determines whether to include the delta :return: the full and delta dataframes """ full = shared.save_search_index_full(tags, hive_db, snapshot, coalesce) target_tags = [shared.P18_TAG, shared.P373_TAG, shared.LEAD_IMAGE_TAG] delta = shared.build_search_index_delta( spark, tags, target_tags, snapshot, is_commons=True ) delta_count = delta.where(F.col(shared.WIKIID_COLNAME) == 'commonswiki').count() if delta_count <= delta_threshold: shared.save_table(delta, hive_db, shared.SEARCH_INDEX_DELTA_TABLE, coalesce) return full, delta
def parse_args() -> argparse.Namespace: # pragma: no cover description = ( "Build full and delta datasets of weighted tags for Commons' search index" ) parser = shared.build_base_arg_parser(description) parser.add_argument( 'delta_threshold', metavar='N', type=int, help='Row count threshold above which the Commons search index delta will be discarded', ) return parser.parse_args() def main(args: argparse.Namespace, spark: SparkSession) -> None: hive_db = args.hive_db snapshot = args.snapshot coalesce = args.coalesce delta_threshold = args.delta_threshold wikidata = shared.load_wikidata(spark, hive_db, snapshot) lead_images = shared.load_lead_images(spark, hive_db, snapshot) weighted_tags = build_weighted_tags(wikidata, lead_images) write_weighted_tags( spark, weighted_tags, hive_db, snapshot, coalesce, delta_threshold ) if __name__ == '__main__': args = parse_args() spark = shared.build_spark_session() # Pass 'dev' to switch environments main(args, spark) spark.stop()