import boto3 S3_BUCKET = "" S3_KEY = "" def download_cache(): s3_client = boto3.client('s3') t0 = time.perf_counter() try: response = s3_client.get_object(Bucket=S3_BUCKET, Key=S3_KEY) artifact_bytes = response['Body'].read() torch.compiler.load_cache_artifacts(artifact_bytes) print(f"Cache restored. Time: {time.perf_counter()-t0} sec") except: return False return True def upload_cache(): s3_client = boto3.client('s3') artifact_bytes, cache_info = torch.compiler.save_cache_artifacts() s3_client.put_object( Bucket=S3_BUCKET, Key=S3_KEY, Body=artifact_bytes ) if __name__ == '__main__': # specify inductor cache dir inductor_cache_dir = '/tmp/inductor_cache' os.environ['TORCHINDUCTOR_CACHE_DIR'] = inductor_cache_dir # clean up compiler cache torch._dynamo.reset() shutil.rmtree(inductor_cache_dir, ignore_errors=True) # upload the compilation artifacts download_cache() # train the model train() # upload the compilation artifacts upload_cache()