Team Ai
Apppublic

GoodWin/Deep-Multi-scale

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
get_data.py116 linesDownload Raw Back to util
1from __future__ import print_function2import os3import tarfile4import requests5from warnings import warn6from zipfile import ZipFile7from bs4 import BeautifulSoup8from os.path import abspath, isdir, join, basename9 10 11class GetData(object):12    """13 14    Download CycleGAN or Pix2Pix Data.15 16    Args:17        technique : str18            One of: 'cyclegan' or 'pix2pix'.19        verbose : bool20            If True, print additional information.21 22    Examples:23        >>> from util.get_data import GetData24        >>> gd = GetData(technique='cyclegan')25        >>> new_data_path = gd.get(save_path='./datasets')  # options will be displayed.26 27    """28 29    def __init__(self, technique='cyclegan', verbose=True):30        url_dict = {31            'pix2pix': 'https://people.eecs.berkeley.edu/~tinghuiz/projects/pix2pix/datasets',32            'cyclegan': 'https://people.eecs.berkeley.edu/~taesung_park/CycleGAN/datasets'33        }34        self.url = url_dict.get(technique.lower())35        self._verbose = verbose36 37    def _print(self, text):38        if self._verbose:39            print(text)40 41    @staticmethod42    def _get_options(r):43        soup = BeautifulSoup(r.text, 'lxml')44        options = [h.text for h in soup.find_all('a', href=True)45                   if h.text.endswith(('.zip', 'tar.gz'))]46        return options47 48    def _present_options(self):49        r = requests.get(self.url)50        options = self._get_options(r)51        print('Options:\n')52        for i, o in enumerate(options):53            print("{0}: {1}".format(i, o))54        choice = input("\nPlease enter the number of the "55                       "dataset above you wish to download:")56        return options[int(choice)]57 58    def _download_data(self, dataset_url, save_path):59        if not isdir(save_path):60            os.makedirs(save_path)61 62        base = basename(dataset_url)63        temp_save_path = join(save_path, base)64 65        with open(temp_save_path, "wb") as f:66            r = requests.get(dataset_url)67            f.write(r.content)68 69        if base.endswith('.tar.gz'):70            obj = tarfile.open(temp_save_path)71        elif base.endswith('.zip'):72            obj = ZipFile(temp_save_path, 'r')73        else:74            raise ValueError("Unknown File Type: {0}.".format(base))75 76        self._print("Unpacking Data...")77        obj.extractall(save_path)78        obj.close()79        os.remove(temp_save_path)80 81    def get(self, save_path, dataset=None):82        """83 84        Download a dataset.85 86        Args:87            save_path : str88                A directory to save the data to.89            dataset : str, optional90                A specific dataset to download.91                Note: this must include the file extension.92                If None, options will be presented for you93                to choose from.94 95        Returns:96            save_path_full : str97                The absolute path to the downloaded data.98 99        """100        if dataset is None:101            selected_dataset = self._present_options()102        else:103            selected_dataset = dataset104 105        save_path_full = join(save_path, selected_dataset.split('.')[0])106 107        if isdir(save_path_full):108            warn("\n'{0}' already exists. Voiding Download.".format(109                save_path_full))110        else:111            self._print('Downloading Data...')112            url = "{0}/{1}".format(self.url, selected_dataset)113            self._download_data(url, save_path=save_path)114 115        return abspath(save_path_full)116