from pyspark.sql import SparkSession, Row from pyspark.sql.functions import udf from pyspark.sql.types import StructType, StructField, StringType, LongType from transformers import AutoTokenizer, AutoModelForSeq2SeqLM # Create a Spark session spark = SparkSession.builder.appName("T5Seq2SeqExample").getOrCreate() # Create an Example Spark DataFrame schema = StructType([ StructField("id", LongType(), nullable=False), StructField("sentence", StringType(), nullable=False) ]) data = [ Row(1, "It is a good test for Spark."), Row(2, "Spark DataFrames are powerful."), Row(3, "LLMs could be very slow."), Row(4, "It is a naive statement.") ] input_df = spark.createDataFrame(data, schema=schema) # Loading t5 Model and Tokenizer model = AutoModelForSeq2SeqLM.from_pretrained("google/flan-t5-small") tokenizer = AutoTokenizer.from_pretrained("google/flan-t5-small") # Defining the Spark UDF def t5_seq2seq_udf(input_text): prompt = f"sentiment of the text: {input_text}" input = tokenizer(prompt, return_tensors="pt") output = model.generate(**input) output_text = tokenizer.decode(output[0], skip_special_tokens=True) return output_text t5_udf = udf(t5_seq2seq_udf, returnType=StringType()) results_df = input_df.withColumn('output_column', t5_udf(input_df['sentence'])) results_df.show(truncate=False)