Views
No views yet
requirements.txt).
All details, including the request format, can be inferred without errors from the code.
The best checkpoint was picked by a maximum ROUGE on Canard conversational QA's ROUGE.1import datasets
2
3canard_train_augm = datasets.load_dataset("gaussalgo/Canard_Wiki-augmented", split="train")
4canard_test_augm = datasets.load_dataset("gaussalgo/Canard_Wiki-augmented", split="test")
5
6canard_df = canard_train_augm.to_pandas()
7canard_test_df = canard_train_augm.to_pandas()
8
9
10### Curation of seq2seq input contexts and labels
11import random
12
13def input_context_from_sample(row: dict, max_length=5) -> str:
14 context = "Previous conversation:"
15 context += "\nQuestion: "
16 context += ", ".join(row["History"][:3])
17 for i in range(3, len(row["History"]), 2):
18 context += "\nAnswer: "
19 context += row["History"][i]
20 if i+1 < len(row["History"]):
21 context += "\nQuestion: "
22 context += row["History"][i+1]
23
24 context += "\n\nCurrent Question: "
25 context += row["Question"]
26
27 context += "\nSearch results:"
28 all_contexts = row["retrieved_contexts"].tolist()[:max_length-1] + [row["true_contexts"]]
29 random.shuffle(all_contexts)
30
31 for i, search_result in enumerate(all_contexts):
32 context += "\n[%s]: " % (i+1)
33 context += search_result.replace("CANNOTANSWER", "")
34
35 context += "\nCurrent Answer: "
36 return context
37
38def rephrasing_context_from_sample(row: dict) -> str:
39 context = "Previous conversation:"
40 context += "\nQuestion: "
41 context += ", ".join(row["History"][:3])
42 for i in range(3, len(row["History"]), 2):
43 context += "\nAnswer: "
44 context += row["History"][i]
45 if i+1 < len(row["History"]):
46 context += "\nQuestion: "
47 context += row["History"][i+1]
48
49 context += "\n\nCurrent Question: "
50 context += row["Question"]
51
52 context += "\nMore specific question: "
53 return context
54
55def hotpotqa_context(row: dict) -> str:
56 context = "Current Question: "
57 context += row["question"]
58
59 context += "\nSearch results:"
60 all_contexts = [" ".join(context) for context in row["context"]["sentences"]]
61
62 for i, search_result in enumerate(all_contexts):
63 context += "\n[%s]: " % (i+1)
64 context += search_result.replace("CANNOTANSWER", "")
65
66 context += "\nCurrent Answer: "
67 return context
68
69# Conversational QA sequences
70input_texts = canard_df.apply(lambda row: input_context_from_sample(row), axis=1).values
71input_val_texts = canard_test_df.iloc[:200].apply(lambda row: input_context_from_sample(row), axis=1).values
72
73too_long_index = [len(t) > 20000 for t in input_texts]
74input_texts = [t for i, t in enumerate(input_texts) if not too_long_index[i]]
75# print(too_long_index)
76print("training on %s samples" % len(input_texts))
77
78labels = canard_df.answer.apply(lambda ans: "No answer" if ans == "CANNOTANSWER" else ans).values
79labels = [l for i, l in enumerate(labels) if not too_long_index[i]]
80val_labels = canard_test_df.answer.apply(lambda ans: "No answer" if ans == "CANNOTANSWER" else ans).values
81
82# Rephrasing sequences
83rephrasing_inputs = canard_df.apply(lambda row: rephrasing_context_from_sample(row), axis=1).values
84rephrasing_val_inputs = canard_test_df.apply(lambda row: rephrasing_context_from_sample(row), axis=1).values
85
86rephrasing_labels = canard_df.Rewrite.values
87rephrasing_val_labels = canard_test_df.Rewrite.values
88
89# HotpotQA sequences
90hotpot_train = datasets.load_dataset("hotpot_qa", "distractor")["train"]
91hotpot_val = datasets.load_dataset("hotpot_qa", "distractor")["validation"]
92
93hotpot_inputs = hotpot_train.to_pandas().apply(hotpotqa_context, axis=1)
94hotpot_val_inputs = hotpot_val.to_pandas().apply(hotpotqa_context, axis=1)
95too_long_index = [len(t) > 20000 for t in hotpot_inputs]
96
97hotpot_inputs = [t for i, t in enumerate(hotpot_inputs) if not too_long_index[i]]
98hotpot_answers = [t for i, t in enumerate(hotpot_train["answer"]) if not too_long_index[i]]
99
100# Training routine
101# see Adaptor's homepage for details:
102# https://github.com/gaussalgo/adaptor
103
104# Base model
105from adaptor.lang_module import LangModule
106lang_module = LangModule("google/t5-large-lm-adapt")
107
108from adaptor.evaluators.generative import ROUGE, BLEU
109
110# Evaluations
111evaluators = [BLEU(), ROUGE(decides_convergence=True)]
112
113# Objectives
114from adaptor.objectives.seq2seq import Sequence2Sequence
115
116seq_qa = Sequence2Sequence(lang_module,
117 texts_or_path=input_texts,
118 labels_or_path=labels,
119 val_texts_or_path=input_val_texts,
120 val_labels_or_path=val_labels,
121 batch_size=4,
122 val_evaluators=evaluators,
123 objective_id="Canard")
124
125seq_additional_qa = Sequence2Sequence(lang_module,
126 texts_or_path=hotpot_inputs,
127 labels_or_path=hotpot_answers,
128 val_texts_or_path=hotpot_val_inputs[:200],
129 val_labels_or_path=hotpot_val["answer"][:200],
130 batch_size=4,
131 val_evaluators=evaluators,
132 objective_id="HotpotQA",
133 share_other_objective_head=seq_qa)
134
135seq_rephrasing = Sequence2Sequence(lang_module,
136 texts_or_path=rephrasing_inputs,
137 labels_or_path=rephrasing_labels,
138 val_texts_or_path=rephrasing_val_inputs[:200],
139 val_labels_or_path=rephrasing_val_labels[:200],
140 batch_size=4,
141 val_evaluators=evaluators,
142 objective_id="rephrasing",
143 share_other_objective_head=seq_qa)
144
145# Training schedule & arguments
146from adaptor.utils import AdaptationArguments, StoppingStrategy
147
148training_arguments = AdaptationArguments(output_dir="checkpoints-chatbot",
149 learning_rate=5e-5,
150 stopping_strategy=StoppingStrategy.ALL_OBJECTIVES_CONVERGED,
151 stopping_patience=8,
152 save_total_limit=8,
153 do_train=True,
154 do_eval=True,
155 bf16=True,
156 warmup_steps=1000,
157 gradient_accumulation_steps=8,
158 logging_steps=10,
159 eval_steps=200,
160 save_steps=1000,
161 num_train_epochs=10,
162 evaluation_strategy="steps")
163from adaptor.schedules import ParallelSchedule
164from adaptor.adapter import Adapter
165
166schedule = ParallelSchedule(objectives=[seq_qa, seq_additional_qa, seq_rephrasing],
167 args=training_arguments)
168adapter = Adapter(lang_module, schedule, args=training_arguments)
169adapter.train() # Training for 63k updates