liuyimeta/training_data_chat
0
1import sys
2from langchain.vectorstores import FAISS
3from pathlib import Path
4from langchain.text_splitter import CharacterTextSplitter
5from langchain.embeddings import OpenAIEmbeddings
6import pickle
7import faiss
8
9
10def train(files_path):
11 # trainingData = list(Path("training/facts/").glob("**/*.*"))
12 trainingData = list(Path(files_path).glob("**/*.*"))
13 if len(trainingData) < 1:
14 print("The folder training/facts should be populated with at least one .txt or .md file.", file=sys.stderr)
15 return
16
17 data = []
18 for training in trainingData:
19 with open(training, "r", encoding='utf-8') as f:
20 print(f"Add {f.name} to dataset")
21 data.append(f.read())
22
23 textSplitter = CharacterTextSplitter(chunk_size=1000, separator="\n", chunk_overlap=0)
24
25 docs = []
26 for sets in data:
27 docs.extend(textSplitter.split_text(sets))
28
29 store1 = FAISS.from_texts(docs, OpenAIEmbeddings())
30 faiss.write_index(store1.index, "after_training/training.index")
31 store1.index = None
32
33 with open("after_training/faiss.pkl", "wb") as f:
34 pickle.dump(store1, f)
35 return "训练完成"
36 