AI-For-Beginners/translations/no/lessons/5-NLP/17-GenerativeNetworks/GenerativePyTorch.ipynb

414 lines
20 KiB
Plaintext

{
"cells": [
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Generative nettverk\n",
"\n",
"Rekurrente nevrale nettverk (RNNs) og deres gatede cellevarianter, som Long Short Term Memory Cells (LSTMs) og Gated Recurrent Units (GRUs), gir en mekanisme for språkmodellering, dvs. de kan lære ordrekkefølge og gi prediksjoner for neste ord i en sekvens. Dette gjør det mulig å bruke RNNs til **generative oppgaver**, som vanlig tekstgenerering, maskinoversettelse og til og med bildetekstgenerering.\n",
"\n",
"I RNN-arkitekturen vi diskuterte i forrige enhet, produserte hver RNN-enhet neste skjulte tilstand som et output. Men vi kan også legge til en annen output til hver rekurrente enhet, som lar oss generere en **sekvens** (som er like lang som den opprinnelige sekvensen). Videre kan vi bruke RNN-enheter som ikke tar imot en input ved hvert steg, men bare tar en initial tilstandsvektor og deretter produserer en sekvens av outputs.\n",
"\n",
"I denne notatboken skal vi fokusere på enkle generative modeller som hjelper oss med å generere tekst. For enkelhets skyld skal vi bygge et **tegnnivå-nettverk**, som genererer tekst bokstav for bokstav. Under trening må vi ta en tekstkorpus og dele den opp i tegnsekvenser.\n"
]
},
{
"cell_type": "code",
"execution_count": 1,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Loading dataset...\n",
"Building vocab...\n"
]
}
],
"source": [
"import torch\n",
"import torchtext\n",
"import numpy as np\n",
"from torchnlp import *\n",
"train_dataset,test_dataset,classes,vocab = load_dataset()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Bygge tegnordforråd\n",
"\n",
"For å bygge et generativt nettverk på tegnnivå, må vi dele opp teksten i individuelle tegn i stedet for ord. Dette kan gjøres ved å definere en annen tokenizer:\n"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Vocabulary size = 82\n",
"Encoding of 'a' is 1\n",
"Character with code 13 is c\n"
]
}
],
"source": [
"def char_tokenizer(words):\n",
" return list(words) #[word for word in words]\n",
"\n",
"counter = collections.Counter()\n",
"for (label, line) in train_dataset:\n",
" counter.update(char_tokenizer(line))\n",
"vocab = torchtext.vocab.vocab(counter)\n",
"\n",
"vocab_size = len(vocab)\n",
"print(f\"Vocabulary size = {vocab_size}\")\n",
"print(f\"Encoding of 'a' is {vocab.get_stoi()['a']}\")\n",
"print(f\"Character with code 13 is {vocab.get_itos()[13]}\")"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"La oss se eksempelet på hvordan vi kan kode teksten fra datasettet vårt:\n"
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"tensor([ 0, 1, 2, 2, 3, 4, 5, 6, 3, 7, 8, 1, 9, 10, 3, 11, 2, 1,\n",
" 12, 3, 7, 1, 13, 14, 3, 15, 16, 5, 17, 3, 5, 18, 8, 3, 7, 2,\n",
" 1, 13, 14, 3, 19, 20, 8, 21, 5, 8, 9, 10, 22, 3, 20, 8, 21, 5,\n",
" 8, 9, 10, 3, 23, 3, 4, 18, 17, 9, 5, 23, 10, 8, 2, 2, 8, 9,\n",
" 10, 24, 3, 0, 1, 2, 2, 3, 4, 5, 9, 8, 8, 5, 25, 10, 3, 26,\n",
" 12, 27, 16, 26, 2, 27, 16, 28, 29, 30, 1, 16, 26, 3, 17, 31, 3, 21,\n",
" 2, 5, 9, 1, 23, 13, 32, 16, 27, 13, 10, 24, 3, 1, 9, 8, 3, 10,\n",
" 8, 8, 27, 16, 28, 3, 28, 9, 8, 8, 16, 3, 1, 28, 1, 27, 16, 6])"
]
},
"execution_count": 3,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"def enc(x):\n",
" return torch.LongTensor(encode(x,voc=vocab,tokenizer=char_tokenizer))\n",
"\n",
"enc(train_dataset[0][1])"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Trening av en generativ RNN\n",
"\n",
"Måten vi skal trene en RNN til å generere tekst på, er som følger. For hvert steg tar vi en sekvens av tegn med lengde `nchars`, og ber nettverket generere neste utgangstegn for hvert inngangstegn:\n",
"\n",
"![Bilde som viser et eksempel på RNN-generering av ordet 'HELLO'.](../../../../../translated_images/no/rnn-generate.56c54afb52f9781d.webp)\n",
"\n",
"Avhengig av det faktiske scenariet, kan vi også ønske å inkludere noen spesialtegn, som *slutt-på-sekvens* `<eos>`. I vårt tilfelle ønsker vi bare å trene nettverket for uendelig tekstgenerering, så vi vil fastsette størrelsen på hver sekvens til å være lik `nchars` tokens. Følgelig vil hvert treningseksempel bestå av `nchars` innganger og `nchars` utganger (som er inngangssekvensen forskjøvet én symbol til venstre). Minibatcher vil bestå av flere slike sekvenser.\n",
"\n",
"Måten vi vil generere minibatcher på, er å ta hver nyhetstekst med lengde `l`, og generere alle mulige inngangs-utgangskombinasjoner fra den (det vil være `l-nchars` slike kombinasjoner). Disse vil danne én minibatch, og størrelsen på minibatchene vil være forskjellig ved hvert treningssteg.\n"
]
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {},
"outputs": [
{
"data": {
"text/plain": [
"(tensor([[ 0, 1, 2, ..., 28, 29, 30],\n",
" [ 1, 2, 2, ..., 29, 30, 1],\n",
" [ 2, 2, 3, ..., 30, 1, 16],\n",
" ...,\n",
" [20, 8, 21, ..., 1, 28, 1],\n",
" [ 8, 21, 5, ..., 28, 1, 27],\n",
" [21, 5, 8, ..., 1, 27, 16]]),\n",
" tensor([[ 1, 2, 2, ..., 29, 30, 1],\n",
" [ 2, 2, 3, ..., 30, 1, 16],\n",
" [ 2, 3, 4, ..., 1, 16, 26],\n",
" ...,\n",
" [ 8, 21, 5, ..., 28, 1, 27],\n",
" [21, 5, 8, ..., 1, 27, 16],\n",
" [ 5, 8, 9, ..., 27, 16, 6]]))"
]
},
"execution_count": 4,
"metadata": {},
"output_type": "execute_result"
}
],
"source": [
"nchars = 100\n",
"\n",
"def get_batch(s,nchars=nchars):\n",
" ins = torch.zeros(len(s)-nchars,nchars,dtype=torch.long,device=device)\n",
" outs = torch.zeros(len(s)-nchars,nchars,dtype=torch.long,device=device)\n",
" for i in range(len(s)-nchars):\n",
" ins[i] = enc(s[i:i+nchars])\n",
" outs[i] = enc(s[i+1:i+nchars+1])\n",
" return ins,outs\n",
"\n",
"get_batch(train_dataset[0][1])"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"La oss definere generatornettverket. Det kan baseres på hvilken som helst rekurrent celle som vi diskuterte i forrige enhet (enkel, LSTM eller GRU). I vårt eksempel vil vi bruke LSTM.\n",
"\n",
"Siden nettverket tar tegn som input, og vokabularstørrelsen er ganske liten, trenger vi ikke et embedding-lag; én-hot-kodet input kan sendes direkte til LSTM-cellen. Men fordi vi sender tegnnummer som input, må vi én-hot-kode dem før vi sender dem til LSTM. Dette gjøres ved å kalle på `one_hot`-funksjonen under `forward`-passet. Utgangskoderen vil være et lineært lag som konverterer skjult tilstand til én-hot-kodet output.\n"
]
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {},
"outputs": [],
"source": [
"class LSTMGenerator(torch.nn.Module):\n",
" def __init__(self, vocab_size, hidden_dim):\n",
" super().__init__()\n",
" self.rnn = torch.nn.LSTM(vocab_size,hidden_dim,batch_first=True)\n",
" self.fc = torch.nn.Linear(hidden_dim, vocab_size)\n",
"\n",
" def forward(self, x, s=None):\n",
" x = torch.nn.functional.one_hot(x,vocab_size).to(torch.float32)\n",
" x,s = self.rnn(x,s)\n",
" return self.fc(x),s"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"Under trening ønsker vi å kunne prøve ut generert tekst. For å gjøre dette, vil vi definere funksjonen `generate` som produserer en utgangsstreng med lengde `size`, som starter med den innledende strengen `start`.\n",
"\n",
"Slik fungerer det: Først sender vi hele startstrengen gjennom nettverket og henter ut tilstanden `s` og neste forutsagte tegn `out`. Siden `out` er én-hot kodet, bruker vi `argmax` for å finne indeksen til tegnet `nc` i vokabularet, og bruker `itos` for å finne det faktiske tegnet og legge det til i den resulterende listen av tegn `chars`. Denne prosessen med å generere ett tegn gjentas `size` ganger for å generere ønsket antall tegn.\n"
]
},
{
"cell_type": "code",
"execution_count": 8,
"metadata": {},
"outputs": [],
"source": [
"def generate(net,size=100,start='today '):\n",
" chars = list(start)\n",
" out, s = net(enc(chars).view(1,-1).to(device))\n",
" for i in range(size):\n",
" nc = torch.argmax(out[0][-1])\n",
" chars.append(vocab.get_itos()[nc])\n",
" out, s = net(nc.view(1,-1),s)\n",
" return ''.join(chars)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"La oss starte treningen! Treningsløkken er nesten den samme som i alle våre tidligere eksempler, men i stedet for nøyaktighet skriver vi ut generert tekst hver 1000. epoke.\n",
"\n",
"Spesiell oppmerksomhet må rettes mot hvordan vi beregner tap. Vi må beregne tap gitt én-hot-kodet output `out`, og forventet tekst `text_out`, som er listen over tegnindekser. Heldigvis forventer `cross_entropy`-funksjonen ikke-normalisert nettverksoutput som første argument, og klassenummer som det andre, noe som er akkurat det vi har. Den utfører også automatisk gjennomsnitt over minibatch-størrelsen.\n",
"\n",
"Vi begrenser også treningen til `samples_to_train` prøver, for å unngå å vente for lenge. Vi oppfordrer deg til å eksperimentere og prøve lengre trening, muligens over flere epoker (i så fall må du lage en ny løkke rundt denne koden).\n"
]
},
{
"cell_type": "code",
"execution_count": 9,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Current loss = 4.398899078369141\n",
"today sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr sr s\n",
"Current loss = 2.161320447921753\n",
"today and to the tor to to the tor to to the tor to to the tor to to the tor to to the tor to to the tor t\n",
"Current loss = 1.6722588539123535\n",
"today and the court to the could to the could to the could to the could to the could to the could to the c\n",
"Current loss = 2.423795223236084\n",
"today and a second to the conternation of the conternation of the conternation of the conternation of the \n",
"Current loss = 1.702607274055481\n",
"today and the company to the company to the company to the company to the company to the company to the co\n",
"Current loss = 1.692358136177063\n",
"today and the company to the company to the company to the company to the company to the company to the co\n",
"Current loss = 1.9722288846969604\n",
"today and the control the control the control the control the control the control the control the control \n",
"Current loss = 1.8705692291259766\n",
"today and the second to the second to the second to the second to the second to the second to the second t\n",
"Current loss = 1.7626899480819702\n",
"today and a security and a security and a security and a security and a security and a security and a secu\n",
"Current loss = 1.5574463605880737\n",
"today and the company and the company and the company and the company and the company and the company and \n",
"Current loss = 1.5620026588439941\n",
"today and the be that the be the be that the be the be that the be the be that the be the be that the be t\n"
]
}
],
"source": [
"net = LSTMGenerator(vocab_size,64).to(device)\n",
"\n",
"samples_to_train = 10000\n",
"optimizer = torch.optim.Adam(net.parameters(),0.01)\n",
"loss_fn = torch.nn.CrossEntropyLoss()\n",
"net.train()\n",
"for i,x in enumerate(train_dataset):\n",
" # x[0] is class label, x[1] is text\n",
" if len(x[1])-nchars<10:\n",
" continue\n",
" samples_to_train-=1\n",
" if not samples_to_train: break\n",
" text_in, text_out = get_batch(x[1])\n",
" optimizer.zero_grad()\n",
" out,s = net(text_in)\n",
" loss = torch.nn.functional.cross_entropy(out.view(-1,vocab_size),text_out.flatten()) #cross_entropy(out,labels)\n",
" loss.backward()\n",
" optimizer.step()\n",
" if i%1000==0:\n",
" print(f\"Current loss = {loss.item()}\")\n",
" print(generate(net))"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"Dette eksempelet genererer allerede ganske god tekst, men det kan forbedres på flere måter:\n",
"\n",
"* **Bedre minibatch-generering**. Måten vi forberedte data for trening på var å generere én minibatch fra én prøve. Dette er ikke ideelt, fordi minibatcher har ulik størrelse, og noen av dem kan ikke engang genereres fordi teksten er mindre enn `nchars`. I tillegg utnytter små minibatcher ikke GPU-en tilstrekkelig. Det ville vært smartere å hente en stor tekstblokk fra alle prøvene, deretter generere alle input-output-par, stokke dem, og lage minibatcher av lik størrelse.\n",
"\n",
"* **Flere lag med LSTM**. Det gir mening å prøve 2 eller 3 lag med LSTM-celler. Som vi nevnte i forrige enhet, trekker hvert lag i LSTM ut visse mønstre fra teksten, og i tilfelle av en generator på tegnnivå kan vi forvente at de lavere LSTM-lagene er ansvarlige for å trekke ut stavelser, mens de høyere lagene håndterer ord og ordkombinasjoner. Dette kan enkelt implementeres ved å sende et parameter for antall lag til LSTM-konstruktøren.\n",
"\n",
"* Du kan også eksperimentere med **GRU-enheter** og se hvilke som gir bedre resultater, samt prøve **forskjellige størrelser på de skjulte lagene**. Et for stort skjult lag kan føre til overtilpasning (f.eks. at nettverket lærer seg teksten nøyaktig), mens et for lite lag kanskje ikke gir gode resultater.\n"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Myk tekstgenerering og temperatur\n",
"\n",
"I den tidligere definisjonen av `generate` valgte vi alltid tegnet med høyest sannsynlighet som neste tegn i den genererte teksten. Dette førte ofte til at teksten \"gikk i sirkel\" mellom de samme tegnsekvensene igjen og igjen, som i dette eksempelet:\n",
"```\n",
"today of the second the company and a second the company ...\n",
"```\n",
"\n",
"Men hvis vi ser på sannsynlighetsfordelingen for neste tegn, kan det hende at forskjellen mellom de høyeste sannsynlighetene ikke er stor, f.eks. ett tegn kan ha sannsynlighet 0.2, et annet - 0.19, osv. For eksempel, når vi ser etter neste tegn i sekvensen '*play*', kan neste tegn like gjerne være et mellomrom eller **e** (som i ordet *player*).\n",
"\n",
"Dette leder oss til konklusjonen at det ikke alltid er \"rettferdig\" å velge tegnet med høyest sannsynlighet, fordi det å velge det nest høyeste fortsatt kan føre til meningsfull tekst. Det er mer fornuftig å **samle** tegn fra sannsynlighetsfordelingen gitt av nettverksutgangen.\n",
"\n",
"Denne samplingen kan gjøres ved hjelp av funksjonen `multinomial`, som implementerer den såkalte **multinomialfordelingen**. En funksjon som implementerer denne **myke** tekstgenereringen er definert nedenfor:\n"
]
},
{
"cell_type": "code",
"execution_count": 10,
"metadata": {
"scrolled": true
},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"--- Temperature = 0.3\n",
"Today and a company and complete an all the land the restrational the as a security and has provers the pay to and a report and the computer in the stand has filities and working the law the stations for a company and with the company and the final the first company and refight of the state and and workin\n",
"\n",
"--- Temperature = 0.8\n",
"Today he oniis its first to Aus bomblaties the marmation a to manan boogot that pirate assaid a relaid their that goverfin the the Cappets Ecrotional Assonia Cition targets it annight the w scyments Blamity #39;s TVeer Diercheg Reserals fran envyuil that of ster said access what succers of Dour-provelith\n",
"\n",
"--- Temperature = 1.0\n",
"Today holy they a 11 will meda a toket subsuaties, engins for Chanos, they's has stainger past to opening orital his thempting new Nattona was al innerforder advan-than #36;s night year his religuled talitatian what the but with Wednesday to Justment will wemen of Mark CCC Camp as Timed Nae wome a leaders\n",
"\n",
"--- Temperature = 1.3\n",
"Today gpone 2.5 fech atcusion poor cocles toparsdorM.cht Line Pamage put 43 his calt lowed to the book, that has authh-the silia rruch ailing to'ory andhes beutirsimi- Aefffive heading offil an auf eacklets is charged evis, Gunymy oy) Mony has it after-sloythyor loveId out filme, the Natabl -Najuntaxiggs \n",
"\n",
"--- Temperature = 1.8\n",
"Today plary, P.slan chly\\401 mardregationly #39;t 8.1Mide) closes ,filtcon alfly playin roven!\\grea.-QFBEP: Iss onfarchQ/itilia CCf Zivesigntwasta orce.-Peul-aw.uicrin of fuglinfsut aftaningwo, MIEX awayew Aice Woiduar Corvagiugge oppo esig ThusBratourid canthly-RyI.co lagitems\\eexciaishes.conBabntusmor I\n",
"\n"
]
}
],
"source": [
"def generate_soft(net,size=100,start='today ',temperature=1.0):\n",
" chars = list(start)\n",
" out, s = net(enc(chars).view(1,-1).to(device))\n",
" for i in range(size):\n",
" #nc = torch.argmax(out[0][-1])\n",
" out_dist = out[0][-1].div(temperature).exp()\n",
" nc = torch.multinomial(out_dist,1)[0]\n",
" chars.append(vocab.get_itos()[nc])\n",
" out, s = net(nc.view(1,-1),s)\n",
" return ''.join(chars)\n",
" \n",
"for i in [0.3,0.8,1.0,1.3,1.8]:\n",
" print(f\"--- Temperature = {i}\\n{generate_soft(net,size=300,start='Today ',temperature=i)}\\n\")"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"Vi har introdusert en ekstra parameter kalt **temperatur**, som brukes til å indikere hvor strengt vi skal holde oss til den høyeste sannsynligheten. Hvis temperaturen er 1.0, gjør vi rettferdig multinomial sampling, og når temperaturen går mot uendelig - blir alle sannsynligheter like, og vi velger neste tegn tilfeldig. I eksempelet nedenfor kan vi observere at teksten blir meningsløs når vi øker temperaturen for mye, og den ligner \"syklisk\" hard-generert tekst når den nærmer seg 0.\n"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"\n---\n\n**Ansvarsfraskrivelse**: \nDette dokumentet er oversatt ved hjelp av AI-oversettelsestjenesten [Co-op Translator](https://github.com/Azure/co-op-translator). Selv om vi streber etter nøyaktighet, vær oppmerksom på at automatiske oversettelser kan inneholde feil eller unøyaktigheter. Det originale dokumentet på sitt opprinnelige språk bør anses som den autoritative kilden. For kritisk informasjon anbefales profesjonell menneskelig oversettelse. Vi er ikke ansvarlige for misforståelser eller feiltolkninger som oppstår ved bruk av denne oversettelsen.\n"
]
}
],
"metadata": {
"interpreter": {
"hash": "16af2a8bbb083ea23e5e41c7f5787656b2ce26968575d8763f2c4b17f9cd711f"
},
"kernelspec": {
"display_name": "Python 3.8.12 ('py38')",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.8.12"
},
"coopTranslator": {
"original_hash": "7673cd150d96c74c6d6011460094efb4",
"translation_date": "2025-08-28T17:28:30+00:00",
"source_file": "lessons/5-NLP/17-GenerativeNetworks/GenerativePyTorch.ipynb",
"language_code": "no"
}
},
"nbformat": 4,
"nbformat_minor": 4
}